{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Auto-generated single-file script for Kaggle. Do not edit directly —\n# edit the kaggle_ood_eval/ package and re-run build_kaggle_script.py.\n\n\n# ==== config.py ====\n# kaggle_ood_eval/config.py\nimport os\n\n# ---- Fill these in with your actual Kaggle dataset attachment paths ----\nFAKEAVCELEB_FRAMES_ROOT = \"/kaggle/input/fakeav/frames/frames\"\nFAKEAVCELEB_AUDIO_ROOT = \"/kaggle/input/fakeav/AUDIO/AUDIO\"\nLAVDF_ROOT = \"/kaggle/input/localized-audio-visual-deepfake-dataset-lav-df/LAV-DF\"\nLAVDF_METADATA = os.path.join(LAVDF_ROOT, \"metadata.min.json\")\nDFDC_SUBSET0_ROOT = \"/kaggle/input/datasets/pranay22077/dfdc-10/dfdc_train_part_00/dfdc_train_part_0\"\nDFDC_SUBSET1_ROOT = \"/kaggle/input/datasets/pranay22077/dfdc-10/dfdc_train_part_01/dfdc_train_part_1\"\nIRISH_ROOT = \"/kaggle/input/genderwise-irish-accent-english-deepfake-detection/GenderWise/GenderWise_Video\"\n\n# ---- Model roster (fixed order: base models first, fusion models last) ----\nBASE_AUDIO_IDS = [\"audio_wavlm\", \"audio_wav2vec2\", \"audio_hubert\"]\nBASE_IMAGE_IDS = [\"image_vit\", \"image_vgg19\", \"image_resnet50\"]\nFUSION_IDS = [\"fusion_concat\", \"fusion_crossattn\", \"fusion_gated\"]\nMODEL_IDS = BASE_AUDIO_IDS + BASE_IMAGE_IDS + FUSION_IDS\n\nAUDIO_PRETRAINED = {\n    \"audio_wavlm\": \"microsoft/wavlm-base\",\n    \"audio_wav2vec2\": \"facebook/wav2vec2-base\",\n    \"audio_hubert\": \"facebook/hubert-base-ls960\",\n}\nIMAGE_BACKBONE = {\n    \"image_vit\": \"vit\",\n    \"image_vgg19\": \"vgg19\",\n    \"image_resnet50\": \"resnet50\",\n}\n\nOOD_DATASETS = [\"LAV-DF\", \"DFDC-0\", \"DFDC-1\", \"Irish\"]\n\n# Cap each OOD dataset's evaluation set size so eval time stays bounded and\n# counted against the weekly GPU budget instead of running unbounded. Kept\n# below 1500 (rather than using the full datasets) partly for wall-clock\n# reasons and partly to keep the on-disk decode cache (see CACHE_DIR below)\n# comfortably under Kaggle's ~19.5GB /kaggle/working output quota once all\n# four OOD sets plus FakeAVCeleb training samples are cached as uint8 frame\n# tensors.\nOOD_MAX_SAMPLES = 800\n\n# Video/audio decode results (cv2 frame extraction, torchaudio load/resample)\n# are the actual eval-time bottleneck, not GPU compute -- profiling showed\n# batches stalling every few iterations while DataLoader workers wait on\n# video decode. Since the same OOD datasets get evaluated once per model\n# across all 9 models in the sweep, caching each decoded sample to disk on\n# first access turns 8 of those 9 passes into a fast disk read instead of a\n# fresh decode. Set to None to disable caching (e.g. if disk space is tight).\nCACHE_DIR = \"/kaggle/working/decode_cache\"\n\n# Cap the FakeAVCeleb training/validation set size so each model's training\n# epoch is fast. Training is capped to a fixed wall-clock budget per model\n# regardless (see MIN/MAX_MODEL_BUDGET_SECONDS below), but a smaller dataset\n# means more optimizer steps actually happen within that budget.\nTRAIN_MAX_SAMPLES = 800\nTRAIN_VAL_MAX_SAMPLES = 150\n\n# ---- Budget: 30 GPU-hr/week Kaggle quota, self-limited with safety margins ----\nWEEKLY_BUDGET_SECONDS = 29 * 3600      # 1hr safety margin under the 30hr quota\nSESSION_BUDGET_SECONDS = 8.5 * 3600    # Kaggle sessions auto-stop around 9-12hrs\n# Training is capped at a fixed ~1 hour per model (not budget-adaptive) so the\n# full 9-model sweep completes in a predictable amount of time. Min and max\n# are set equal to make this a fixed cap rather than a range; the shared\n# _run_phase/BudgetManager logic still clamps this down if the weekly or\n# session budget remaining is less than an hour.\nMIN_MODEL_BUDGET_SECONDS = 1800\nMAX_MODEL_BUDGET_SECONDS = 1800\n\nCHECKPOINT_DIR = \"/kaggle/working/checkpoints\"\nRUN_STATE_PATH = \"/kaggle/working/run_state.json\"\nRESULTS_PATH = \"/kaggle/working/ood_results.json\"\n# If resuming in a new Kaggle session, attach the previous run's output as an input\n# dataset and point this at it so checkpoints + run_state carry over.\nRESUME_INPUT_DIR = os.environ.get(\"OOD_RESUME_INPUT_DIR\", None)\n\n\n# ==== budget.py ====\nimport json\nimport os\nimport time\n\n\nclass RunState:\n    def __init__(self, done=None, metrics=None, seconds_used_total=0.0,\n                 best_audio_id=None, best_image_id=None):\n        self.done = done or {}\n        self.metrics = metrics or {}\n        self.seconds_used_total = seconds_used_total\n        self.best_audio_id = best_audio_id\n        self.best_image_id = best_image_id\n\n    @classmethod\n    def load(cls, path, resume_input_dir=None):\n        if os.path.exists(path):\n            with open(path) as f:\n                data = json.load(f)\n            return cls(**data)\n        if resume_input_dir:\n            prev_path = os.path.join(resume_input_dir, \"run_state.json\")\n            if os.path.exists(prev_path):\n                with open(prev_path) as f:\n                    data = json.load(f)\n                return cls(**data)\n        return cls()\n\n    def is_done(self, model_id):\n        return self.done.get(model_id, False)\n\n    def mark_done(self, model_id, metrics):\n        self.done[model_id] = True\n        self.metrics[model_id] = metrics\n\n    def set_best_pair(self, audio_id, image_id):\n        self.best_audio_id = audio_id\n        self.best_image_id = image_id\n\n    def save(self, path):\n        os.makedirs(os.path.dirname(path), exist_ok=True)\n        with open(path, \"w\") as f:\n            json.dump({\n                \"done\": self.done,\n                \"metrics\": self.metrics,\n                \"seconds_used_total\": self.seconds_used_total,\n                \"best_audio_id\": self.best_audio_id,\n                \"best_image_id\": self.best_image_id,\n            }, f, indent=2)\n\n\nclass BudgetManager:\n    def __init__(self, run_state, weekly_budget_s, session_budget_s,\n                 min_model_s, max_model_s, clock=time.monotonic):\n        self.run_state = run_state\n        self.weekly_budget_s = weekly_budget_s\n        self.session_budget_s = session_budget_s\n        self.min_model_s = min_model_s\n        self.max_model_s = max_model_s\n        self.clock = clock\n        self._session_start = None\n\n    def session_start(self):\n        self._session_start = self.clock()\n\n    def session_elapsed(self):\n        return self.clock() - self._session_start\n\n    def session_remaining(self):\n        return self.session_budget_s - self.session_elapsed()\n\n    def weekly_remaining(self):\n        return self.weekly_budget_s - self.run_state.seconds_used_total\n\n    def budget_for_next_model(self, remaining_model_ids):\n        remaining = min(self.weekly_remaining(), self.session_remaining())\n        remaining = max(0.0, remaining)\n        even_share = remaining / max(1, len(remaining_model_ids))\n        # Try to give at least min_model_s, but never grant more than what's\n        # actually remaining -- otherwise a near-exhausted budget (e.g. 1s left)\n        # would still hand out a full MIN_MODEL_BUDGET_SECONDS and blow the\n        # weekly/session safety margin.\n        floored = max(self.min_model_s, even_share)\n        capped = min(self.max_model_s, floored)\n        return min(capped, remaining)\n\n    def should_stop_session(self):\n        return self.session_remaining() <= 0\n\n\n# ==== data.py ====\n# kaggle_ood_eval/data.py\nimport hashlib\nimport json\nimport os\nimport random\nfrom pathlib import Path\n\nimport cv2\nimport numpy as np\nimport torch\nimport torchaudio\nimport torchaudio.transforms as T\nfrom torch.utils.data import Dataset\n\n\ndef _decode_cache_path(cache_dir, path, tag):\n    \"\"\"Stable cache filename for a decoded (video or audio) sample.\n\n    The same OOD dataset is decoded once per model across a 9-model sweep\n    (evaluate_ood is called once per model against the same underlying\n    files), so caching the decoded tensor on first access and reusing it on\n    every later access turns 8 of those 9 passes' decode cost into a fast\n    disk read instead of a fresh cv2/torchaudio decode.\"\"\"\n    key = hashlib.sha1(f\"{path}|{tag}\".encode()).hexdigest()\n    return os.path.join(cache_dir, f\"{key}.pt\")\n\n\ndef load_audio_waveform(path, sr=16000, max_seconds=5, random_crop=False, cache_dir=None):\n    max_len = sr * max_seconds\n\n    # Only the deterministic (non-augmented) path is cacheable: random_crop\n    # intentionally varies per call for training-time augmentation, so\n    # caching it would freeze the crop position across epochs.\n    use_cache = cache_dir is not None and not random_crop\n    cache_path = None\n    if use_cache:\n        cache_path = _decode_cache_path(cache_dir, path, f\"audio_sr{sr}_len{max_seconds}\")\n        if os.path.exists(cache_path):\n            try:\n                cached = torch.load(cache_path)\n                return cached[\"waveform\"], cached[\"total_len\"]\n            except Exception:\n                pass\n\n    try:\n        waveform, orig_sr = torchaudio.load(path, normalize=True)\n        if orig_sr != sr:\n            waveform = torchaudio.functional.resample(waveform, orig_sr, sr)\n        if waveform.shape[0] > 1:\n            waveform = waveform.mean(dim=0, keepdim=True)\n        waveform = waveform.squeeze(0)\n        total_len = waveform.shape[0]\n        if random_crop and total_len > max_len:\n            start = random.randint(0, total_len - max_len)\n            waveform = waveform[start:start + max_len]\n        else:\n            waveform = waveform[:max_len]\n        if waveform.shape[0] < max_len:\n            waveform = torch.nn.functional.pad(waveform, (0, max_len - waveform.shape[0]))\n        total_len = min(total_len, max_len)\n        if use_cache:\n            os.makedirs(cache_dir, exist_ok=True)\n            try:\n                torch.save({\"waveform\": waveform, \"total_len\": total_len}, cache_path)\n            except Exception:\n                pass\n        return waveform, total_len\n    except Exception:\n        return torch.zeros(max_len), max_len\n\n\ndef load_video_frames(path, n_frames=16, size=224, cache_dir=None):\n    cache_path = None\n    if cache_dir is not None:\n        cache_path = _decode_cache_path(cache_dir, path, f\"video_n{n_frames}_s{size}\")\n        if os.path.exists(cache_path):\n            try:\n                cached = torch.load(cache_path)\n                return cached.float() / 255.0\n            except Exception:\n                pass\n\n    frames = _decode_video_frames(path, n_frames, size)\n\n    if cache_dir is not None:\n        os.makedirs(cache_dir, exist_ok=True)\n        try:\n            # Store as uint8 (0-255) rather than normalized float32 to keep\n            # the on-disk cache roughly 4x smaller, given Kaggle's working\n            # directory has a limited output quota.\n            torch.save((frames * 255.0).to(torch.uint8), cache_path)\n        except Exception:\n            pass\n    return frames\n\n\ndef _decode_video_frames(path, n_frames, size):\n    try:\n        cap = cv2.VideoCapture(path)\n        if not cap.isOpened():\n            return torch.zeros(n_frames, 3, size, size)\n        total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))\n        if total_frames <= 0:\n            cap.release()\n            return torch.zeros(n_frames, 3, size, size)\n        indices = set(np.linspace(0, total_frames - 1, n_frames).astype(int))\n        frames = {}\n        frame_idx = 0\n        while len(frames) < n_frames:\n            ret, frame = cap.read()\n            if not ret:\n                break\n            if frame_idx in indices:\n                frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)\n                frame = cv2.resize(frame, (size, size))\n                frames[frame_idx] = torch.from_numpy(frame).permute(2, 0, 1).float() / 255.0\n            frame_idx += 1\n        cap.release()\n        if len(frames) == 0:\n            return torch.zeros(n_frames, 3, size, size)\n        frame_list = [frames[k] for k in sorted(frames.keys())]\n        while len(frame_list) < n_frames:\n            frame_list.append(frame_list[-1])\n        return torch.stack(frame_list[:n_frames])\n    except Exception:\n        return torch.zeros(n_frames, 3, size, size)\n\n\nclass FakeAVCelebDataset(Dataset):\n    def set_train_mode(self, mode: bool):\n        self.train_mode = mode\n\n    def __init__(self, frames_root, audio_root, transform=None, mel_transform=None,\n                 downsample_frames=2, sample_rate=16000, n_mels=64, max_audio_time=500,\n                 cache_dir=None):\n        self.augment_prob = 0.0\n        self.max_audio_time = max_audio_time\n        self.frames_root = frames_root\n        self.audio_root = audio_root\n        self.cache_dir = cache_dir\n        self.transform = transform\n        self.mel_transform = mel_transform if mel_transform is not None else torchaudio.transforms.MelSpectrogram(\n            sample_rate=sample_rate, n_mels=n_mels)\n        self.downsample_frames = downsample_frames\n        self.train_mode = True\n        self.freq_mask = T.FrequencyMasking(freq_mask_param=15)\n        self.time_mask = T.TimeMasking(time_mask_param=50)\n        self.label_map = {\n            \"FakeVideo-FakeAudio\": 0,\n            \"FakeVideo-RealAudio\": 1,\n            \"RealVideo-FakeAudio\": 2,\n            \"RealVideo-RealAudio\": 3\n        }\n        self.samples = []\n        for label_name, label_idx in self.label_map.items():\n            frames_category = os.path.join(frames_root, label_name)\n            audio_category = os.path.join(audio_root, label_name)\n            if not os.path.exists(frames_category) or not os.path.exists(audio_category):\n                continue\n            for ethnicity in os.listdir(frames_category):\n                ethnicity_path = os.path.join(frames_category, ethnicity)\n                for gender in os.listdir(ethnicity_path):\n                    gender_path = os.path.join(ethnicity_path, gender)\n                    audio_gender_path = os.path.join(audio_category, ethnicity, gender)\n                    for person_id in os.listdir(gender_path):\n                        person_frames_path = os.path.join(gender_path, person_id)\n                        person_audio_path = os.path.join(audio_gender_path, person_id)\n                        if not os.path.exists(person_audio_path):\n                            continue\n                        for sample_id in os.listdir(person_frames_path):\n                            frame_path = os.path.join(person_frames_path, sample_id)\n                            audio_path = os.path.join(person_audio_path, sample_id)\n                            if not os.path.exists(audio_path):\n                                # Some FakeAVCeleb layouts name the audio file\n                                # \"<sample_id>.wav\" rather than a bare match.\n                                wav_audio_path = audio_path + \".wav\"\n                                if os.path.exists(wav_audio_path):\n                                    audio_path = wav_audio_path\n                                else:\n                                    continue\n                            self.samples.append({\n                                \"frame_path\": frame_path,\n                                \"audio_path\": audio_path,\n                                \"label\": label_idx,\n                            })\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        video = load_video_frames(sample[\"frame_path\"], n_frames=16, cache_dir=self.cache_dir)\n        audio, audio_len = load_audio_waveform(sample[\"audio_path\"], cache_dir=self.cache_dir)\n        label = 0 if sample[\"label\"] == 3 else 1\n        return {\n            \"video\": video,\n            \"audio\": audio,\n            \"audio_len\": torch.tensor(audio_len, dtype=torch.long),\n            \"label\": torch.tensor(label, dtype=torch.long),\n        }\n\n\nclass DFDCDataset(torch.utils.data.Dataset):\n    def __init__(self, root, cache_dir=None):\n        self.cache_dir = cache_dir\n        self.samples = []\n        meta_path = os.path.join(root, \"metadata.json\")\n        assert os.path.exists(meta_path), f\"metadata.json not found in {root}\"\n        with open(meta_path, \"r\") as f:\n            meta = json.load(f)\n        for fname, info in meta.items():\n            if info.get(\"split\", \"train\") != \"train\":\n                continue\n            video_path = os.path.join(root, fname)\n            if not os.path.exists(video_path):\n                continue\n            label = 0 if info[\"label\"] == \"REAL\" else 1\n            self.samples.append((video_path, label))\n        print(f\"[DFDC] Loaded {len(self.samples)} samples from {root}\")\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx, _attempt=0):\n        video_path, label = self.samples[idx]\n        try:\n            video = load_video_frames(video_path, n_frames=16, cache_dir=self.cache_dir)\n            audio, audio_len = load_audio_waveform(video_path, cache_dir=self.cache_dir)\n            if audio.abs().sum() < 1e-6:\n                raise ValueError(\"Silent audio\")\n        except Exception:\n            if _attempt >= 10:\n                # Bounded retries exhausted -- fall back to a zero-tensor sample\n                # rather than risking unbounded recursion on repeated failures.\n                video = torch.zeros(16, 3, 224, 224)\n                audio = torch.zeros(16000 * 5)\n                return {\n                    \"video\": video,\n                    \"audio\": audio,\n                    \"audio_len\": torch.tensor(16000 * 5, dtype=torch.long),\n                    \"label\": torch.tensor(label, dtype=torch.long),\n                }\n            new_idx = np.random.randint(0, len(self.samples))\n            return self.__getitem__(new_idx, _attempt=_attempt + 1)\n        return {\n            \"video\": video,\n            \"audio\": audio,\n            \"audio_len\": torch.tensor(audio_len, dtype=torch.long),\n            \"label\": torch.tensor(label, dtype=torch.long),\n        }\n\n\nclass LAVDFOODDataset(torch.utils.data.Dataset):\n    def __init__(self, root, metadata_path, split=\"test\", cache_dir=None):\n        self.root = root\n        self.cache_dir = cache_dir\n        with open(metadata_path, \"r\") as f:\n            self.meta = json.load(f)\n        self.samples = [m for m in self.meta if m[\"split\"] == split]\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        m = self.samples[idx]\n        video_path = os.path.join(self.root, m[\"file\"])\n        if not os.path.exists(video_path):\n            raise FileNotFoundError(video_path)\n        video = load_video_frames(video_path, n_frames=16, cache_dir=self.cache_dir)\n        audio, audio_len = load_audio_waveform(video_path, cache_dir=self.cache_dir)\n        label = int(m[\"modify_audio\"] or m[\"modify_video\"])\n        return {\n            \"video\": video,\n            \"audio\": audio,\n            \"audio_len\": torch.tensor(audio_len, dtype=torch.long),\n            \"label\": torch.tensor(label, dtype=torch.long),\n        }\n\n\nclass IrishOODDataset(torch.utils.data.Dataset):\n    def __init__(self, root, cache_dir=None):\n        self.cache_dir = cache_dir\n        self.samples = []\n        for label_name, label in [('Real', 0), ('Fake', 1)]:\n            for path in Path(root, label_name).rglob(\"*.mp4\"):\n                self.samples.append((str(path), label))\n\n    def __getitem__(self, idx):\n        video_path, label = self.samples[idx]\n        video = load_video_frames(video_path, n_frames=16, cache_dir=self.cache_dir)\n        audio, audio_len = load_audio_waveform(video_path, cache_dir=self.cache_dir)\n        return {\n            \"video\": video,\n            \"audio\": audio,\n            \"audio_len\": torch.tensor(audio_len, dtype=torch.long),\n            \"label\": torch.tensor(label, dtype=torch.long),\n        }\n\n    def __len__(self):\n        return len(self.samples)\n\n\ndef make_fakeavceleb_train_val(frames_root, audio_root, val_fraction=0.15, seed=42,\n                                max_train_samples=None, max_val_samples=None, cache_dir=None):\n    from torch.utils.data import random_split\n    full = FakeAVCelebDataset(frames_root, audio_root, cache_dir=cache_dir)\n    n_val = int(len(full) * val_fraction)\n    n_train = len(full) - n_val\n    generator = torch.Generator().manual_seed(seed)\n    train_ds, val_ds = random_split(full, [n_train, n_val], generator=generator)\n    train_ds = _cap_dataset(train_ds, max_train_samples)\n    val_ds = _cap_dataset(val_ds, max_val_samples)\n    return train_ds, val_ds\n\n\ndef _cap_dataset(ds, max_samples):\n    \"\"\"Deterministically cap a dataset to its first `max_samples` indices so\n    OOD evaluation time stays bounded (and countable against the GPU budget).\"\"\"\n    if max_samples is not None and len(ds) > max_samples:\n        return torch.utils.data.Subset(ds, list(range(max_samples)))\n    return ds\n\n\ndef build_ood_loaders(cfg):\n    from torch.utils.data import DataLoader\n    cache_dir = getattr(cfg, \"CACHE_DIR\", None)\n    lavdf_ds = LAVDFOODDataset(root=LAVDF_ROOT, metadata_path=LAVDF_METADATA, split=\"test\", cache_dir=cache_dir)\n    dfdc0_ds = DFDCDataset(DFDC_SUBSET0_ROOT, cache_dir=cache_dir)\n    dfdc1_ds = DFDCDataset(DFDC_SUBSET1_ROOT, cache_dir=cache_dir)\n    irish_ds = IrishOODDataset(root=IRISH_ROOT, cache_dir=cache_dir)\n\n    max_samples = getattr(cfg, \"OOD_MAX_SAMPLES\", None)\n    lavdf_ds = _cap_dataset(lavdf_ds, max_samples)\n    dfdc0_ds = _cap_dataset(dfdc0_ds, max_samples)\n    dfdc1_ds = _cap_dataset(dfdc1_ds, max_samples)\n    irish_ds = _cap_dataset(irish_ds, max_samples)\n\n    def _loader(ds):\n        return DataLoader(ds, batch_size=8, shuffle=False, num_workers=4, pin_memory=True)\n\n    return {\n        \"LAV-DF\": _loader(lavdf_ds),\n        \"DFDC-0\": _loader(dfdc0_ds),\n        \"DFDC-1\": _loader(dfdc1_ds),\n        \"Irish\": _loader(irish_ds),\n    }\n\n\n# ==== models_audio.py ====\nimport torch\nimport torch.nn as nn\nfrom transformers import Wav2Vec2Model, WavLMModel, HubertModel\n\n_MODEL_CLASSES = {\n    \"wav2vec2\": Wav2Vec2Model,\n    \"wavlm\": WavLMModel,\n    \"hubert\": HubertModel,\n}\n\n\nclass AudioBaseClassifier(nn.Module):\n    def __init__(self, pretrained_name, model_type, unfreeze_last_n=2, proj_dim=256):\n        super().__init__()\n        model_cls = _MODEL_CLASSES[model_type]\n        self.encoder = model_cls.from_pretrained(pretrained_name)\n\n        for p in self.encoder.parameters():\n            p.requires_grad = False\n        n_layers = len(self.encoder.encoder.layers)\n        for layer in self.encoder.encoder.layers[max(0, n_layers - unfreeze_last_n):]:\n            for p in layer.parameters():\n                p.requires_grad = True\n\n        hidden_size = self.encoder.config.hidden_size\n        self.proj = nn.Sequential(nn.Linear(hidden_size, proj_dim), nn.LayerNorm(proj_dim))\n        self.head = nn.Sequential(\n            nn.Linear(proj_dim, 128), nn.ReLU(), nn.Dropout(0.2), nn.Linear(128, 1)\n        )\n\n    def forward(self, audio_waveform, attention_mask=None):\n        outputs = self.encoder(\n            audio_waveform, attention_mask=attention_mask,\n            output_hidden_states=False, return_dict=True,\n        )\n        hidden = outputs.last_hidden_state  # (B, T, H)\n\n        if attention_mask is not None:\n            feat_mask = self.encoder._get_feature_vector_attention_mask(\n                hidden.shape[1], attention_mask\n            ).unsqueeze(-1).float()\n            emb = (hidden * feat_mask).sum(dim=1) / feat_mask.sum(dim=1).clamp(min=1e-6)\n        else:\n            emb = hidden.mean(dim=1)\n\n        emb = self.proj(emb)\n        logits = self.head(emb)\n        return {\"audio_logits\": logits, \"embedding\": emb}\n\n\ndef build_audio_model(model_id, cfg, unfreeze_last_n=2, proj_dim=256):\n    return AudioBaseClassifier(\n        pretrained_name=AUDIO_PRETRAINED[model_id],\n        model_type=model_id.split(\"_\", 1)[1],\n        unfreeze_last_n=unfreeze_last_n,\n        proj_dim=proj_dim,\n    )\n\n\n# ==== models_image.py ====\nimport torch\nimport torch.nn as nn\nimport torchvision.models as tvm\nfrom transformers import ViTModel, ViTConfig\n\n\ndef _locate_vit_blocks(vit_model):\n    \"\"\"Returns the ViT transformer block list, tolerating the two attribute\n    layouts seen across transformers releases: some versions expose the\n    blocks directly as ViTModel.layers, others nest them under\n    ViTModel.encoder.layer (singular). This avoids a silent AttributeError if\n    a different transformers version is preinstalled on Kaggle than the one\n    used in development.\"\"\"\n    if hasattr(vit_model, \"layers\"):\n        return list(vit_model.layers)\n    if hasattr(vit_model, \"encoder\") and hasattr(vit_model.encoder, \"layer\"):\n        return list(vit_model.encoder.layer)\n    raise AttributeError(\n        \"Could not locate ViT transformer block list on this transformers \"\n        \"version (checked .layers and .encoder.layer).\"\n    )\n\n\nclass ImageBaseClassifier(nn.Module):\n    def __init__(self, backbone, unfreeze_last_n=1, proj_dim=256, pretrained=True):\n        super().__init__()\n        self.backbone_name = backbone\n\n        if backbone == \"vit\":\n            self.encoder = ViTModel.from_pretrained(\"google/vit-base-patch16-224\") \\\n                if pretrained else ViTModel(ViTConfig())\n            out_dim = self.encoder.config.hidden_size\n            self._forward_backbone = self._forward_vit\n            unfreeze_targets = _locate_vit_blocks(self.encoder)\n        elif backbone == \"vgg19\":\n            weights = tvm.VGG19_Weights.IMAGENET1K_V1 if pretrained else None\n            full = tvm.vgg19(weights=weights)\n            self.encoder = full.features\n            self.avgpool = full.avgpool\n            out_dim = 512 * 7 * 7\n            self._forward_backbone = self._forward_vgg\n            # Group features into conv blocks split on MaxPool2d boundaries\n            # (VGG19 has 5 such blocks), so unfreeze_last_n counts real\n            # conv blocks rather than individual (often parameterless) layers.\n            blocks = []\n            current_block = []\n            for layer in self.encoder:\n                current_block.append(layer)\n                if isinstance(layer, nn.MaxPool2d):\n                    blocks.append(current_block)\n                    current_block = []\n            if current_block:\n                blocks.append(current_block)\n            unfreeze_targets = [nn.ModuleList(block) for block in blocks]\n        elif backbone == \"resnet50\":\n            weights = tvm.ResNet50_Weights.IMAGENET1K_V2 if pretrained else None\n            full = tvm.resnet50(weights=weights)\n            self.encoder = nn.Sequential(*list(full.children())[:-1])  # drop fc\n            out_dim = full.fc.in_features\n            self._forward_backbone = self._forward_resnet\n            # Unfreeze from the actual residual stage modules, not from\n            # self.encoder (whose last element is the parameterless avgpool).\n            unfreeze_targets = [full.layer1, full.layer2, full.layer3, full.layer4]\n        else:\n            raise ValueError(f\"Unknown backbone {backbone}\")\n\n        for p in self.encoder.parameters():\n            p.requires_grad = False\n        for module in unfreeze_targets[-unfreeze_last_n:]:\n            for p in module.parameters():\n                p.requires_grad = True\n\n        self.proj = nn.Sequential(nn.Linear(out_dim, proj_dim), nn.LayerNorm(proj_dim))\n        self.head = nn.Sequential(\n            nn.Linear(proj_dim, 128), nn.ReLU(), nn.Dropout(0.2), nn.Linear(128, 1)\n        )\n\n    def _forward_vit(self, flat_frames):\n        return self.encoder(flat_frames).last_hidden_state[:, 0]  # CLS token\n\n    def _forward_vgg(self, flat_frames):\n        feats = self.encoder(flat_frames)\n        feats = self.avgpool(feats)\n        return feats.flatten(1)\n\n    def _forward_resnet(self, flat_frames):\n        return self.encoder(flat_frames).flatten(1)\n\n    def forward(self, video):\n        B, T, C, H, W = video.shape\n        flat = video.view(B * T, C, H, W)\n        feats = self._forward_backbone(flat)          # (B*T, out_dim)\n        feats = feats.view(B, T, -1).mean(dim=1)       # (B, out_dim)\n        emb = self.proj(feats)\n        logits = self.head(emb)\n        return {\"video_logits\": logits, \"embedding\": emb}\n\n\ndef build_image_model(model_id, cfg, unfreeze_last_n=1, proj_dim=256):\n    return ImageBaseClassifier(backbone=IMAGE_BACKBONE[model_id], unfreeze_last_n=unfreeze_last_n, proj_dim=proj_dim, pretrained=True)\n\n\n# ==== models_fusion.py ====\nimport torch\nimport torch.nn as nn\n\n\n\n\ndef attention_entropy(attn_weights, eps=1e-9):\n    attn_weights = torch.nan_to_num(attn_weights, nan=0.0)\n    attn_weights = attn_weights / (attn_weights.sum(dim=-1, keepdim=True) + eps)\n    H = -(attn_weights * torch.log(attn_weights + eps)).sum(dim=-1)\n    return H.mean(dim=1)\n\n\nclass ModalityReliabilityGate(nn.Module):\n    def __init__(self, dim=256):\n        super().__init__()\n        self.video_rel = nn.Sequential(nn.Linear(dim + 1, 64), nn.ReLU(), nn.Linear(64, 1))\n        self.audio_rel = nn.Sequential(nn.Linear(dim + 1, 64), nn.ReLU(), nn.Linear(64, 1))\n\n    def forward(self, v2a, a2v, ent_v, ent_a):\n        v2a_in = torch.cat([v2a, ent_v.unsqueeze(1)], dim=1)\n        a2v_in = torch.cat([a2v, ent_a.unsqueeze(1)], dim=1)\n        score_v = self.video_rel(v2a_in)\n        score_a = self.audio_rel(a2v_in)\n        weights = torch.softmax(torch.cat([score_v, score_a], dim=1) / 1.5, dim=1)\n        w_v, w_a = weights[:, 0:1], weights[:, 1:2]\n        fused = w_v * v2a + w_a * a2v\n        return fused, w_v, w_a\n\n\nclass _BaseFusion(nn.Module):\n    \"\"\"Shared setup: builds audio+image encoders/heads reused by all fusion variants.\"\"\"\n\n    def __init__(self, audio_model_id, image_model_id, cfg, proj_dim=256, pretrained_image=True):\n        super().__init__()\n        self.audio_base = AudioBaseClassifier(\n            pretrained_name=AUDIO_PRETRAINED[audio_model_id],\n            model_type=audio_model_id.split(\"_\", 1)[1],\n            unfreeze_last_n=2, proj_dim=proj_dim,\n        )\n        self.image_base = ImageBaseClassifier(\n            backbone=IMAGE_BACKBONE[image_model_id],\n            unfreeze_last_n=1, proj_dim=proj_dim, pretrained=pretrained_image,\n        )\n        self.audio_head = nn.Sequential(nn.Linear(proj_dim, 128), nn.ReLU(), nn.Dropout(0.2), nn.Linear(128, 1))\n        self.video_head = nn.Sequential(nn.Linear(proj_dim, 128), nn.ReLU(), nn.Dropout(0.2), nn.Linear(128, 1))\n\n    def _embed(self, video, audio, attention_mask=None):\n        audio_out = self.audio_base(audio, attention_mask=attention_mask)\n        video_out = self.image_base(video)\n        return audio_out[\"embedding\"], video_out[\"embedding\"]\n\n\nclass ConcatFusionModel(_BaseFusion):\n    def __init__(self, audio_model_id, image_model_id, cfg, proj_dim=256, pretrained_image=True):\n        super().__init__(audio_model_id, image_model_id, cfg, proj_dim, pretrained_image)\n        self.fusion_fc = nn.Linear(proj_dim * 2, proj_dim)\n        self.fusion_head = nn.Sequential(nn.Linear(proj_dim, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, 1))\n\n    def forward(self, video, audio, attention_mask=None):\n        a_emb, v_emb = self._embed(video, audio, attention_mask)\n        fused = self.fusion_fc(torch.cat([a_emb, v_emb], dim=1))\n        return {\n            \"audio_logits\": self.audio_head(a_emb),\n            \"video_logits\": self.video_head(v_emb),\n            \"fusion_logits\": self.fusion_head(fused),\n        }\n\n\nclass CrossAttnFusionModel(_BaseFusion):\n    def __init__(self, audio_model_id, image_model_id, cfg, proj_dim=256, pretrained_image=True):\n        super().__init__(audio_model_id, image_model_id, cfg, proj_dim, pretrained_image)\n        self.video_to_audio_attn = nn.MultiheadAttention(embed_dim=proj_dim, num_heads=4, batch_first=True)\n        self.audio_to_video_attn = nn.MultiheadAttention(embed_dim=proj_dim, num_heads=4, batch_first=True)\n        self.fusion_fc = nn.Linear(proj_dim * 2, proj_dim)\n        self.fusion_head = nn.Sequential(nn.Linear(proj_dim, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, 1))\n\n    def forward(self, video, audio, attention_mask=None):\n        a_emb, v_emb = self._embed(video, audio, attention_mask)\n        a_seq, v_seq = a_emb.unsqueeze(1), v_emb.unsqueeze(1)  # (B,1,D) — pooled tokens as length-1 sequences\n        v2a, _ = self.video_to_audio_attn(v_seq, a_seq, a_seq)\n        a2v, _ = self.audio_to_video_attn(a_seq, v_seq, v_seq)\n        fused = self.fusion_fc(torch.cat([v2a.squeeze(1), a2v.squeeze(1)], dim=1))\n        return {\n            \"audio_logits\": self.audio_head(a_emb),\n            \"video_logits\": self.video_head(v_emb),\n            \"fusion_logits\": self.fusion_head(fused),\n        }\n\n\nclass GatedFusionModel(_BaseFusion):\n    def __init__(self, audio_model_id, image_model_id, cfg, proj_dim=256, pretrained_image=True):\n        super().__init__(audio_model_id, image_model_id, cfg, proj_dim, pretrained_image)\n        self.video_to_audio_attn = nn.MultiheadAttention(embed_dim=proj_dim, num_heads=4, batch_first=True)\n        self.audio_to_video_attn = nn.MultiheadAttention(embed_dim=proj_dim, num_heads=4, batch_first=True)\n        self.rel_gate = ModalityReliabilityGate(dim=proj_dim)\n        self.fusion_head = nn.Sequential(nn.Linear(proj_dim, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, 1))\n\n    def forward(self, video, audio, attention_mask=None):\n        a_emb, v_emb = self._embed(video, audio, attention_mask)\n        a_seq, v_seq = a_emb.unsqueeze(1), v_emb.unsqueeze(1)\n        v2a, w_v2a = self.video_to_audio_attn(v_seq, a_seq, a_seq, average_attn_weights=True)\n        a2v, w_a2v = self.audio_to_video_attn(a_seq, v_seq, v_seq, average_attn_weights=True)\n        ent_v = attention_entropy(w_v2a)\n        ent_a = attention_entropy(w_a2v)\n        fused, w_v, w_a = self.rel_gate(v2a.squeeze(1), a2v.squeeze(1), ent_v, ent_a)\n        return {\n            \"audio_logits\": self.audio_head(a_emb),\n            \"video_logits\": self.video_head(v_emb),\n            \"fusion_logits\": self.fusion_head(fused),\n            \"w_v\": w_v, \"w_a\": w_a, \"ent_v\": ent_v, \"ent_a\": ent_a,\n        }\n\n\ndef build_fusion_model(fusion_id, audio_model_id, image_model_id, cfg, proj_dim=256):\n    return {\n        \"fusion_concat\": ConcatFusionModel,\n        \"fusion_crossattn\": CrossAttnFusionModel,\n        \"fusion_gated\": GatedFusionModel,\n    }[fusion_id](audio_model_id, image_model_id, cfg, proj_dim=proj_dim)\n\n\n# ==== metrics.py ====\nimport torch\nfrom tqdm import tqdm\nfrom sklearn.metrics import balanced_accuracy_score, roc_auc_score, roc_curve, f1_score, precision_score, recall_score\n\n\ndef evaluate_ood(model, loader, device, dataset_name=\"Dataset\"):\n    model.eval()\n    all_labels, all_probs_f = [], []\n    all_preds_f, all_preds_a, all_preds_v = [], [], []\n\n    with torch.no_grad():\n        for batch in tqdm(loader, desc=f\"Evaluating {dataset_name}\"):\n            videos = batch[\"video\"].to(device)\n            audios = batch[\"audio\"].to(device)\n            labels = batch[\"label\"].to(device).long().view(-1)\n\n            audio_lens = batch.get(\"audio_len\", None)\n            if audio_lens is not None:\n                audio_lens = audio_lens.to(device)\n                B, L = audios.shape\n                mask = torch.arange(L, device=device).unsqueeze(0)\n                attention_mask = (mask < audio_lens.unsqueeze(1))\n            else:\n                attention_mask = None\n\n            out = model(videos, audios, attention_mask=attention_mask)\n\n            probs_f = torch.sigmoid(out[\"fusion_logits\"]).view(-1)\n            probs_a = torch.sigmoid(out[\"audio_logits\"]).view(-1)\n            probs_v = torch.sigmoid(out[\"video_logits\"]).view(-1)\n\n            all_preds_f.append((probs_f > 0.5).long().cpu())\n            all_preds_a.append((probs_a > 0.5).long().cpu())\n            all_preds_v.append((probs_v > 0.5).long().cpu())\n            all_probs_f.append(probs_f.cpu())\n            all_labels.append(labels.cpu())\n\n    y_true = torch.cat(all_labels).numpy()\n    pred_f = torch.cat(all_preds_f).numpy()\n    pred_a = torch.cat(all_preds_a).numpy()\n    pred_v = torch.cat(all_preds_v).numpy()\n    probs_f = torch.cat(all_probs_f).numpy()\n\n    acc_f = balanced_accuracy_score(y_true, pred_f)\n    acc_a = balanced_accuracy_score(y_true, pred_a)\n    acc_v = balanced_accuracy_score(y_true, pred_v)\n    roc = roc_auc_score(y_true, probs_f) if len(set(y_true.tolist())) > 1 else float(\"nan\")\n\n    f1_f = f1_score(y_true, pred_f, zero_division=0)\n    f1_a = f1_score(y_true, pred_a, zero_division=0)\n    f1_v = f1_score(y_true, pred_v, zero_division=0)\n    prec_f = precision_score(y_true, pred_f, zero_division=0)\n    rec_f = recall_score(y_true, pred_f, zero_division=0)\n\n    fpr, tpr, thresholds = roc_curve(y_true, probs_f)\n    j_scores = tpr - fpr\n    best_thresh = float(thresholds[j_scores.argmax()])\n    pred_opt = (probs_f > best_thresh).astype(int)\n    acc_opt = balanced_accuracy_score(y_true, pred_opt)\n    f1_opt = f1_score(y_true, pred_opt, zero_division=0)\n\n    return {\n        \"acc_f\": float(acc_f), \"acc_a\": float(acc_a), \"acc_v\": float(acc_v),\n        \"roc\": float(roc), \"f1_f\": float(f1_f), \"f1_a\": float(f1_a), \"f1_v\": float(f1_v),\n        \"prec_f\": float(prec_f), \"rec_f\": float(rec_f),\n        \"best_thresh\": best_thresh, \"acc_opt\": float(acc_opt), \"f1_opt\": float(f1_opt),\n    }\n\n\nclass _SingleModalityAdapter(torch.nn.Module):\n    \"\"\"Wraps a single-modality base model so its one logits branch is exposed as\n    fusion/audio/video_logits, letting evaluate_ood run unmodified.\"\"\"\n\n    def __init__(self, base_model, modality):\n        super().__init__()\n        self.base_model = base_model\n        self.modality = modality\n\n    def forward(self, video, audio, attention_mask=None):\n        if self.modality == \"audio\":\n            out = self.base_model(audio, attention_mask=attention_mask)\n            logits = out[\"audio_logits\"]\n        else:\n            out = self.base_model(video)\n            logits = out[\"video_logits\"]\n        return {\"fusion_logits\": logits, \"audio_logits\": logits, \"video_logits\": logits}\n\n\ndef evaluate_single_modality(model, loader, device, modality, dataset_name=\"Dataset\"):\n    adapter = _SingleModalityAdapter(model, modality).to(device)\n    result = evaluate_ood(adapter, loader, device, dataset_name=dataset_name)\n    other = \"video\" if modality == \"audio\" else \"audio\"\n    result[f\"acc_{other[0]}\"] = None\n    result[f\"f1_{other[0]}\"] = None\n    return result\n\n\ndef results_to_markdown_table(results, ood_datasets):\n    header = \"| Model |\"\n    sep = \"|---|\"\n    for ds in ood_datasets:\n        header += f\" {ds} Acc | {ds} F1 |\"\n        sep += \"---|---|\"\n    rows = [header, sep]\n    for model_id, per_ds in results.items():\n        row = f\"| {model_id} |\"\n        for ds in ood_datasets:\n            metrics_dict = per_ds.get(ds, {})\n            acc = metrics_dict.get(\"acc_f\")\n            f1 = metrics_dict.get(\"f1_f\")\n            acc_str = f\"{acc:.2f}\" if acc is not None else \"N/A\"\n            f1_str = f\"{f1:.2f}\" if f1 is not None else \"N/A\"\n            row += f\" {acc_str} | {f1_str} |\"\n        rows.append(row)\n    return \"\\n\".join(rows)\n\n\n# ==== train.py ====\nimport time\n\nimport torch\nimport torch.nn as nn\n\n\ndef _default_forward_fn(model, batch, device):\n    video = batch[\"video\"].to(device)\n    audio = batch[\"audio\"].to(device)\n    return model(video, audio)\n\n\ndef train_one_model(model, train_loader, val_loader, device, budget_seconds,\n                     clock=time.monotonic, lr=1e-4, forward_fn=None):\n    forward_fn = forward_fn or _default_forward_fn\n    model.to(device)\n    model.train()\n    optimizer = torch.optim.Adam(\n        (p for p in model.parameters() if p.requires_grad), lr=lr\n    )\n    criterion = nn.BCEWithLogitsLoss()\n\n    if len(train_loader) == 0:\n        raise ValueError(\"train_loader is empty — no batches to train on\")\n\n    start = clock()\n    stop = False\n    while not stop:\n        for batch in train_loader:\n            if clock() - start >= budget_seconds:\n                stop = True\n                break\n            optimizer.zero_grad()\n            out = forward_fn(model, batch, device)\n            labels = batch[\"label\"].to(device).float().view(-1, 1)\n            loss = criterion(out[\"fusion_logits\"], labels)\n            loss.backward()\n            optimizer.step()\n        if clock() - start >= budget_seconds:\n            stop = True\n\n    model.eval()\n    return model\n\n\n# ==== main.py ====\n# kaggle_ood_eval/main.py\nimport json\nimport os\n\nimport torch\n\n\n\n\n\n\ndef select_best_pair(run_state):\n    def _mean_acc(model_id):\n        per_ds = run_state.metrics[model_id]\n        vals = [v[\"acc_f\"] for v in per_ds.values() if v.get(\"acc_f\") is not None]\n        return sum(vals) / len(vals) if vals else float(\"-inf\")\n\n    audio_ids = [m for m in run_state.metrics if m.startswith(\"audio_\")]\n    image_ids = [m for m in run_state.metrics if m.startswith(\"image_\")]\n    best_audio = max(audio_ids, key=_mean_acc)\n    best_image = max(image_ids, key=_mean_acc)\n    return best_audio, best_image\n\n\ndef _default_build_model_fn(model_id, cfg, audio_model_id=None, image_model_id=None):\n    if model_id in BASE_AUDIO_IDS:\n        return build_audio_model(model_id, cfg)\n    if model_id in BASE_IMAGE_IDS:\n        return build_image_model(model_id, cfg)\n    return build_fusion_model(model_id, audio_model_id, image_model_id, cfg)\n\n\ndef _build_attention_mask(batch, audios, device):\n    \"\"\"Replicates metrics.evaluate_ood's attention_mask construction so\n    training-time forward calls see the same masking as eval-time ones.\"\"\"\n    audio_lens = batch.get(\"audio_len\", None)\n    if audio_lens is not None:\n        audio_lens = audio_lens.to(device)\n        B, L = audios.shape\n        mask = torch.arange(L, device=device).unsqueeze(0)\n        return mask < audio_lens.unsqueeze(1)\n    return None\n\n\ndef _default_train_fn(model, model_id, budget_s, cfg, train_loader, val_loader, device):\n    forward_fn = None\n    if model_id in BASE_AUDIO_IDS:\n        def forward_fn(m, batch, dev):\n            audios = batch[\"audio\"].to(dev)\n            attention_mask = _build_attention_mask(batch, audios, dev)\n            out = m(audios, attention_mask=attention_mask)\n            out[\"fusion_logits\"] = out[\"audio_logits\"]\n            return out\n    elif model_id in BASE_IMAGE_IDS:\n        def forward_fn(m, batch, dev):\n            out = m(batch[\"video\"].to(dev))\n            out[\"fusion_logits\"] = out[\"video_logits\"]\n            return out\n    else:\n        def forward_fn(m, batch, dev):\n            video = batch[\"video\"].to(dev)\n            audios = batch[\"audio\"].to(dev)\n            attention_mask = _build_attention_mask(batch, audios, dev)\n            return m(video, audios, attention_mask=attention_mask)\n    return train_one_model(\n        model, train_loader, val_loader, device, budget_seconds=budget_s, forward_fn=forward_fn,\n    )\n\n\ndef _default_eval_fn(model, model_id, cfg, ood_loaders, device):\n    results = {}\n    for ds_name, loader in ood_loaders.items():\n        if model_id in BASE_AUDIO_IDS:\n            results[ds_name] = evaluate_single_modality(model, loader, device, \"audio\", ds_name)\n        elif model_id in BASE_IMAGE_IDS:\n            results[ds_name] = evaluate_single_modality(model, loader, device, \"video\", ds_name)\n        else:\n            results[ds_name] = evaluate_ood(model, loader, device, ds_name)\n    return results\n\n\ndef _run_phase(model_ids, cfg, run_state, budget_mgr, build_model_fn, train_fn, eval_fn,\n               ood_loaders, train_loader, val_loader, device, clock, extra_build_kwargs=None):\n    extra_build_kwargs = extra_build_kwargs or {}\n    remaining = [m for m in model_ids if not run_state.is_done(m)]\n    for model_id in list(remaining):\n        if budget_mgr.should_stop_session() or budget_mgr.weekly_remaining() <= 0:\n            break\n        model_budget = budget_mgr.budget_for_next_model(remaining)\n        model = build_model_fn(model_id, cfg, **extra_build_kwargs)\n\n        checkpoint_path = os.path.join(CHECKPOINT_DIR, f\"{model_id}.pth\")\n        if not os.path.exists(checkpoint_path):\n            resume_input_dir = getattr(cfg, \"RESUME_INPUT_DIR\", None)\n            if resume_input_dir:\n                prev_checkpoint_path = os.path.join(resume_input_dir, \"checkpoints\", f\"{model_id}.pth\")\n                if os.path.exists(prev_checkpoint_path):\n                    checkpoint_path = prev_checkpoint_path\n\n        start = clock()\n        loaded_from_checkpoint = False\n        if os.path.exists(checkpoint_path) and not run_state.is_done(model_id):\n            # A checkpoint exists from an interrupted prior session but the model\n            # wasn't marked done -- reload its trained weights instead of retraining,\n            # then go straight to eval.\n            try:\n                state_dict = torch.load(checkpoint_path, map_location=device)\n                model.load_state_dict(state_dict)\n                loaded_from_checkpoint = True\n            except Exception as exc:\n                print(f\"WARNING: failed to load checkpoint for {model_id} from \"\n                      f\"{checkpoint_path} ({exc!r}); training from scratch instead.\")\n\n        if not loaded_from_checkpoint:\n            model = train_fn(model, model_id, model_budget, cfg=cfg, train_loader=train_loader, val_loader=val_loader, device=device)\n            own_checkpoint_path = os.path.join(CHECKPOINT_DIR, f\"{model_id}.pth\")\n            os.makedirs(CHECKPOINT_DIR, exist_ok=True)\n            torch.save(model.state_dict(), own_checkpoint_path)\n        ds_results = eval_fn(model, model_id, cfg=cfg, ood_loaders=ood_loaders, device=device)\n        elapsed = clock() - start\n        run_state.seconds_used_total += elapsed\n        run_state.mark_done(model_id, ds_results)\n        run_state.save(RUN_STATE_PATH)\n        _save_results_and_report(cfg, run_state, model_id, ds_results)\n        remaining.remove(model_id)\n\n\ndef _save_results_and_report(cfg, run_state, model_id, ds_results):\n    \"\"\"Persists the growing results file and prints this model's metrics\n    immediately, so results are available on disk and visible in the log the\n    moment each model's evaluation finishes rather than only at the very end\n    of the whole sweep.\"\"\"\n    results_dir = os.path.dirname(RESULTS_PATH)\n    if results_dir:\n        os.makedirs(results_dir, exist_ok=True)\n    with open(RESULTS_PATH, \"w\") as f:\n        json.dump(run_state.metrics, f, indent=2)\n\n    print(f\"\\n==== Completed: {model_id} ====\")\n    for ds_name, m in ds_results.items():\n        acc = m.get(\"acc_f\")\n        f1 = m.get(\"f1_f\")\n        acc_opt = m.get(\"acc_opt\")\n        roc = m.get(\"roc\")\n        acc_str = f\"{acc:.3f}\" if acc is not None else \"N/A\"\n        f1_str = f\"{f1:.3f}\" if f1 is not None else \"N/A\"\n        acc_opt_str = f\"{acc_opt:.3f}\" if acc_opt is not None else \"N/A\"\n        roc_str = f\"{roc:.3f}\" if roc is not None else \"N/A\"\n        print(f\"  {ds_name:10s}  acc@0.5={acc_str}  acc@opt={acc_opt_str}  f1@0.5={f1_str}  roc_auc={roc_str}\")\n    print(f\"Saved results so far to {RESULTS_PATH}\\n\")\n\n\ndef run(cfg, build_model_fn=None, train_fn=None, eval_fn=None, ood_loaders_fn=None, train_val_fn=None,\n        device=None, clock=None):\n    import time\n\n    build_model_fn = build_model_fn or _default_build_model_fn\n    train_fn = train_fn or _default_train_fn\n    eval_fn = eval_fn or _default_eval_fn\n    ood_loaders_fn = ood_loaders_fn or build_ood_loaders\n    train_val_fn = train_val_fn or (lambda c: make_fakeavceleb_train_val(\n        c.FAKEAVCELEB_FRAMES_ROOT, c.FAKEAVCELEB_AUDIO_ROOT,\n        max_train_samples=getattr(c, \"TRAIN_MAX_SAMPLES\", None),\n        max_val_samples=getattr(c, \"TRAIN_VAL_MAX_SAMPLES\", None),\n        cache_dir=getattr(c, \"CACHE_DIR\", None),\n    ))\n    device = device or (torch.device(\"cuda\") if torch.cuda.is_available() else torch.device(\"cpu\"))\n    clock = clock or time.monotonic\n\n    run_state = RunState.load(RUN_STATE_PATH, resume_input_dir=getattr(cfg, \"RESUME_INPUT_DIR\", None))\n    budget_mgr = BudgetManager(\n        run_state, WEEKLY_BUDGET_SECONDS, SESSION_BUDGET_SECONDS,\n        MIN_MODEL_BUDGET_SECONDS, MAX_MODEL_BUDGET_SECONDS, clock=clock,\n    )\n    budget_mgr.session_start()\n\n    ood_loaders = ood_loaders_fn(cfg)\n    train_ds, val_ds = train_val_fn(cfg)\n    train_loader = torch.utils.data.DataLoader(train_ds, batch_size=8, shuffle=True) if train_ds is not None else None\n    val_loader = torch.utils.data.DataLoader(val_ds, batch_size=8, shuffle=False) if val_ds is not None else None\n\n    base_ids = BASE_AUDIO_IDS + BASE_IMAGE_IDS\n    _run_phase(base_ids, cfg, run_state, budget_mgr, build_model_fn, train_fn, eval_fn,\n               ood_loaders, train_loader, val_loader, device, clock)\n\n    if all(run_state.is_done(m) for m in base_ids) and run_state.best_audio_id is None:\n        best_audio, best_image = select_best_pair(run_state)\n        run_state.set_best_pair(best_audio, best_image)\n        run_state.save(RUN_STATE_PATH)\n\n    if run_state.best_audio_id is not None:\n        _run_phase(FUSION_IDS, cfg, run_state, budget_mgr, build_model_fn, train_fn, eval_fn,\n                   ood_loaders, train_loader, val_loader, device, clock,\n                   extra_build_kwargs={\"audio_model_id\": run_state.best_audio_id, \"image_model_id\": run_state.best_image_id})\n\n    os.makedirs(os.path.dirname(RESULTS_PATH), exist_ok=True) if os.path.dirname(RESULTS_PATH) else None\n    with open(RESULTS_PATH, \"w\") as f:\n        json.dump(run_state.metrics, f, indent=2)\n\n    if all(run_state.is_done(m) for m in MODEL_IDS):\n        table = results_to_markdown_table(run_state.metrics, OOD_DATASETS)\n        print(table)\n\n\nif __name__ == \"__main__\":\n    import types as _types\n    _cfg = _types.SimpleNamespace(**{\n        k: v for k, v in globals().items()\n        if k.isupper() or k in (\"AUDIO_PRETRAINED\", \"IMAGE_BACKBONE\")\n    })\n    run(_cfg)","metadata":{"_uuid":"c60f5683-6643-455e-860c-f252309210ce","_cell_guid":"55d5568c-1a29-40da-82e3-8e299828ea0b","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}