{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":106680,"databundleVersionId":13374319,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nAIRR-ML-25 Challenge - FINAL GPU PRODUCTION (with fixed submission)\n==================================================================\n- GPU-based feature selection (XGBoost) to avoid CPU mutual-info crashes.\n- Correct submission generation:\n  * Start from sample_submissions.csv\n  * Update ONLY Task-1 rows for test_dataset_* with predicted probabilities\n  * Leave Task-2 rows (train_dataset_*) exactly as-is\n\"\"\"\n\nimport os\nimport gc\nimport warnings\nimport subprocess\nfrom pathlib import Path\nfrom collections import Counter\nfrom typing import Dict, Optional\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.linear_model import LogisticRegression\n\nimport xgboost as xgb\nimport lightgbm as lgb\nfrom joblib import Parallel, delayed\n\nwarnings.filterwarnings(\"ignore\")\n\n\n# =====================================================================\n# CONFIGURATION\n# =====================================================================\nclass Config:\n    DATA_ROOT = Path(\"/kaggle/input/adaptive-immune-profiling-challenge-2025/\")\n    TRAIN_DIR = DATA_ROOT / \"train_datasets\" / \"train_datasets\"\n    TEST_DIR = DATA_ROOT / \"test_datasets\" / \"test_datasets\"\n    SAMPLE_SUBMISSION = DATA_ROOT / \"sample_submissions.csv\"\n\n    # Keep strictly to [3, 4] to prevent RAM explosion\n    K_LIST = [3, 4]\n    TOP_KMER = 400\n    MAX_SEQUENCES_PER_FILE = 50000\n\n    # Public clone settings\n    PUB_MAX_FILES = 20\n    PUB_MIN_FREQ = 0.18\n    PUB_ENRICH = 6.0\n    PUB_TOP_N = {7: 5000, 8: 3000, \"default\": 2000}\n\n    N_SPLITS = 5\n    RANDOM_STATE = 42\n    EARLY_STOP = 100\n\n    SCALE_POS_WEIGHT = {\n        1: 1.0, 2: 1.0, 3: 1.0, 4: 1.0, 5: 1.0, 6: 1.0, 7: 5.04, 8: 2.05\n    }\n\n\n# =====================================================================\n# AMINO ACID PROPERTIES\n# =====================================================================\nAA_PROPERTIES = {\n    \"A\": {\"hydro\": 1.8, \"vol\": 88.6, \"charge\": 0, \"polar\": 0},\n    \"R\": {\"hydro\": -4.5, \"vol\": 173.4, \"charge\": 1, \"polar\": 1},\n    \"N\": {\"hydro\": -3.5, \"vol\": 114.1, \"charge\": 0, \"polar\": 1},\n    \"D\": {\"hydro\": -3.5, \"vol\": 111.1, \"charge\": -1, \"polar\": 1},\n    \"C\": {\"hydro\": 2.5, \"vol\": 108.5, \"charge\": 0, \"polar\": 0},\n    \"Q\": {\"hydro\": -3.5, \"vol\": 143.8, \"charge\": 0, \"polar\": 1},\n    \"E\": {\"hydro\": -3.5, \"vol\": 138.4, \"charge\": -1, \"polar\": 1},\n    \"G\": {\"hydro\": -0.4, \"vol\": 60.1, \"charge\": 0, \"polar\": 0},\n    \"H\": {\"hydro\": -3.2, \"vol\": 153.2, \"charge\": 0.5, \"polar\": 1},\n    \"I\": {\"hydro\": 4.5, \"vol\": 166.7, \"charge\": 0, \"polar\": 0},\n    \"L\": {\"hydro\": 3.8, \"vol\": 166.7, \"charge\": 0, \"polar\": 0},\n    \"K\": {\"hydro\": -3.9, \"vol\": 168.6, \"charge\": 1, \"polar\": 1},\n    \"M\": {\"hydro\": 1.9, \"vol\": 162.9, \"charge\": 0, \"polar\": 0},\n    \"F\": {\"hydro\": 2.8, \"vol\": 189.9, \"charge\": 0, \"polar\": 0},\n    \"P\": {\"hydro\": -1.6, \"vol\": 112.7, \"charge\": 0, \"polar\": 0},\n    \"S\": {\"hydro\": -0.8, \"vol\": 89.0, \"charge\": 0, \"polar\": 1},\n    \"T\": {\"hydro\": -0.7, \"vol\": 116.1, \"charge\": 0, \"polar\": 1},\n    \"W\": {\"hydro\": -0.9, \"vol\": 227.8, \"charge\": 0, \"polar\": 0},\n    \"Y\": {\"hydro\": -1.3, \"vol\": 193.6, \"charge\": 0, \"polar\": 1},\n    \"V\": {\"hydro\": 4.2, \"vol\": 140.0, \"charge\": 0, \"polar\": 0},\n}\n\n\n# =====================================================================\n# GPU UTILITIES\n# =====================================================================\ndef check_gpu() -> bool:\n    try:\n        r = subprocess.run([\"nvidia-smi\"], capture_output=True, text=True, timeout=5)\n        return r.returncode == 0\n    except Exception:\n        return False\n\n\ndef get_gpu_memory() -> str:\n    try:\n        result = subprocess.run(\n            [\"nvidia-smi\", \"--query-gpu=memory.used,memory.total\", \"--format=csv,noheader,nounits\"],\n            capture_output=True, text=True, timeout=5\n        )\n        if result.returncode == 0 and result.stdout.strip():\n            used, total = map(int, result.stdout.strip().split(\",\"))\n            return f\"{used}/{total} MB\"\n    except Exception:\n        pass\n    return \"N/A\"\n\n\n# =====================================================================\n# DATA UTILITIES\n# =====================================================================\ndef dataset_id_from_name(name: str) -> int:\n    for part in name.replace(\"_\", \" \").split():\n        if part.isdigit():\n            return int(part)\n    return 1\n\n\ndef read_repertoire(tsv_path: Path, max_seqs: Optional[int] = None) -> pd.DataFrame:\n    cols = [\"junction_aa\", \"v_call\", \"j_call\", \"templates\"]\n    try:\n        header = pd.read_csv(tsv_path, sep=\"\\t\", nrows=0)\n        usecols = [c for c in cols if c in header.columns]\n        df = pd.read_csv(tsv_path, sep=\"\\t\", usecols=usecols)\n    except Exception:\n        return pd.DataFrame(columns=cols)\n\n    if max_seqs and len(df) > max_seqs:\n        if \"templates\" in df.columns:\n            weights = pd.to_numeric(df[\"templates\"], errors=\"coerce\").fillna(1.0).values\n            s = weights.sum()\n            if s <= 0:\n                df = df.sample(n=max_seqs, random_state=42).reset_index(drop=True)\n            else:\n                weights = weights / s\n                idx = np.random.choice(len(df), max_seqs, replace=False, p=weights)\n                df = df.iloc[idx].reset_index(drop=True)\n        else:\n            df = df.sample(n=max_seqs, random_state=42).reset_index(drop=True)\n\n    for col in cols:\n        if col not in df.columns:\n            df[col] = \"\" if col != \"templates\" else 1.0\n\n    df[\"junction_aa\"] = df[\"junction_aa\"].fillna(\"\").astype(str)\n    df[\"templates\"] = pd.to_numeric(df[\"templates\"], errors=\"coerce\").fillna(1.0)\n    return df\n\n\n# =====================================================================\n# FEATURE ENGINEERING\n# =====================================================================\nclass FeatureExtractor:\n    def __init__(self, k_list=None):\n        self.k_list = k_list or [3, 4]\n\n    def gene_family(self, gene_call: str) -> str:\n        if not isinstance(gene_call, str) or not gene_call:\n            return \"UNK\"\n        return gene_call.split(\"*\")[0].split(\"-\")[0].upper() or \"UNK\"\n\n    def extract_all(\n        self,\n        df: pd.DataFrame,\n        pub_dict: Optional[Dict] = None,\n        meta_row: Optional[pd.Series] = None,\n        ds_id: int = 1\n    ) -> Dict[str, float]:\n        seqs = df[\"junction_aa\"].dropna().astype(str).tolist()\n        seqs = [s for s in seqs if len(s) > 0]\n\n        features: Dict[str, float] = {}\n\n        # 1) K-mers\n        for k in self.k_list:\n            c = Counter()\n            total = 0\n            for seq in seqs:\n                if len(seq) < k:\n                    continue\n                for i in range(len(seq) - k + 1):\n                    kmer = seq[i:i + k]\n                    if all(ch in AA_PROPERTIES for ch in kmer):\n                        c[kmer] += 1\n                        total += 1\n            if total > 0:\n                features.update({f\"kmer_{k}_{km}\": v / total for km, v in c.items()})\n\n        # 2) Positional k-mers (start/end)\n        k_pos = 3\n        start_c, end_c = Counter(), Counter()\n        ns, ne = 0, 0\n        for seq in seqs:\n            if len(seq) < k_pos:\n                continue\n            sk, ek = seq[:k_pos], seq[-k_pos:]\n            if all(ch in AA_PROPERTIES for ch in sk):\n                start_c[sk] += 1\n                ns += 1\n            if all(ch in AA_PROPERTIES for ch in ek):\n                end_c[ek] += 1\n                ne += 1\n        if ns > 0:\n            features.update({f\"pos_start_{km}\": v / ns for km, v in start_c.most_common(20)})\n        if ne > 0:\n            features.update({f\"pos_end_{km}\": v / ne for km, v in end_c.most_common(20)})\n\n        # 3) Physicochemical\n        hydro, vol = [], []\n        for seq in seqs:\n            h, v = 0.0, 0.0\n            cnt = 0\n            for aa in seq:\n                if aa in AA_PROPERTIES:\n                    h += AA_PROPERTIES[aa][\"hydro\"]\n                    v += AA_PROPERTIES[aa][\"vol\"]\n                    cnt += 1\n            if cnt > 0:\n                hydro.append(h / cnt)\n                vol.append(v / cnt)\n        if hydro:\n            features[\"phys_hydro_mean\"] = float(np.mean(hydro))\n            features[\"phys_vol_mean\"] = float(np.mean(vol))\n\n        # 4) V families\n        if \"v_call\" in df.columns:\n            v_fam = df[\"v_call\"].apply(self.gene_family)\n            for fam, freq in v_fam.value_counts(normalize=True).head(30).items():\n                features[f\"v_fam_{fam}\"] = float(freq)\n\n        # 5) Length stats\n        lens = [len(s) for s in seqs]\n        if lens:\n            features[\"len_mean\"] = float(np.mean(lens))\n            features[\"len_std\"] = float(np.std(lens))\n\n        # 6) Metadata\n        if meta_row is not None:\n            if \"sex\" in meta_row.index:\n                features[\"meta_sex_male\"] = 1.0 if str(meta_row[\"sex\"]).upper() in [\"M\", \"MALE\"] else 0.0\n            if ds_id == 7 and \"race\" in meta_row.index:\n                features[\"meta_race_white\"] = 1.0 if \"white\" in str(meta_row[\"race\"]).lower() else 0.0\n            if ds_id == 7 and \"sequencing_run_id\" in meta_row.index:\n                features[\"meta_run_hash\"] = (hash(str(meta_row[\"sequencing_run_id\"])) % 100) / 100.0\n            if ds_id == 8:\n                for hla in [\"A\", \"B\", \"C\", \"DRB1\"]:\n                    if hla in meta_row.index:\n                        features[f\"meta_hla_{hla}\"] = 1.0 if pd.notna(meta_row[hla]) else 0.0\n\n        # 7) Public clones\n        if pub_dict:\n            seq_set = set(seqs)\n            hits = [pub_dict[s][\"score\"] for s in seq_set if s in pub_dict]\n            features[\"pub_score_sum\"] = float(sum(hits))\n            features[\"pub_hits\"] = float(len(hits))\n\n        return features\n\n\n# =====================================================================\n# PUBLIC CLONE MINING\n# =====================================================================\ndef mine_public_clones(\n    dataset_path: Path,\n    max_files: int = 20,\n    min_freq: float = 0.18,\n    enrichment: float = 6.0,\n    top_n: int = 2000\n) -> Dict[str, Dict]:\n    meta = pd.read_csv(dataset_path / \"metadata.csv\")\n    pos_files = meta[meta[\"label_positive\"] == True][\"filename\"].tolist()[:max_files]  # noqa: E712\n    neg_files = meta[meta[\"label_positive\"] == False][\"filename\"].tolist()[:max_files]  # noqa: E712\n\n    if not pos_files:\n        return {}\n\n    def get_seqs(files):\n        c = Counter()\n        for f in files:\n            try:\n                df = pd.read_csv(dataset_path / f, sep=\"\\t\", usecols=[\"junction_aa\"])\n                c.update(df[\"junction_aa\"].dropna().unique())\n            except Exception:\n                pass\n        return c\n\n    pos_c = get_seqs(pos_files)\n    neg_c = get_seqs(neg_files)\n\n    scored = []\n    n_pos, n_neg = max(1, len(pos_files)), max(1, len(neg_files))\n\n    for seq, count in pos_c.items():\n        pf = count / n_pos\n        nf = neg_c.get(seq, 0) / n_neg\n        if pf >= min_freq and pf > nf * enrichment:\n            score = float(np.log((pf + 1e-6) / (nf + 1e-6)))\n            scored.append({\"seq\": seq, \"score\": score})\n\n    scored.sort(key=lambda x: -x[\"score\"])\n    return {item[\"seq\"]: item for item in scored[:top_n]}\n\n\n# =====================================================================\n# ENSEMBLE TRAINER (GPU FIXED)\n# =====================================================================\nclass EnsembleTrainer:\n    def __init__(self, use_gpu: bool = True, random_state: int = 42):\n        self.use_gpu = use_gpu\n        self.random_state = random_state\n        self.models = {}\n        self.weights = {\"xgb\": 0.5, \"lgb\": 0.5}\n        self.feature_cols = []\n\n    def select_features_gpu(self, X_df: pd.DataFrame, y: np.ndarray, top_k: int = 400):\n        print(f\"   Selecting top {top_k} features using GPU XGBoost... \", end=\"\")\n\n        all_cols = X_df.columns.tolist()\n        dtrain = xgb.DMatrix(X_df, label=y)\n\n        params = {\n            \"tree_method\": \"hist\",\n            \"device\": \"cuda\",\n            \"max_depth\": 4,\n            \"learning_rate\": 0.1,\n            \"reg_lambda\": 1.0,\n            \"verbosity\": 0,\n        }\n\n        bst = xgb.train(params, dtrain, num_boost_round=20)\n        scores = bst.get_score(importance_type=\"gain\")  # {feature_name: gain}\n\n        sorted_feats = sorted(scores.items(), key=lambda x: x[1], reverse=True)\n        selected = [f[0] for f in sorted_feats[:top_k]]\n\n        if len(selected) < top_k:\n            remaining = [c for c in all_cols if c not in selected]\n            selected.extend(remaining[:top_k - len(selected)])\n\n        print(f\"Done ({len(selected)}).\")\n        return selected\n\n    def train(self, df: pd.DataFrame, ds_id: int):\n        y = df[\"label_positive\"].values.astype(np.float32)\n        X_df = df.drop(columns=[\"ID\", \"dataset\", \"label_positive\"], errors=\"ignore\")\n\n        self.feature_cols = self.select_features_gpu(X_df, y, top_k=Config.TOP_KMER)\n        X = X_df[self.feature_cols].values.astype(np.float32)\n\n        print(f\"  Training ensemble on {len(X)} samples, {len(self.feature_cols)} cols\")\n\n        xgb_params = {\n            \"objective\": \"binary:logistic\",\n            \"eval_metric\": \"auc\",\n            \"max_depth\": 6,\n            \"learning_rate\": 0.03,\n            \"subsample\": 0.8,\n            \"colsample_bytree\": 0.8,\n            \"min_child_weight\": 15,\n            \"seed\": self.random_state,\n            \"scale_pos_weight\": Config.SCALE_POS_WEIGHT.get(ds_id, 1.0),\n            \"tree_method\": \"hist\",\n            \"device\": \"cuda\",\n            \"verbosity\": 0,\n        }\n\n        lgb_params = {\n            \"objective\": \"binary\",\n            \"metric\": \"auc\",\n            \"device\": \"gpu\",\n            \"max_depth\": 6,\n            \"learning_rate\": 0.02,\n            \"num_leaves\": 31,\n            \"min_child_samples\": 20,\n            \"scale_pos_weight\": Config.SCALE_POS_WEIGHT.get(ds_id, 1.0),\n            \"verbosity\": -1,\n        }\n\n        # robust n_splits\n        pos = int((y == 1).sum())\n        neg = int((y == 0).sum())\n        min_class = max(2, min(pos, neg)) if (pos > 0 and neg > 0) else 2\n        n_splits = min(Config.N_SPLITS, min_class)\n\n        kf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=self.random_state)\n\n        oof_xgb = np.zeros(len(y), dtype=np.float32)\n        oof_lgb = np.zeros(len(y), dtype=np.float32)\n\n        cv_xgb, cv_lgb, best_iters = [], [], []\n\n        for tr_idx, va_idx in kf.split(X, y):\n            X_tr, X_val = X[tr_idx], X[va_idx]\n            y_tr, y_val = y[tr_idx], y[va_idx]\n\n            # XGB\n            dtr = xgb.DMatrix(X_tr, label=y_tr)\n            dval = xgb.DMatrix(X_val, label=y_val)\n            bst = xgb.train(\n                xgb_params,\n                dtr,\n                num_boost_round=1000,\n                evals=[(dval, \"v\")],\n                early_stopping_rounds=Config.EARLY_STOP,\n                verbose_eval=False,\n            )\n            oof_xgb[va_idx] = bst.predict(dval)\n            cv_xgb.append(roc_auc_score(y_val, oof_xgb[va_idx]))\n            best_iters.append(int(bst.best_iteration or 0))\n\n            # LGB\n            lgb_tr = lgb.Dataset(X_tr, label=y_tr)\n            lgb_val = lgb.Dataset(X_val, label=y_val, reference=lgb_tr)\n            lgb_bst = lgb.train(\n                lgb_params,\n                lgb_tr,\n                num_boost_round=1000,\n                valid_sets=[lgb_val],\n                callbacks=[lgb.early_stopping(Config.EARLY_STOP, verbose=False)],\n            )\n            oof_lgb[va_idx] = lgb_bst.predict(X_val)\n            cv_lgb.append(roc_auc_score(y_val, oof_lgb[va_idx]))\n\n        print(f\"  CV: XGB={float(np.mean(cv_xgb)):.4f}, LGB={float(np.mean(cv_lgb)):.4f}\")\n\n        # Stacking weights (safe)\n        meta = LogisticRegression(max_iter=2000)\n        meta.fit(np.column_stack([oof_xgb, oof_lgb]), y)\n        w = np.clip(meta.coef_[0], 0, None)\n        if float(w.sum()) <= 0:\n            self.weights = {\"xgb\": 0.5, \"lgb\": 0.5}\n        else:\n            self.weights = {\"xgb\": float(w[0] / w.sum()), \"lgb\": float(w[1] / w.sum())}\n\n        rounds = int(np.mean(best_iters)) + 50\n        rounds = max(rounds, 100)\n\n        self.models[\"xgb\"] = xgb.train(xgb_params, xgb.DMatrix(X, label=y), num_boost_round=rounds)\n        self.models[\"lgb\"] = lgb.train(lgb_params, lgb.Dataset(X, label=y), num_boost_round=800)\n\n        return self, self.feature_cols, float(np.mean(cv_xgb))\n\n    def predict(self, X: np.ndarray) -> np.ndarray:\n        X = X.astype(np.float32)\n        p1 = self.models[\"xgb\"].predict(xgb.DMatrix(X))\n        p2 = self.models[\"lgb\"].predict(X)\n        return p1 * self.weights[\"xgb\"] + p2 * self.weights[\"lgb\"]\n\n\n# =====================================================================\n# PARALLEL HELPERS\n# =====================================================================\ndef process_file_parallel(row, path: Path, ds_id: int, pub_dict: Dict, extractor: FeatureExtractor):\n    try:\n        df = read_repertoire(path / row[\"filename\"], Config.MAX_SEQUENCES_PER_FILE)\n        feats = extractor.extract_all(df, pub_dict, row, ds_id)\n        return {\n            **feats,\n            \"ID\": row.get(\"repertoire_id\", Path(row[\"filename\"]).stem),\n            \"label_positive\": int(row[\"label_positive\"]),\n            \"dataset\": path.name,\n        }\n    except Exception:\n        return None\n\n\ndef process_test_parallel(tsv: Path, path: Path, ds_id: int, pub_dict: Dict, extractor: FeatureExtractor):\n    try:\n        df = read_repertoire(tsv, Config.MAX_SEQUENCES_PER_FILE)\n        feats = extractor.extract_all(df, pub_dict, None, ds_id)\n        return {**feats, \"ID\": tsv.stem, \"dataset\": path.name}\n    except Exception:\n        return None\n\n\n# =====================================================================\n# SUBMISSION CREATION (FIXED)\n# =====================================================================\ndef create_submission_from_final(\n    final_pred_df: pd.DataFrame,\n    sample_path: str | Path,\n    output_path: str = \"submission.csv\",\n) -> pd.DataFrame:\n    \"\"\"\n    final_pred_df must contain: ['ID','dataset','label_positive_probability'] for test_dataset_* rows.\n    This updates ONLY Task-1 rows (test_dataset_*) in the sample submission and leaves Task-2 intact.\n    \"\"\"\n    sample = pd.read_csv(sample_path)\n\n    pred_map = (\n        final_pred_df.drop_duplicates(subset=[\"dataset\", \"ID\"])\n        .set_index([\"dataset\", \"ID\"])[\"label_positive_probability\"]\n    )\n\n    test_mask = sample[\"dataset\"].astype(str).str.startswith(\"test_dataset_\")\n    idx = pd.MultiIndex.from_frame(sample.loc[test_mask, [\"dataset\", \"ID\"]])\n\n    new_vals = pred_map.reindex(idx).to_numpy()\n    old_vals = sample.loc[test_mask, \"label_positive_probability\"].to_numpy()\n\n    sample.loc[test_mask, \"label_positive_probability\"] = np.where(pd.isna(new_vals), old_vals, new_vals)\n    sample.to_csv(output_path, index=False)\n    return sample\n\n\n# =====================================================================\n# MAIN PIPELINE\n# =====================================================================\ndef main():\n    print(\"AIRR-ML-25: GPU OPTIMIZED PIPELINE (submission fixed)\")\n    gpu_ok = check_gpu()\n    print(f\"GPU detected: {gpu_ok} | GPU mem: {get_gpu_memory()}\")\n\n    train_sets = sorted([d.name for d in Config.TRAIN_DIR.glob(\"train_dataset_*\")])\n    test_sets = sorted([d.name for d in Config.TEST_DIR.glob(\"test_dataset_*\")])\n\n    extractor = FeatureExtractor(Config.K_LIST)\n    bundles = {}\n\n    for ds_name in train_sets:\n        ds_id = dataset_id_from_name(ds_name)\n        print(f\"\\nTraining on {ds_name} (id={ds_id})\")\n\n        pub_dict = mine_public_clones(\n            Config.TRAIN_DIR / ds_name,\n            max_files=Config.PUB_MAX_FILES,\n            min_freq=Config.PUB_MIN_FREQ,\n            enrichment=Config.PUB_ENRICH,\n            top_n=Config.PUB_TOP_N.get(ds_id, Config.PUB_TOP_N[\"default\"]),\n        )\n\n        meta = pd.read_csv(Config.TRAIN_DIR / ds_name / \"metadata.csv\")\n        print(\"  Extracting features...\")\n        res = Parallel(n_jobs=12, backend=\"loky\")(\n            delayed(process_file_parallel)(row, Config.TRAIN_DIR / ds_name, ds_id, pub_dict, extractor)\n            for _, row in tqdm(meta.iterrows(), total=len(meta), leave=False)\n        )\n        df = pd.DataFrame([r for r in res if r is not None])\n\n        trainer = EnsembleTrainer(use_gpu=True, random_state=Config.RANDOM_STATE)\n        trainer, fcols, score = trainer.train(df, ds_id)\n\n        bundles[ds_name] = {\"trainer\": trainer, \"cols\": fcols, \"pub\": pub_dict, \"score\": score}\n        print(f\"  Stored bundle for {ds_name} | AUC proxy: {score:.4f}\")\n\n        del df, res\n        gc.collect()\n\n    print(\"\\nPREDICTING...\")\n    preds = []\n    default_train = train_sets[0] if train_sets else None\n\n    for test_name in test_sets:\n        ds_id = dataset_id_from_name(test_name)\n        train_key = f\"train_dataset_{ds_id}\"\n        if train_key not in bundles:\n            train_key = default_train\n\n        if train_key is None:\n            raise RuntimeError(\"No training bundles available.\")\n\n        bundle = bundles[train_key]\n        print(f\"  {test_name} -> {train_key}\")\n\n        files = sorted((Config.TEST_DIR / test_name).glob(\"*.tsv\"))\n        res = Parallel(n_jobs=12, backend=\"loky\")(\n            delayed(process_test_parallel)(f, Config.TEST_DIR / test_name, ds_id, bundle[\"pub\"], extractor)\n            for f in tqdm(files, leave=False)\n        )\n        test_df = pd.DataFrame([r for r in res if r is not None])\n\n        # Align columns to training features\n        X = pd.DataFrame(0.0, index=np.arange(len(test_df)), columns=bundle[\"cols\"])\n        for c in bundle[\"cols\"]:\n            if c in test_df.columns:\n                X[c] = test_df[c].astype(np.float32)\n\n        p = bundle[\"trainer\"].predict(X.values)\n\n        sub_part = test_df[[\"ID\", \"dataset\"]].copy()\n        sub_part[\"label_positive_probability\"] = p.astype(float)\n        preds.append(sub_part)\n\n        del test_df, X, res\n        gc.collect()\n\n    # ---- Submission (FIXED) ----\n    final = pd.concat(preds, ignore_index=True) if preds else pd.DataFrame(\n        columns=[\"ID\", \"dataset\", \"label_positive_probability\"]\n    )\n    create_submission_from_final(final, Config.SAMPLE_SUBMISSION, \"best_submission.csv\")\n    print(\"\\nDone! Saved submission.csv\")\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T15:12:00.129942Z","iopub.execute_input":"2025-12-13T15:12:00.130260Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}