Team Ai
Apppublic

NormalCable/brainrot-detector-api

sourceHugging Facecc-by-nc-nd-4.0updated 5mo agoView on Hugging Face
0likes
api_server.py1299 linesDownload Raw Back to root
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"])

Showing the first 1,200 of 1299 lines. Download the file for the rest.