NormalCable/brainrot-detector-api
0
1"""
2=============================================================================
3BRAINROT DETECTOR — FASTAPI SERVER
4=============================================================================
5Setup:
6 pip install fastapi uvicorn python-multipart torch torchvision
7 pip install opencv-python librosa transformers==4.40.0 tensorflow
8
9Run:
10 uvicorn api_server:app --host 0.0.0.0 --port 8000 --reload
11
12API Docs (auto-generated):
13 http://localhost:8000/docs
14
15Endpoints:
16 POST /predict — classify a single uploaded video
17 POST /predict/ensemble — classify using all fold models (more accurate)
18 GET /health — check server status
19 GET /model/info — get loaded model info
20=============================================================================
21"""
22
23import os
24import gc
25import re
26import uuid
27import time
28import shutil
29import tempfile
30import warnings
31warnings.filterwarnings("ignore")
32
33import asyncio
34
35import numpy as np
36import torch
37import torch.nn as nn
38import torchvision.models as models
39import torchvision.transforms as transforms
40
41from fastapi import FastAPI, File, UploadFile, HTTPException, Form
42from fastapi.middleware.cors import CORSMiddleware
43from pydantic import BaseModel
44from typing import Optional, List
45
46# Video downloader engine
47from video_downloader import get_scraper, detect_platform
48
49# ── CONFIG — update to match your setup ─────────────────────────────────────
50VISUAL_FEAT_DIM = 1280
51AUDIO_FEAT_DIM = 2048
52TEXT_FEAT_DIM = 768
53FUSION_DIM = 256
54NUM_CLASSES = 2
55MAX_TEXT_LEN = 128
56N_MELS = 128
57AUDIO_MAX_FRAMES= 300
58IMG_SIZE = 224
59MAX_FRAMES = 16
60
61DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
62
63
64# ── MODEL VERSIONS ──────────────────────────────────────────────────────────
65MODEL_VERSIONS = {
66 "default": {
67 "name": "Default (Full Dataset)",
68 "checkpoint_dir": "./model1_checkpoints",
69 "n_folds": 3,
70 },
71 "no_yt": {
72 "name": "No YouTube",
73 "checkpoint_dir": "./model2_checkpoints",
74 "n_folds": 3,
75 },
76 "model3": {
77 "name": "Model 3 (V3 Architecture)",
78 "checkpoint_dir": "./model3_checkpoints",
79 "n_folds": 3,
80 },
81 "model4": {
82 "name": "Model 4 (Refined Cross-Attention)",
83 "checkpoint_dir": "./model4_checkpoints",
84 "n_folds": 3,
85 },
86}
87DEFAULT_VERSION = "model4"
88
89# State for tracking background tasks
90ACTIVE_TASKS_STATE = {}
91
92
93# =============================================================================
94# MODEL ARCHITECTURE (must match training)
95# =============================================================================
96
97class ModalityProjector(nn.Module):
98 def __init__(self, in_dim, out_dim, dropout=0.3):
99 super().__init__()
100 self.net = nn.Sequential(
101 nn.Linear(in_dim, out_dim),
102 nn.LayerNorm(out_dim),
103 nn.GELU(),
104 nn.Dropout(dropout),
105 )
106
107 def forward(self, x):
108 if x.dim() == 3:
109 x = x.mean(dim=1)
110 return self.net(x)
111
112
113class AttentionFusion(nn.Module):
114 def __init__(self, fusion_dim, num_modalities=3):
115 super().__init__()
116 self.attn = nn.Linear(fusion_dim * num_modalities, num_modalities)
117
118 def forward(self, feats):
119 cat = torch.cat(feats, dim=-1)
120 weights = torch.softmax(self.attn(cat), dim=-1)
121 self.last_weights = weights
122 stacked = torch.stack(feats, dim=1)
123 fused = (weights.unsqueeze(-1) * stacked).sum(dim=1)
124 return fused
125
126
127class BrainrotModelV1(nn.Module):
128 def __init__(self):
129 super().__init__()
130 self.visual_proj = ModalityProjector(VISUAL_FEAT_DIM, FUSION_DIM)
131 self.audio_proj = ModalityProjector(AUDIO_FEAT_DIM, FUSION_DIM)
132 self.text_proj = ModalityProjector(TEXT_FEAT_DIM, FUSION_DIM)
133 self.fusion = AttentionFusion(FUSION_DIM, num_modalities=3)
134 self.classifier = nn.Sequential(
135 nn.Linear(FUSION_DIM, FUSION_DIM // 2),
136 nn.GELU(),
137 nn.Dropout(0.3),
138 nn.Linear(FUSION_DIM // 2, NUM_CLASSES),
139 )
140
141 def forward(self, visual, audio, text):
142 v = self.visual_proj(visual)
143 a = self.audio_proj(audio)
144 t = self.text_proj(text)
145 return self.classifier(self.fusion([v, a, t]))
146
147class BrainrotModelV3(nn.Module):
148 def __init__(self):
149 super().__init__()
150 fd = 256
151 self.visual_proj = ModalityProjector(2560, fd)
152 self.audio_proj = ModalityProjector(384, fd)
153 self.text_proj = ModalityProjector(3, fd)
154 self.fusion = AttentionFusion(fd)
155 self.classifier = nn.Sequential(
156 nn.Linear(fd, fd // 2), nn.GELU(),
157 nn.Dropout(0.2), nn.Linear(fd // 2, 1)
158 )
159 def forward(self, v, a, t):
160 return self.classifier(self.fusion([
161 self.visual_proj(v), self.audio_proj(a), self.text_proj(t)
162 ]))
163
164class CrossAttentionFusion(nn.Module):
165 def __init__(self, fusion_dim, num_modalities=3):
166 super().__init__()
167 self.modality_embed = nn.Parameter(torch.randn(1, num_modalities, fusion_dim))
168 self.weights = nn.Parameter(torch.ones(num_modalities))
169 self.attn = nn.MultiheadAttention(fusion_dim, num_heads=4, dropout=0.1, batch_first=True)
170 self.norm = nn.LayerNorm(fusion_dim)
171
172 def forward(self, feats):
173 # feats is list of (B, D)
174 stacked = torch.stack(feats, dim=1) # (B, 3, D)
175 # Add learned embeddings
176 x = stacked + self.modality_embed
177
178 # Self-attention across modalities
179 attn_out, weights = self.attn(x, x, x)
180 self.last_weights = weights.mean(dim=1) # (B, 3) - for visualization
181
182 x = self.norm(x + attn_out)
183
184 # Weighted average based on learned weights
185 w = torch.softmax(self.weights, dim=0)
186 out = (x * w.unsqueeze(0).unsqueeze(-1)).sum(dim=1)
187 return out
188
189class BrainrotModelV4(nn.Module):
190 def __init__(self):
191 super().__init__()
192 fd = 256
193 self.visual_proj = ModalityProjector(2560, fd)
194 self.audio_proj = ModalityProjector(384, fd)
195 self.text_proj = ModalityProjector(5, fd) # Slang, Words, AvgLen, Phonetic, Clarity
196 self.fusion = CrossAttentionFusion(fd)
197 self.classifier = nn.Sequential(
198 nn.Linear(fd, fd // 2), nn.GELU(),
199 nn.Dropout(0.2), nn.Linear(fd // 2, 1)
200 )
201 def forward(self, v, a, t):
202 return self.classifier(self.fusion([
203 self.visual_proj(v), self.audio_proj(a), self.text_proj(t)
204 ]))
205
206
207# =============================================================================
208# MODEL MANAGER (loads once on startup, reuses for every request)
209# =============================================================================
210
211class ModelManager:
212 def __init__(self):
213 self.models = {}
214 self.eff_model = None
215 self.bert_model = None
216 self.tokenizer = None
217 self._loaded = False
218 self.current_version = None
219 self.whisper_model = None
220
221 # Image transforms matching EfficientNet V1 ImageNet training
222 self.img_transform = transforms.Compose([
223 transforms.ToPILImage(),
224 transforms.Resize((IMG_SIZE, IMG_SIZE)),
225 transforms.ToTensor(),
226 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
227 ])
228
229 def _load_feature_extractors(self):
230 """Load shared feature extractors (only once, reused across versions)."""
231 from transformers import DistilBertModel, DistilBertTokenizer
232
233 print(" [OK] Loading EfficientNet-B0 (PyTorch)...")
234 self.eff_model = models.efficientnet_b0(weights=models.EfficientNet_B0_Weights.DEFAULT).to(DEVICE)
235 self.eff_model.classifier = nn.Identity() # Remove top layer to get features
236 self.eff_model.eval()
237
238 print(" [OK] Loading DistilBert (HuggingFace)...")
239 self.tokenizer = DistilBertTokenizer.from_pretrained('distilbert-base-uncased')
240 self.bert_model = DistilBertModel.from_pretrained('distilbert-base-uncased').to(DEVICE)
241 self.bert_model.eval()
242
243 try:
244 import whisper
245 print(" [OK] Loading Whisper (Tiny)...")
246 self.whisper_model = whisper.load_model("tiny")
247 except Exception as e:
248 print(f" [!] Whisper load failed: {e}")
249
250 def _load_fold_models(self, version_id: str):
251 """Load fold checkpoint models for a specific version."""
252 version_cfg = MODEL_VERSIONS[version_id]
253 checkpoint_dir = version_cfg["checkpoint_dir"]
254 n_folds = version_cfg["n_folds"]
255
256 # Unload existing fold models
257 self.models.clear()
258 gc.collect()
259 if torch.cuda.is_available():
260 torch.cuda.empty_cache()
261
262 print(f"\n[ModelManager] Loading version '{version_id}' ({version_cfg['name']}) from {checkpoint_dir}...")
263
264 for fold_num in range(1, n_folds + 1):
265 ckpt_path = os.path.join(checkpoint_dir, f"BEST_fold{fold_num}.pt")
266 if os.path.exists(ckpt_path):
267 if version_id == "model4":
268 model = BrainrotModelV4().to(DEVICE)
269 elif version_id == "model3":
270 model = BrainrotModelV3().to(DEVICE)
271 else:
272 model = BrainrotModelV1().to(DEVICE)
273 ckpt = torch.load(ckpt_path, map_location=DEVICE)
274 model.load_state_dict(ckpt["model"])
275 model.eval()
276 self.models[fold_num] = model
277 print(f" [OK] Loaded BEST_fold{fold_num}.pt")
278 else:
279 print(f" [OK] Found checkpoint: {ckpt_path}")
280
281 if not self.models:
282 raise RuntimeError(f"No checkpoints found in: {checkpoint_dir}")
283
284 self.current_version = version_id
285 print(f"[ModelManager] Version '{version_id}' ready. Folds available: {list(self.models.keys())}")
286
287 def load_all(self, version_id: str = None):
288 """Initial load: feature extractors + default version fold models."""
289 if version_id is None:
290 version_id = DEFAULT_VERSION
291
292 if not self._loaded:
293 print("\n[ModelManager] Models loaded successfully!")
294 self._load_feature_extractors()
295 self._loaded = True
296
297 self._load_fold_models(version_id)
298
299 def switch_version(self, version_id: str):
300 """Hot-swap to a different model version (keeps feature extractors loaded)."""
301 if version_id not in MODEL_VERSIONS:
302 raise ValueError(f"Unknown model version: '{version_id}'. Available: {list(MODEL_VERSIONS.keys())}")
303
304 if version_id == self.current_version:
305 print(f"[ModelManager] Version '{version_id}' is already loaded.")
306 return
307
308 self._load_fold_models(version_id)
309
310 def ensure_version(self, version_id: str):
311 """Ensure the requested version is loaded (auto-swap if needed)."""
312 if version_id and version_id != self.current_version:
313 self.switch_version(version_id)
314
315 @property
316 def available_folds(self):
317 return list(self.models.keys())
318
319 @property
320 def version_name(self):
321 if self.current_version and self.current_version in MODEL_VERSIONS:
322 return MODEL_VERSIONS[self.current_version]["name"]
323 return "Unknown"
324
325
326model_manager = ModelManager()
327
328
329# =============================================================================
330# FEATURE EXTRACTION
331# =============================================================================
332
333def extract_visual(video_path: str, ) -> np.ndarray:
334 import cv2
335
336 cap = cv2.VideoCapture(video_path)
337 total = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
338 frames = []
339
340 indices = np.linspace(0, max(total - 1, 0), MAX_FRAMES, dtype=int)
341 for idx in indices:
342 cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
343 ret, frame = cap.read()
344 if ret:
345 frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
346 pixel_values = model_manager.img_transform(frame)
347 frames.append(pixel_values)
348 cap.release()
349
350 if not frames:
351 return np.zeros((1, VISUAL_FEAT_DIM), dtype=np.float32)
352
353 batch = torch.stack(frames).to(DEVICE)
354 with torch.no_grad():
355 feats = model_manager.eff_model(batch) # (N, 1280)
356
357 return feats.cpu().numpy().astype(np.float32)
358
359
360def extract_audio(video_path: str, ) -> np.ndarray:
361 import librosa
362
363 audio_path = video_path + "_audio.wav"
364 ret = os.system(f'ffmpeg -i "{video_path}" -ac 1 -ar 22050 "{audio_path}" -y -loglevel quiet')
365
366 if ret != 0 or not os.path.exists(audio_path):
367 return np.zeros((1, AUDIO_FEAT_DIM), dtype=np.float32)
368
369 try:
370 y, sr = librosa.load(audio_path, sr=22050, mono=True)
371 os.remove(audio_path)
372
373 mel = librosa.feature.melspectrogram(y=y, sr=sr, n_mels=N_MELS)
374 mel_db = librosa.power_to_db(mel, ref=np.max).astype(np.float32).T
375
376 if mel_db.shape[0] >= AUDIO_MAX_FRAMES:
377 mel_db = mel_db[:AUDIO_MAX_FRAMES, :]
378 else:
379 pad = np.zeros((AUDIO_MAX_FRAMES - mel_db.shape[0], N_MELS), dtype=np.float32)
380 mel_db = np.vstack([mel_db, pad])
381
382 flat = mel_db.flatten()[None]
383 if flat.shape[1] >= AUDIO_FEAT_DIM:
384 return flat[:, :AUDIO_FEAT_DIM].astype(np.float32)
385 else:
386 feat = np.zeros((1, AUDIO_FEAT_DIM), dtype=np.float32)
387 feat[:, :flat.shape[1]] = flat
388 return feat
389 except Exception:
390 if os.path.exists(audio_path):
391 os.remove(audio_path)
392 return np.zeros((1, AUDIO_FEAT_DIM), dtype=np.float32)
393
394
395def extract_text(video_path: str, ) -> tuple:
396 """Returns (feature_array, transcript_string)"""
397 transcript = "no speech detected"
398
399 try:
400 import whisper
401 audio_path = video_path + "_whisper.wav"
402 os.system(f'ffmpeg -i "{video_path}" -ac 1 -ar 16000 "{audio_path}" -y -loglevel quiet')
403 if os.path.exists(audio_path):
404 wm = whisper.load_model("tiny")
405 result = wm.transcribe(audio_path)
406 transcript = result.get("text", "").strip() or "no speech detected"
407 os.remove(audio_path)
408 del wm
409 gc.collect()
410 except ImportError:
411 pass
412
413 encoded = model_manager.tokenizer(
414 [transcript],
415 padding='max_length',
416 truncation=True,
417 max_length=MAX_TEXT_LEN,
418 return_tensors='pt'
419 )
420 input_ids = encoded['input_ids'].to(DEVICE)
421 attention_mask = encoded['attention_mask'].to(DEVICE)
422
423 with torch.no_grad():
424 output = model_manager.bert_model(input_ids=input_ids, attention_mask=attention_mask)
425
426 feat = output.last_hidden_state[:, 0, :].cpu().numpy()
427 return feat.astype(np.float32), transcript
428
429
430def run_inference(video_path: str, fold_nums: Optional[List[int]] = None, model_version: str = None, task_id: str = None) -> dict:
431 if task_id:
432 ACTIVE_TASKS_STATE[task_id] = {"stage": "Starting Inference", "log": "Initializing pipeline..."}
433
434 if not model_version:
435 model_version = model_manager.current_version
436
437 model_manager.ensure_version(model_version)
438 is_v3 = (model_version == "model3")
439 is_v4 = (model_version == "model4")
440
441 if fold_nums is None:
442 fold_nums = model_manager.available_folds
443
444 print(f"[Inference] Starting pipeline for version: {model_version}")
445 t0 = time.time()
446
447 # 1. TRANSCRIPTION (Common)
448 transcript = "no speech detected"
449 try:
450 if model_manager.whisper_model:
451 print(" [Step 1/4] Transcribing audio with Whisper...")
452 audio_path = video_path + "_whisper.wav"
453 # Extract 16kHz mono audio for whisper
454 os.system(f'ffmpeg -i "{video_path}" -ac 1 -ar 16000 "{audio_path}" -y -loglevel quiet')
455 if os.path.exists(audio_path):
456 result = model_manager.whisper_model.transcribe(audio_path)
457 transcript = result.get("text", "").strip() or "no speech detected"
458 os.remove(audio_path)
459 except Exception as e:
460 print(f" [!] Transcription failed: {e}")
461
462 # 2. EXTRACTION
463 try:
464 if is_v3:
465 # VISUAL V3
466 print(" [Step 2/4] Extracting Visual V3 (32 frames, mean-max)...")
467 import cv2
468 cap = cv2.VideoCapture(video_path)
469 total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
470 indices = np.linspace(0, max(total_frames - 1, 0), 32, dtype=int)
471 frames = []
472 for idx in indices:
473 cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
474 ret, frame = cap.read()
475 if ret:
476 frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
477 pixel_values = model_manager.img_transform(frame)
478 frames.append(pixel_values)
479 cap.release()
480 if not frames:
481 vis_feat = np.zeros((1, 2560), dtype=np.float32)
482 else:
483 batch = torch.stack(frames).to(DEVICE)
484 with torch.no_grad():
485 feats = model_manager.eff_model(batch).unsqueeze(0)
486 vis_feat = torch.cat([feats.mean(dim=1), feats.max(dim=1).values], dim=-1).cpu().numpy().astype(np.float32)
487 t1 = time.time()
488
489 # AUDIO V3
490 print(" [Step 3/4] Extracting Audio V3 (stats)...")
491 import librosa
492 try:
493 y, _ = librosa.load(video_path, sr=16000, mono=True, duration=60)
494 if len(y) > 0:
495 mel_db = librosa.power_to_db(librosa.feature.melspectrogram(y=y, sr=16000, n_mels=128, fmax=8000), ref=np.max)
496 aud_feat = np.concatenate([mel_db.mean(axis=1), mel_db.std(axis=1), mel_db.max(axis=1)]).astype(np.float32)
497 else:
498 aud_feat = np.zeros(384, dtype=np.float32)
499 except Exception as e:
500 print(f" [!] Audio V3 failed: {e}")
501 aud_feat = np.zeros(384, dtype=np.float32)
502 aud_feat = np.expand_dims(aud_feat, axis=0)
503 t2 = time.time()
504
505 # TEXT V3
506 print(" [Step 4/4] Extracting Text V3 (slang dictionary)...")
507 SLANG_WORDS = ["skibidi", "rizz", "gyatt", "sigma", "fanum", "mewing", "edge", "rizzler", "aura", "mog", "looksmaxxing", "ohio", "delulu", "bussin", "no cap", "fr", "ong", "goated", "cooked", "mid", "sus", "beta", "alpha", "pookie", "hawk tuah", "fanum tax", "negative aura", "aura farming", "mew", "goon", "grimace", "6-7", "sixty seven", "diddyblud", "nonchalant", "mango", "zesty", "ate", "slay", "brat", "lock in", "tweaking", "yapping", "sybau", "chuzz", "huzz", "gurt", "womp womp", "skibidi toilet", "ohio rizz", "level 10 gyatt", "baby gronk", "let him cook", "ratio", "L", "goofinator", "brainrot", "slop", "chronocore", "main npc", "tung tung", "sahur", "ballerina", "cappuccina"]
508 txt = transcript.lower()
509 s_count = sum(txt.count(w) for w in SLANG_WORDS)
510 w_count = len(txt.split())
511 avg_len = len(txt) / max(1, w_count)
512 txt_feat = np.expand_dims(np.array([float(s_count), float(w_count), float(avg_len)], dtype=np.float32), axis=0)
513 t3 = time.time()
514 vis_seq = None
515
516 elif is_v4:
517 # VISUAL V4 (Same as V3 but potentially different stats if needed)
518 print(" [Step 2/4] Extracting Visual V4 (32 frames, mean-max)...")
519 import cv2
520 cap = cv2.VideoCapture(video_path)
521 total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
522 indices = np.linspace(0, max(total_frames - 1, 0), 32, dtype=int)
523 frames = []
524 for idx in indices:
525 cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
526 ret, frame = cap.read()
527 if ret:
528 frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
529 pixel_values = model_manager.img_transform(frame)
530 frames.append(pixel_values)
531 cap.release()
532 if not frames:
533 vis_feat = np.zeros((1, 2560), dtype=np.float32)
534 else:
535 batch = torch.stack(frames).to(DEVICE)
536 with torch.no_grad():
537 feats = model_manager.eff_model(batch).unsqueeze(0)
538 vis_feat = torch.cat([feats.mean(dim=1), feats.max(dim=1).values], dim=-1).cpu().numpy().astype(np.float32)
539 t1 = time.time()
540
541 # AUDIO V4
542 print(" [Step 3/4] Extracting Audio V4 (stats)...")
543 import librosa
544 try:
545 y, _ = librosa.load(video_path, sr=16000, mono=True, duration=60)
546 if len(y) > 0:
547 mel_db = librosa.power_to_db(librosa.feature.melspectrogram(y=y, sr=16000, n_mels=128, fmax=8000), ref=np.max)
548 aud_feat = np.concatenate([mel_db.mean(axis=1), mel_db.std(axis=1), mel_db.max(axis=1)]).astype(np.float32)
549 else:
550 aud_feat = np.zeros(384, dtype=np.float32)
551 except Exception as e:
552 print(f" [!] Audio V4 failed: {e}")
553 aud_feat = np.zeros(384, dtype=np.float32)
554 aud_feat = np.expand_dims(aud_feat, axis=0)
555 t2 = time.time()
556
557 # TEXT V4 (Phonetic & Clarity)
558 print(" [Step 4/4] Extracting Text V4 (Phonetic + Clarity)...")
559 SLANG_WORDS = ["skibidi", "rizz", "gyatt", "sigma", "fanum", "mewing", "edge", "rizzler", "aura", "mog", "looksmaxxing", "ohio", "delulu", "bussin", "no cap", "fr", "ong", "goated", "cooked", "mid", "sus", "beta", "alpha", "pookie", "hawk tuah", "fanum tax", "negative aura", "aura farming", "mew", "goon", "grimace", "6-7", "sixty seven", "diddyblud", "nonchalant", "mango", "zesty", "ate", "slay", "brat", "lock in", "tweaking", "yapping", "sybau", "chuzz", "huzz", "gurt", "womp womp", "skibidi toilet", "ohio rizz", "level 10 gyatt", "baby gronk", "let him cook", "ratio", "L", "goofinator", "brainrot", "slop", "chronocore", "main npc", "tung tung", "sahur", "ballerina", "cappuccina"]
560 txt = transcript.lower()
561 words = txt.split()
562 s_count = sum(txt.count(w) for w in SLANG_WORDS)
563 w_count = len(words)
564 avg_len = len(txt) / max(1, w_count)
565
566 # 4. Phonetic Slang (Repeated characters common in brainrot)
567 phonetic_score = len(re.findall(r'(.)\1{2,}', txt))
568
569 # 5. Speech Clarity (Ratio of clean words)
570 clarity_score = 1.0 - (s_count / max(1, w_count))
571
572 txt_feat = np.expand_dims(np.array([
573 float(s_count), float(w_count), float(avg_len),
574 float(phonetic_score), float(clarity_score)
575 ], dtype=np.float32), axis=0)
576 t3 = time.time()
577 vis_seq = None
578
579 else:
580 # V1 Extraction
581 print(" [Step 2/4] Extracting Visual V1...")
582 vis_seq = extract_visual(video_path)
583 vis_feat = np.mean(vis_seq, axis=0, keepdims=True)
584 t1 = time.time()
585
586 print(" [Step 3/4] Extracting Audio V1...")
587 aud_feat = extract_audio(video_path)
588 t2 = time.time()
589
590 print(" [Step 4/4] Extracting Text V1 (BERT)...")
591 encoded = model_manager.tokenizer([transcript], padding='max_length', truncation=True, max_length=MAX_TEXT_LEN, return_tensors='pt')
592 input_ids = encoded['input_ids'].to(DEVICE)
593 attention_mask = encoded['attention_mask'].to(DEVICE)
594 with torch.no_grad():
595 output = model_manager.bert_model(input_ids=input_ids, attention_mask=attention_mask)
596 txt_feat = output.last_hidden_state[:, 0, :].cpu().numpy().astype(np.float32)
597 t3 = time.time()
598 except Exception as e:
599 print(f" [!] Feature extraction failed: {e}")
600 raise e
601
602 # 3. FORWARD PASS
603 print(f" [Inference] Running forward pass on {len(fold_nums)} folds...")
604 vis_t = torch.tensor(vis_feat, dtype=torch.float32).to(DEVICE)
605 aud_t = torch.tensor(aud_feat, dtype=torch.float32).to(DEVICE)
606 txt_t = torch.tensor(txt_feat, dtype=torch.float32).to(DEVICE)
607
608 per_fold_probs = []
609 attentions_list = []
610
611 # NEW: Calculate Modality Importance based on Projected Norms (L2) + Attention
612 # This helps fix the "stiffening" bias by showing both active presence and model attention
613 modality_norms = []
614
615 for fold_num in fold_nums:
616 if fold_num not in model_manager.models: continue
617 model = model_manager.models[fold_num]
618 with torch.no_grad():
619 # Get projected features to calculate norms
620 v_proj = model.visual_proj(vis_t)
621 a_proj = model.audio_proj(aud_t)
622 t_proj = model.text_proj(txt_t)
623
624 # L2 Norms of projected features (represents signal strength)
625 norms = [
626 torch.norm(v_proj).item(),
627 torch.norm(a_proj).item(),
628 torch.norm(t_proj).item()
629 ]
630 modality_norms.append(norms)
631
632 logits = model(vis_t, aud_t, txt_t)
633 # Try to get attention weights
634 try:
635 w = model.fusion.last_weights[0].cpu().numpy().tolist()
636 attentions_list.append(w)
637 except: pass
638
639 if is_v4:
640 raw_logit = logits.squeeze().item()
641 # Model 4 Calibration (assumed similar to V3 but with V4 specific distribution)
642 # Using 1.2 mean and 0.8 std as a refined starting point for V4 if not provided
643 standardized = (raw_logit - 1.2010) / (0.8500 + 1e-8)
644 prob = 1 / (1 + np.exp(-standardized))
645 per_fold_probs.append(float(prob))
646 elif is_v3:
647 raw_logit = logits.squeeze().item()
648 # Fuzzy calibration
649 standardized = (raw_logit - 1.4614) / (1.1190 + 1e-8)
650 prob = 1 / (1 + np.exp(-standardized))
651 per_fold_probs.append(float(prob))
652 else:
653 probs = torch.softmax(logits, dim=-1)
654 per_fold_probs.append(float(probs[0][1].cpu()))
655
656 if not per_fold_probs:
657 raise ValueError("No predictions were generated. Check fold models.")
658
659 # Combined weights: Attention * Normalized Norms (balanced across modalities)
660 # Raw norms are biased by input dimensionality (Visual:2560→256 >> Text:3→256),
661 # so we normalize each norm to [0,1] before combining with attention.
662 avg_norms = np.mean(modality_norms, axis=0)
663 max_norm = np.max(avg_norms) + 1e-8
664 normalized_norms = avg_norms / max_norm # Scale all to [0, 1]
665 avg_attn = np.mean(attentions_list, axis=0) if attentions_list else np.array([0.33, 0.33, 0.34])
666 # Geometric mean balances "signal presence" (norm) with "model focus" (attention)
667 combined_importance = np.sqrt(normalized_norms * avg_attn)
668 modality_weights = (combined_importance / (np.sum(combined_importance) + 1e-8)).tolist()
669
670 # 4. TEMPORAL PROBS (V1 sequence or V3 simulated sequence)
671 temporal_probs = []
672 first_fold = model_manager.models.get(fold_nums[0])
673
674 if first_fold:
675 if not is_v3 and vis_seq is not None:
676 # V1 logic
677 try:
678 vis_seq_t = torch.tensor(vis_seq, dtype=torch.float32).to(DEVICE)
679 N = vis_seq_t.size(0)
680 t_logits = first_fold(vis_seq_t, aud_t.repeat(N, 1), txt_t.repeat(N, 1))
681 temporal_probs = torch.softmax(t_logits, dim=-1)[:, 1].cpu().numpy().tolist()
682 except: pass
683 elif is_v3 and 'frames' in locals() and len(frames) > 0:
684 # V3 logic: Simulated temporal sequence by running inference on frame windows
685 # We divide the 32 frames into 8 chunks of 4 frames
686 try:
687 windows = []
688 for i in range(0, 32, 4):
689 window_frames = frames[i:i+4]
690 if not window_frames: continue
691 w_batch = torch.stack(window_frames).to(DEVICE)
692 with torch.no_grad():
693 w_feats = model_manager.eff_model(w_batch).unsqueeze(0)
694 w_vis = torch.cat([w_feats.mean(dim=1), w_feats.max(dim=1).values], dim=-1)
695 w_logits = first_fold(w_vis, aud_t, txt_t)
696 w_standardized = (w_logits.squeeze().item() - 1.4614) / (1.1190 + 1e-8)
697 w_prob = 1 / (1 + np.exp(-w_standardized))
698 windows.append(float(w_prob))
699 temporal_probs = windows
700 except: pass
701 elif is_v4 and 'frames' in locals() and len(frames) > 0:
702 # V4 temporal sequence
703 try:
704 windows = []
705 for i in range(0, 32, 4):
706 window_frames = frames[i:i+4]
707 if not window_frames: continue
708 w_batch = torch.stack(window_frames).to(DEVICE)
709 with torch.no_grad():
710 w_feats = model_manager.eff_model(w_batch).unsqueeze(0)
711 w_vis = torch.cat([w_feats.mean(dim=1), w_feats.max(dim=1).values], dim=-1)
712 w_logits = first_fold(w_vis, aud_t, txt_t)
713 w_standardized = (w_logits.squeeze().item() - 1.2010) / (0.8500 + 1e-8)
714 w_prob = 1 / (1 + np.exp(-w_standardized))
715 windows.append(float(w_prob))
716 temporal_probs = windows
717 except: pass
718
719 if not temporal_probs:
720 temporal_probs = [float(np.mean(per_fold_probs))]
721
722 # 5. TOKEN SALIENCY (V3/V4 slang heatmap)
723 token_saliency = []
724 if is_v3 or is_v4:
725 words = transcript.split()
726 SLANG_WORDS = ["skibidi", "rizz", "gyatt", "sigma", "fanum", "mewing", "edge", "rizzler", "aura", "mog", "looksmaxxing", "ohio", "delulu", "bussin", "no cap", "fr", "ong", "goated", "cooked", "mid", "sus", "beta", "alpha", "pookie", "hawk tuah", "fanum tax", "negative aura", "aura farming", "mew", "goon", "grimace", "6-7", "sixty seven", "diddyblud", "nonchalant", "mango", "zesty", "ate", "slay", "brat", "lock in", "tweaking", "yapping", "sybau", "chuzz", "huzz", "gurt", "womp womp", "skibidi toilet", "ohio rizz", "level 10 gyatt", "baby gronk", "let him cook", "ratio", "L", "goofinator", "brainrot", "slop", "chronocore", "main npc", "tung tung", "sahur", "ballerina", "cappuccina"]
727 for w in words:
728 clean_w = re.sub(r'[^a-zA-Z]', '', w).lower()
729 score = 0.0
730 if clean_w in SLANG_WORDS:
731 score = 0.8 + (0.2 * (len(clean_w) / 15)) # higher score for longer/rarer slang
732 token_saliency.append({"word": w, "score": round(score, 3)})
733
734 # 6. FEATURE IMPORTANCE (SHAP-style)
735 feature_importance = []
736 if is_v3 or is_v4:
737 # Simplified SHAP calculation based on feature presence and modality weights
738 v_imp = modality_weights[0] * (vis_feat.max() / 10.0)
739 a_imp = modality_weights[1] * (aud_feat.mean() / 5.0)
740 t_imp = modality_weights[2] * (s_count / 5.0)
741
742 # Normalize impacts to percentages
743 total_imp = v_imp + a_imp + t_imp + 1e-8
744
745 if is_v4:
746 feature_importance = [
747 {"feature": "Phonetic Slang", "impact": round((phonetic_score / 5.0) * 20, 1)},
748 {"feature": "Slang Density", "impact": round((t_imp / total_imp) * 100, 1)},
749 {"feature": "Speech Clarity", "impact": round((clarity_score) * -15, 1)}, # negative impact usually
750 {"feature": "Visual Peaks", "impact": round((v_imp / total_imp) * 80, 1)},
751 {"feature": "Audio Energy", "impact": round((a_imp / total_imp) * 60, 1)},
752 ]
753 else:
754 feature_importance = [
755 {"feature": "Slang Density", "impact": round((t_imp / total_imp) * 100, 1)},
756 {"feature": "Visual Peak Intensity", "impact": round((v_imp / total_imp) * 80, 1)},
757 {"feature": "Audio Energy", "impact": round((a_imp / total_imp) * 60, 1)},
758 {"feature": "Text Complexity", "impact": round((w_count / 100) * 10, 1)},
759 {"feature": "Temporal Variance", "impact": round(np.std(temporal_probs) * 20, 1) if len(temporal_probs) > 1 else 0.5}
760 ]
761 # Sort by impact
762 feature_importance.sort(key=lambda x: abs(x["impact"]), reverse=True)
763
764 t4 = time.time()
765 avg_prob = float(np.mean(per_fold_probs))
766 prediction = "BRAINROT" if avg_prob >= 0.5 else "NON_BRAINROT"
767 certainty = abs(avg_prob - 0.5) * 2 # 0.0 to 1.0
768
769 print(f"[Inference] Done! Result: {prediction} ({avg_prob:.2%})")
770
771 return {
772 "prediction": prediction,
773 "confidence": round(avg_prob if prediction == "BRAINROT" else 1 - avg_prob, 4),
774 "certainty": round(certainty, 4),
775 "prob_brainrot": round(avg_prob, 4),
776 "prob_non_brainrot": round(1 - avg_prob, 4),
777 "transcript": transcript,
778 "token_saliency": token_saliency,
779 "feature_importance": feature_importance,
780 "folds_used": len(per_fold_probs),
781 "per_fold_probs": [round(p, 4) for p in per_fold_probs],
782 "model_version": model_manager.current_version,
783 "model_version_name": model_manager.version_name,
784 "pipeline_metrics": {
785 "visual_ext_s": round(t1 - t0, 3),
786 "audio_ext_s": round(t2 - t1, 3),
787 "text_ext_s": round(t3 - t2, 3),
788 "inference_s": round(t4 - t3, 3)
789 },
790 "modality_weights": [round(w, 4) for w in modality_weights],
791 "temporal_probs": [round(p, 4) for p in temporal_probs],
792 "modality_features": {
793 "text": {
794 "slang_count": s_count,
795 "word_count": w_count,
796 "avg_word_len": round(avg_len, 2),
797 "phonetic_score": phonetic_score if is_v4 else None,
798 "clarity_score": round(clarity_score, 3) if is_v4 else None
799 },
800 "visual": {
801 "mean_intensity": round(float(np.mean(vis_feat)), 4),
802 "max_spike": round(float(np.max(vis_feat)), 4)
803 }
804 } if (is_v3 or is_v4) else None
805 }
806
807app = FastAPI(
808 title="Brainrot Detector API",
809 description="Multimodal brainrot video classification using Visual + Audio + Text fusion",
810 version="1.0.0",
811)
812
813app.add_middleware(
814 CORSMiddleware,
815 allow_origins=["*"],
816 allow_methods=["*"],
817 allow_headers=["*"],
818)
819
820
821# ── Response schemas ──────────────────────────────────────────────────────────
822
823class PipelineMetrics(BaseModel):
824 visual_ext_s: float
825 audio_ext_s: float
826 text_ext_s: float
827 inference_s: float
828
829class PredictionResponse(BaseModel):
830 video_name: str
831 prediction: str
832 confidence: float
833 prob_brainrot: float
834 prob_non_brainrot: float
835 transcript: str
836 folds_used: int
837 per_fold_probs: List[float]
838 processing_time_s: float
839 model_version: str
840 model_version_name: str
841 pipeline_metrics: PipelineMetrics
842 modality_weights: List[float]
843 temporal_probs: List[float]
844 modality_features: Optional[dict] = None
845
846
847class HealthResponse(BaseModel):
848 status: str
849 device: str
850 folds_loaded: List[int]
851 models_ready: bool
852 model_version: str
853 model_version_name: str
854
855
856class ModelInfoResponse(BaseModel):
857 model_version: str
858 model_version_name: str
859 checkpoint_dir: str
860 folds_available: List[int]
861 device: str
862 visual_dim: int
863 audio_dim: int
864 text_dim: int
865 fusion_dim: int
866
867
868class ModelVersionInfo(BaseModel):
869 version_id: str
870 name: str
871 checkpoint_dir: str
872 n_folds: int
873 is_active: bool
874
875
876class ModelsListResponse(BaseModel):
877 active_version: str
878 versions: List[ModelVersionInfo]
879
880
881class SwitchResponse(BaseModel):
882 message: str
883 active_version: str
884 folds_loaded: List[int]
885
886
887# ── Startup ───────────────────────────────────────────────────────────────────
888
889@app.on_event("startup")
890async def startup_event():
891 print("[API] Starting up — loading models...")
892 model_manager.load_all()
893 print("[API] Ready to serve requests.")
894
895
896# ── Endpoints ─────────────────────────────────────────────────────────────────
897
898@app.get("/health", response_model=HealthResponse, tags=["Status"])
899async def health_check():
900 """Check if the API server and models are ready."""
901 return HealthResponse(
902 status = "ok",
903 device = str(DEVICE),
904 folds_loaded = model_manager.available_folds,
905 models_ready = model_manager._loaded,
906 model_version = model_manager.current_version or "",
907 model_version_name = model_manager.version_name,
908 )
909
910
911@app.get("/model/info", response_model=ModelInfoResponse, tags=["Status"])
912async def model_info():
913 """Get information about the loaded model configuration."""
914 version_cfg = MODEL_VERSIONS.get(model_manager.current_version, {})
915 return ModelInfoResponse(
916 model_version = model_manager.current_version or "",
917 model_version_name = model_manager.version_name,
918 checkpoint_dir = version_cfg.get("checkpoint_dir", ""),
919 folds_available = model_manager.available_folds,
920 device = str(DEVICE),
921 visual_dim = VISUAL_FEAT_DIM,
922 audio_dim = AUDIO_FEAT_DIM,
923 text_dim = TEXT_FEAT_DIM,
924 fusion_dim = FUSION_DIM,
925 )
926
927
928@app.get("/models", response_model=ModelsListResponse, tags=["Model Versions"])
929async def list_models():
930 """List all available model versions and which one is currently active."""
931 versions = []
932 for vid, cfg in MODEL_VERSIONS.items():
933 versions.append(ModelVersionInfo(
934 version_id = vid,
935 name = cfg["name"],
936 checkpoint_dir = cfg["checkpoint_dir"],
937 n_folds = cfg["n_folds"],
938 is_active = (vid == model_manager.current_version),
939 ))
940 return ModelsListResponse(
941 active_version = model_manager.current_version or "",
942 versions = versions,
943 )
944
945
946@app.post("/models/switch", response_model=SwitchResponse, tags=["Model Versions"])
947async def switch_model(version: str):
948 """
949 Switch the active model version.
950 Available versions can be listed via GET /models.
951 """
952 if version not in MODEL_VERSIONS:
953 raise HTTPException(
954 status_code=400,
955 detail=f"Unknown version: '{version}'. Available: {list(MODEL_VERSIONS.keys())}"
956 )
957 try:
958 model_manager.switch_version(version)
959 return SwitchResponse(
960 message = f"Switched to '{version}' ({MODEL_VERSIONS[version]['name']})",
961 active_version = model_manager.current_version,
962 folds_loaded = model_manager.available_folds,
963 )
964 except Exception as e:
965 raise HTTPException(status_code=500, detail=str(e))
966
967
968
969
970@app.post("/predict", response_model=PredictionResponse, tags=["Inference"])
971async def predict(
972 video: UploadFile = File(..., description="Video file to classify (mp4, avi, mov)"),
973 model_version: Optional[str] = None,
974):
975 """
976 Classify a video using the best model from Fold 1.
977 Optionally specify model_version ('default' or 'no_yt') as a query parameter.
978 Faster than ensemble — good for quick classification.
979 """
980 if not model_manager._loaded:
981 raise HTTPException(status_code=503, detail="Models not loaded yet.")
982
983 if model_version and model_version not in MODEL_VERSIONS:
984 raise HTTPException(
985 status_code=400,
986 detail=f"Unknown model version: '{model_version}'. Available: {list(MODEL_VERSIONS.keys())}"
987 )
988
989 # Validate file type
990 allowed = {".mp4", ".avi", ".mov", ".mkv", ".webm", ".flv"}
991 ext = os.path.splitext(video.filename)[1].lower()
992 if ext not in allowed:
993 raise HTTPException(
994 status_code=400,
995 detail=f"Unsupported file type: {ext}. Allowed: {allowed}"
996 )
997
998 # Save uploaded file to temp location
999 tmp_path = os.path.join(tempfile.gettempdir(), f"{uuid.uuid4()}{ext}")
1000 try:
1001 with open(tmp_path, "wb") as f:
1002 content = await video.read()
1003 f.write(content)
1004
1005 start_time = time.time()
1006 result = await asyncio.to_thread(run_inference, tmp_path, [min(model_manager.available_folds)], model_version)
1007 elapsed = time.time() - start_time
1008
1009 return PredictionResponse(
1010 video_name = video.filename,
1011 processing_time_s = round(elapsed, 2),
1012 **result,
1013 )
1014 except HTTPException:
1015 raise
1016 except Exception as e:
1017 raise HTTPException(status_code=500, detail=str(e))
1018 finally:
1019 if os.path.exists(tmp_path):
1020 os.remove(tmp_path)
1021
1022
1023@app.post("/predict/ensemble", response_model=PredictionResponse, tags=["Inference"])
1024async def predict_ensemble(
1025 video: UploadFile = File(..., description="Video file to classify (mp4, avi, mov)"),
1026 model_version: Optional[str] = None,
1027):
1028 """
1029 Classify a video using ALL fold models and average their predictions.
1030 Optionally specify model_version ('default' or 'no_yt') as a query parameter.
1031 More accurate than single-fold — recommended for final results.
1032 """
1033 if not model_manager._loaded:
1034 raise HTTPException(status_code=503, detail="Models not loaded yet.")
1035
1036 if model_version and model_version not in MODEL_VERSIONS:
1037 raise HTTPException(
1038 status_code=400,
1039 detail=f"Unknown model version: '{model_version}'. Available: {list(MODEL_VERSIONS.keys())}"
1040 )
1041
1042 allowed = {".mp4", ".avi", ".mov", ".mkv", ".webm", ".flv"}
1043 ext = os.path.splitext(video.filename)[1].lower()
1044 if ext not in allowed:
1045 raise HTTPException(
1046 status_code=400,
1047 detail=f"Unsupported file type: {ext}. Allowed: {allowed}"
1048 )
1049
1050 tmp_path = os.path.join(tempfile.gettempdir(), f"{uuid.uuid4()}{ext}")
1051 try:
1052 with open(tmp_path, "wb") as f:
1053 content = await video.read()
1054 f.write(content)
1055
1056 start_time = time.time()
1057 result = await asyncio.to_thread(run_inference, tmp_path, model_manager.available_folds, model_version)
1058 elapsed = time.time() - start_time
1059
1060 return PredictionResponse(
1061 video_name = video.filename,
1062 processing_time_s = round(elapsed, 2),
1063 **result,
1064 )
1065 except HTTPException:
1066 raise
1067 except Exception as e:
1068 raise HTTPException(status_code=500, detail=str(e))
1069 finally:
1070 if os.path.exists(tmp_path):
1071 os.remove(tmp_path)
1072
1073
1074# =============================================================================
1075# URL-BASED INFERENCE (Download from TikTok / Instagram / YouTube)
1076# =============================================================================
1077
1078class URLRequest(BaseModel):
1079 url: str
1080 task_id: Optional[str] = None
1081 model_version: Optional[str] = None
1082 mode: Optional[str] = "ensemble" # 'single' or 'ensemble'
1083
1084
1085class URLValidateResponse(BaseModel):
1086 valid: bool
1087 platform: Optional[str] = None
1088 url: str
1089
1090
1091@app.post("/validate/url", response_model=URLValidateResponse, tags=["URL Download"])
1092async def validate_url(payload: URLRequest):
1093 """Validate a URL and detect the platform (youtube, tiktok, instagram)."""
1094 url = payload.url.strip()
1095 if not url:
1096 return URLValidateResponse(valid=False, platform=None, url=url)
1097
1098 url_lower = url.lower()
1099 supported_patterns = [
1100 'youtube.com', 'youtu.be',
1101 'tiktok.com',
1102 'instagram.com', 'instagr.am',
1103 ]
1104 is_valid = any(p in url_lower for p in supported_patterns)
1105 platform = detect_platform(url) if is_valid else None
1106 return URLValidateResponse(valid=is_valid, platform=platform, url=url)
1107
1108
1109# ── Video Preview (metadata extraction without downloading) ──────────────────
1110
1111class PreviewResponse(BaseModel):
1112 title: Optional[str] = None
1113 thumbnail: Optional[str] = None
1114 duration: Optional[float] = None
1115 uploader: Optional[str] = None
1116 platform: Optional[str] = None
1117 success: bool = False
1118
1119
1120@app.post("/preview/url", response_model=PreviewResponse, tags=["URL Download"])
1121async def preview_url(payload: URLRequest):
1122 """
1123 Extract video metadata (title, thumbnail, duration) from a URL
1124 without downloading the full video. Powers the frontend preview card.
1125 """
1126 import yt_dlp
1127 url = payload.url.strip()
1128 if not url:
1129 return PreviewResponse(success=False)
1130
1131 platform = detect_platform(url)
1132
1133 def _extract():
1134 opts = {
1135 'quiet': True,
1136 'no_warnings': True,
1137 'skip_download': True,
1138 'ignoreerrors': True,
1139 'socket_timeout': 15,
1140 }
1141 with yt_dlp.YoutubeDL(opts) as ydl:
1142 info = ydl.extract_info(url, download=False)
1143 if info is None:
1144 return None
1145 return {
1146 'title': info.get('title'),
1147 'thumbnail': info.get('thumbnail'),
1148 'duration': info.get('duration'),
1149 'uploader': info.get('uploader') or info.get('channel'),
1150 }
1151
1152 try:
1153 result = await asyncio.to_thread(_extract)
1154 if result is None:
1155 return PreviewResponse(success=False, platform=platform)
1156
1157 # Proxy the thumbnail URL through our server to avoid CORS issues
1158 thumb_url = result.get('thumbnail')
1159 if thumb_url:
1160 # Encode and route through our proxy
1161 import urllib.parse
1162 proxied = f"/proxy/thumbnail?url={urllib.parse.quote(thumb_url, safe='')}"
1163 result['thumbnail'] = proxied
1164
1165 return PreviewResponse(
1166 success=True,
1167 platform=platform,
1168 **result,
1169 )
1170 except Exception as e:
1171 print(f"[Preview] Error extracting metadata: {e}")
1172 return PreviewResponse(success=False, platform=platform)
1173
1174
1175@app.get("/proxy/thumbnail", tags=["URL Download"])
1176async def proxy_thumbnail(url: str):
1177 """
1178 Proxy a thumbnail image to avoid CORS issues.
1179 The frontend calls this with the encoded thumbnail URL.
1180 """
1181 import httpx
1182 from fastapi.responses import Response
1183
1184 try:
1185 async with httpx.AsyncClient(follow_redirects=True, timeout=10.0) as client:
1186 resp = await client.get(url, headers={
1187 'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36'
1188 })
1189 resp.raise_for_status()
1190 content_type = resp.headers.get('content-type', 'image/jpeg')
1191 return Response(
1192 content=resp.content,
1193 media_type=content_type,
1194 headers={'Cache-Control': 'public, max-age=3600'}
1195 )
1196 except Exception:
1197 raise HTTPException(status_code=502, detail="Failed to fetch thumbnail")
1198
1199
1200@app.post("/predict/url", response_model=PredictionResponse, tags=["URL Download"])
