{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.0"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"datasetVersion","sourceId":5152020,"datasetId":2993489,"databundleVersionId":5223746},{"sourceType":"datasetVersion","sourceId":4059877,"datasetId":2395943,"databundleVersionId":4115935},{"sourceType":"datasetVersion","sourceId":13779956,"datasetId":8771108,"databundleVersionId":14533918},{"sourceType":"datasetVersion","sourceId":3319673,"datasetId":1870444,"databundleVersionId":3370595},{"sourceType":"datasetVersion","sourceId":897617,"datasetId":480187,"databundleVersionId":924581}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":5,"nbformat":4,"cells":[{"id":"title_cell","cell_type":"markdown","source":"# TCH-Net: Multi-Branch IoT Botnet Detection\n\n**T**emporal + **C**ontext + **H**igh-level statistical branches with Cross-Branch Gated Attention Fusion (CB-GAF).\n\nThis notebook trains TCH-Net across 5 seeds on the BRIDGE benchmark and saves best checkpoints.\n\n**Architecture highlights:**\n- **CB-GAF** — Cross-Branch Gated Attention Fusion\n- **MSTE** — Multi-Scale Temporal Encoding (Conv-GRU + Stride-GRU + Transformer)\n- **Auxiliary Feature Reconstruction** regularisation\n\n**Paper:** BRIDGE and TCH-Net — arXiv:2604.11324","metadata":{}},{"id":"cell1_md","cell_type":"markdown","source":"## Cell 1 — Imports","metadata":{}},{"id":"cell1_code","cell_type":"code","source":"import os, glob, json, time, warnings, copy, gc, math\nwarnings.filterwarnings('ignore')\nos.environ['CUDA_LAUNCH_BLOCKING'] = '1'\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib\ntry:\n get_ipython()\n import matplotlib.pyplot as plt\nexcept NameError:\n matplotlib.use('Agg')\n import matplotlib.pyplot as plt\n\nfrom sklearn.preprocessing import RobustScaler\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n classification_report, roc_auc_score,\n f1_score, accuracy_score, precision_score, recall_score,\n roc_curve, precision_recall_curve, auc, matthews_corrcoef)\n\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\n\nSEED = 42\nnp.random.seed(SEED); torch.manual_seed(SEED)\nif torch.cuda.is_available():\n torch.cuda.manual_seed_all(SEED)\n torch.backends.cudnn.benchmark = True\n torch.backends.cudnn.deterministic = False\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"PyTorch {torch.__version__} | Device: {device}\")\nprint(f\"CUDA: {torch.cuda.is_available()}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell2_md","cell_type":"markdown","source":"## Cell 2 — Canonical Feature Vocabulary (46 CICFlowMeter features)\n\nEvery feature keeps its name and units regardless of source dataset.\nFeatures not present in a dataset are zero-filled — never fabricated.\n\n- **T-branch** (indices 0–16): rates, durations, counts\n- **H-branch** (indices 17–37): packet size & IAT distributions\n- **Both** (indices 38–45): TCP flags, header, window","metadata":{}},{"id":"cell2_code","cell_type":"code","source":"SEMANTIC_FEATURES = [\n # ── Flow-level counts & rates (T-branch primary) ──────────────────────\n \"flow_duration\", # 0\n \"pkt_count_fwd\", # 1\n \"pkt_count_bwd\", # 2\n \"byte_count_fwd\", # 3\n \"byte_count_bwd\", # 4\n \"pkt_rate\", # 5\n \"byte_rate\", # 6\n \"fwd_pkt_rate\", # 7\n \"bwd_pkt_rate\", # 8\n \"fwd_byte_rate\", # 9\n \"bwd_byte_rate\", # 10\n \"pkt_count_total\", # 11\n \"byte_count_total\", # 12\n \"fwd_pkt_len_total\", # 13\n \"bwd_pkt_len_total\", # 14\n \"subflow_fwd_pkts\", # 15\n \"subflow_bwd_pkts\", # 16\n # ── Packet size statistics (H-branch primary) ─────────────────────────\n \"pkt_len_min\", # 17\n \"pkt_len_max\", # 18\n \"pkt_len_mean\", # 19\n \"pkt_len_std\", # 20\n \"pkt_len_var\", # 21\n \"fwd_pkt_len_min\", # 22\n \"fwd_pkt_len_max\", # 23\n \"fwd_pkt_len_mean\", # 24\n \"fwd_pkt_len_std\", # 25\n \"bwd_pkt_len_min\", # 26\n \"bwd_pkt_len_max\", # 27\n \"bwd_pkt_len_mean\", # 28\n \"bwd_pkt_len_std\", # 29\n # ── IAT statistics (T + H) ────────────────────────────────────────────\n \"iat_mean\", # 30\n \"iat_std\", # 31\n \"iat_max\", # 32\n \"iat_min\", # 33\n \"fwd_iat_mean\", # 34\n \"fwd_iat_std\", # 35\n \"bwd_iat_mean\", # 36\n \"bwd_iat_std\", # 37\n # ── TCP flags (H-branch) ──────────────────────────────────────────────\n \"flag_syn\", # 38\n \"flag_ack\", # 39\n \"flag_fin\", # 40\n \"flag_rst\", # 41\n \"flag_psh\", # 42\n \"flag_urg\", # 43\n # ── Header / window ───────────────────────────────────────────────────\n \"fwd_header_len\", # 44\n \"init_win_fwd\", # 45\n]\nN_SEM = len(SEMANTIC_FEATURES)\nSEM_IDX = {name: i for i, name in enumerate(SEMANTIC_FEATURES)}\n\n# ── Per-dataset alias maps ─────────────────────────────────────────────────\nALIAS_CICIDS = {\n \"flow_duration\": [\"Flow Duration\"],\n \"pkt_count_fwd\": [\"Total Fwd Packets\"],\n \"pkt_count_bwd\": [\"Total Backward Packets\"],\n \"byte_count_fwd\": [\"Total Length of Fwd Packets\"],\n \"byte_count_bwd\": [\"Total Length of Bwd Packets\"],\n \"pkt_rate\": [\"Flow Packets/s\"],\n \"byte_rate\": [\"Flow Bytes/s\"],\n \"fwd_pkt_rate\": [\"Fwd Packets/s\"],\n \"bwd_pkt_rate\": [\"Bwd Packets/s\"],\n \"fwd_pkt_len_total\": [\"Total Length of Fwd Packets\"],\n \"bwd_pkt_len_total\": [\"Total Length of Bwd Packets\"],\n \"subflow_fwd_pkts\": [\"Subflow Fwd Packets\"],\n \"subflow_bwd_pkts\": [\"Subflow Bwd Packets\"],\n \"pkt_len_min\": [\"Min Packet Length\"],\n \"pkt_len_max\": [\"Max Packet Length\"],\n \"pkt_len_mean\": [\"Packet Length Mean\"],\n \"pkt_len_std\": [\"Packet Length Std\"],\n \"pkt_len_var\": [\"Packet Length Variance\"],\n \"fwd_pkt_len_min\": [\"Fwd Packet Length Min\"],\n \"fwd_pkt_len_max\": [\"Fwd Packet Length Max\"],\n \"fwd_pkt_len_mean\": [\"Fwd Packet Length Mean\"],\n \"fwd_pkt_len_std\": [\"Fwd Packet Length Std\"],\n \"bwd_pkt_len_min\": [\"Bwd Packet Length Min\"],\n \"bwd_pkt_len_max\": [\"Bwd Packet Length Max\"],\n \"bwd_pkt_len_mean\": [\"Bwd Packet Length Mean\"],\n \"bwd_pkt_len_std\": [\"Bwd Packet Length Std\"],\n \"iat_mean\": [\"Flow IAT Mean\"],\n \"iat_std\": [\"Flow IAT Std\"],\n \"iat_max\": [\"Flow IAT Max\"],\n \"iat_min\": [\"Flow IAT Min\"],\n \"fwd_iat_mean\": [\"Fwd IAT Mean\"],\n \"fwd_iat_std\": [\"Fwd IAT Std\"],\n \"bwd_iat_mean\": [\"Bwd IAT Mean\"],\n \"bwd_iat_std\": [\"Bwd IAT Std\"],\n \"flag_syn\": [\"SYN Flag Count\"],\n \"flag_ack\": [\"ACK Flag Count\"],\n \"flag_fin\": [\"FIN Flag Count\"],\n \"flag_rst\": [\"RST Flag Count\"],\n \"flag_psh\": [\"PSH Flag Count\"],\n \"flag_urg\": [\"URG Flag Count\"],\n \"fwd_header_len\": [\"Fwd Header Length\"],\n \"init_win_fwd\": [\"Init_Win_bytes_forward\"],\n}\n\nALIAS_CICIOT = {\n \"flow_duration\": [\"flow_duration\", \"Flow Duration\", \"duration\"],\n \"pkt_count_total\": [\"Number\", \"number\", \"total_fwd_packets\", \"Total Fwd Packets\"],\n \"pkt_count_fwd\": [\"total_fwd_packets\", \"Total Fwd Packets\"],\n \"pkt_count_bwd\": [\"total_bwd_packets\", \"Total Backward Packets\"],\n \"byte_count_fwd\": [\"total_length_of_fwd_packets\", \"Total Length of Fwd Packets\"],\n \"byte_count_bwd\": [\"total_length_of_bwd_packets\", \"Total Length of Bwd Packets\"],\n \"byte_count_total\": [\"Tot sum\", \"tot sum\", \"Tot size\", \"tot size\"],\n \"pkt_rate\": [\"Rate\", \"rate\", \"flow_packets/s\", \"Flow Packets/s\"],\n \"byte_rate\": [\"flow_bytes/s\", \"Flow Bytes/s\"],\n \"fwd_pkt_rate\": [\"Srate\", \"srate\", \"fwd_packets/s\", \"Fwd Packets/s\"],\n \"bwd_pkt_rate\": [\"Drate\", \"drate\", \"bwd_packets/s\", \"Bwd Packets/s\"],\n \"pkt_len_min\": [\"Min\", \"min\", \"min_packet_length\"],\n \"pkt_len_max\": [\"Max\", \"max\", \"max_packet_length\"],\n \"pkt_len_mean\": [\"AVG\", \"avg\", \"packet_length_mean\"],\n \"pkt_len_std\": [\"Std\", \"std\", \"packet_length_std\"],\n \"pkt_len_var\": [\"Variance\", \"variance\", \"packet_length_variance\"],\n \"fwd_pkt_len_mean\": [\"fwd_packet_length_mean\"],\n \"fwd_pkt_len_std\": [\"fwd_packet_length_std\"],\n \"bwd_pkt_len_mean\": [\"bwd_packet_length_mean\"],\n \"iat_mean\": [\"IAT\", \"iat\", \"flow_iat_mean\", \"Flow IAT Mean\"],\n \"iat_std\": [\"flow_iat_std\"],\n \"iat_max\": [\"flow_iat_max\"],\n \"iat_min\": [\"flow_iat_min\"],\n \"fwd_iat_mean\": [\"fwd_iat_mean\"],\n \"fwd_iat_std\": [\"fwd_iat_std\"],\n \"bwd_iat_mean\": [\"bwd_iat_mean\"],\n \"bwd_iat_std\": [\"bwd_iat_std\"],\n \"flag_syn\": [\"syn_flag_number\", \"syn_flag_count\", \"SYN Flag Count\"],\n \"flag_ack\": [\"ack_flag_number\", \"ack_flag_count\", \"ACK Flag Count\"],\n \"flag_fin\": [\"fin_flag_number\", \"fin_flag_count\", \"FIN Flag Count\"],\n \"flag_rst\": [\"rst_flag_number\", \"rst_flag_count\", \"RST Flag Count\"],\n \"flag_psh\": [\"psh_flag_number\", \"psh_flag_count\", \"PSH Flag Count\"],\n \"flag_urg\": [\"urg_flag_number\", \"urg_flag_count\", \"URG Flag Count\"],\n \"fwd_header_len\": [\"Header_Length\", \"header_length\", \"Fwd Header Length\"],\n \"init_win_fwd\": [\"init_win_bytes_forward\", \"Init_Win_bytes_forward\"],\n}\n\nALIAS_BOTIOT = {\n \"flow_duration\": [\"dur\", \"duration\"],\n \"pkt_count_fwd\": [\"spkts\", \"src_pkts\"],\n \"pkt_count_bwd\": [\"dpkts\", \"dst_pkts\"],\n \"byte_count_fwd\": [\"sbytes\", \"src_bytes\"],\n \"byte_count_bwd\": [\"dbytes\", \"dst_bytes\"],\n \"pkt_rate\": [\"rate\"],\n \"fwd_pkt_rate\": [\"srate\", \"src_rate\"],\n \"bwd_pkt_rate\": [\"drate\", \"dst_rate\"],\n \"pkt_count_total\": [\"pkts\", \"totpkts\", \"total_pkts\"],\n \"byte_count_total\": [\"bytes\", \"totbytes\", \"total_bytes\"],\n \"iat_mean\": [\"sintpkt\"],\n \"bwd_iat_mean\": [\"dintpkt\"],\n \"flag_syn\": [\"syn\"],\n \"flag_ack\": [\"ack\"],\n \"flag_fin\": [\"fin\"],\n \"flag_rst\": [\"rst\"],\n \"flag_psh\": [\"push\"],\n \"init_win_fwd\": [\"swin\"],\n}\n\nALIAS_EDGE = {\n \"pkt_len_mean\": [\"tcp.len\"],\n \"byte_count_total\": [\"tcp.len\"],\n \"flag_syn\": [\"tcp.flags.syn\", \"tcp.connection.syn\"],\n \"flag_ack\": [\"tcp.flags.ack\"],\n \"flag_fin\": [\"tcp.connection.fin\"],\n \"flag_rst\": [\"tcp.connection.rst\"],\n \"flag_psh\": [\"tcp.flags.push\"],\n \"iat_mean\": [\"udp.time_delta\", \"tcp.time_delta\"],\n \"fwd_header_len\": [\"tcp.hdr_len\"],\n}\n\nALIAS_NBAIOT = {\n \"pkt_count_total\": [\"MI_dir_L5_weight\"],\n \"fwd_pkt_rate\": [\"MI_dir_L0.1_weight\"],\n \"pkt_len_mean\": [\"H_L5_mean\"],\n \"pkt_len_std\": [\"H_L5_std\"],\n \"pkt_len_var\": [\"H_L5_variance\"],\n \"fwd_iat_mean\": [\"HpHp_L5_mean\"],\n \"fwd_iat_std\": [\"HpHp_L5_std\"],\n}\n\nALIAS_MAPS = {\n \"CICIDS-2017\": ALIAS_CICIDS,\n \"CIC-IoT-2023\": ALIAS_CICIOT,\n \"Bot-IoT\": ALIAS_BOTIOT,\n \"Edge-IIoTset\": ALIAS_EDGE,\n \"N-BaIoT\": ALIAS_NBAIOT,\n}\n\ndef build_semantic_vector(df, ds_name):\n alias_map = ALIAS_MAPS.get(ds_name, {})\n col_lower = {c.strip().lower(): c for c in df.columns}\n out = np.zeros((len(df), N_SEM), dtype=np.float32)\n matched = []\n for sem_name in SEMANTIC_FEATURES:\n si = SEM_IDX[sem_name]\n col = None\n if sem_name in col_lower:\n col = col_lower[sem_name]\n elif sem_name in alias_map:\n for alias in alias_map[sem_name]:\n al = alias.strip().lower()\n if al in col_lower:\n col = col_lower[al]; break\n if col is None and sem_name in alias_map:\n for alias in alias_map[sem_name]:\n al = alias.strip().lower()\n ms = [c for c in df.columns if al in c.strip().lower()]\n if ms: col = ms[0]; break\n if col is None and sem_name in alias_map:\n for alias in alias_map[sem_name]:\n al = alias.strip().lower().replace(' ','').replace('_','')\n for c in df.columns:\n cn = c.strip().lower().replace(' ','').replace('_','')\n if al == cn or (len(al) > 4 and al in cn):\n col = c; break\n if col: break\n if col is not None:\n vals = pd.to_numeric(df[col], errors='coerce').fillna(0).values\n out[:, si] = vals.astype(np.float32)\n matched.append(sem_name)\n return out, matched\n\nprint(f\"Canonical vocabulary: {N_SEM} features\")\nprint(f\"Alias maps defined for {len(ALIAS_MAPS)} datasets\")\nfor ds, am in ALIAS_MAPS.items():\n print(f\" {ds}: {len(am)} mapped features / {N_SEM} = {len(am)/N_SEM*100:.0f}%\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell3_md","cell_type":"markdown","source":"## Cell 3 — Configuration","metadata":{}},{"id":"cell3_code","cell_type":"code","source":"class Config:\n # ── Dataset paths (update to your paths) ──────────────────────────────\n CICIDS_PATH = \"/kaggle/input/datasets/dhoogla/cicids2017\"\n CICIOT_PATH = \"/kaggle/input/datasets/raqeeb24/ciciot-2023-stratified-dataset\"\n BOTIOT_PATH = \"/kaggle/input/datasets/vigneshvenkateswaran/bot-iot-5-data\"\n EDGE_PATH = \"/kaggle/input/datasets/mohamedamineferrag/edgeiiotset-cyber-security-dataset-of-iot-iiot\"\n NBAIOT_PATH = \"/kaggle/input/datasets/mkashifn/nbaiot-dataset\"\n\n USE_CICIDS = os.path.isdir(CICIDS_PATH) if CICIDS_PATH else False\n USE_CICIOT = os.path.isdir(CICIOT_PATH) if CICIOT_PATH else False\n USE_BOTIOT = os.path.isdir(BOTIOT_PATH) if BOTIOT_PATH else False\n USE_EDGE = os.path.isdir(EDGE_PATH) if EDGE_PATH else False\n USE_NBAIOT = os.path.isdir(NBAIOT_PATH) if NBAIOT_PATH else False\n\n # ── Row caps per dataset ───────────────────────────────────────────────\n CICIDS_MAX = 3_000_000\n CICIOT_MAX = 3_000_000\n BOTIOT_MAX = 3_000_000\n EDGE_MAX = 2_000_000\n NBAIOT_MAX = 3_000_000\n\n # ── Sequence parameters ────────────────────────────────────────────────\n WINDOW_SIZE = 32\n STRIDE = 4\n MAX_TRAIN_SEQ = 800_000\n MAX_TEST_SEQ = 200_000\n\n # ── Class balance ──────────────────────────────────────────────────────\n TARGET_ATK_BEN_RATIO = 1.0\n\n # ── Training ───────────────────────────────────────────────────────────\n EPOCHS = 30\n WARMUP = 3\n EARLY_STOP = 7\n BATCH_SIZE = 512\n LR = 5e-4\n WD = 1e-4\n FOCAL_GAMMA = 2.5\n LABEL_SMOOTH = 0.01\n AUX_WT = 0.05\n\n # ── Architecture ───────────────────────────────────────────────────────\n EMBED_DIM = 32\n CONV_CH = [64, 128, 128]\n GRU_HIDDEN = 128\n GRU_LAYERS = 2\n ATTN_HEADS = 8\n DROPOUT = 0.20\n CBGAF_DIM = 128\n\n # ── Seeds ──────────────────────────────────────────────────────────────\n EVAL_SEEDS = [42, 123, 456, 789, 2024]\n\n N_DEV_CATS = 6\n N_CLASSES = 2\n N_DS_SRC = 5\n\n DEVICE_CAT_MAP = {\n \"CICIDS-2017\": {12: 4},\n \"CIC-IoT-2023\": {11: 0},\n \"Bot-IoT\": {9: 4},\n \"Edge-IIoTset\": {10: 3},\n \"N-BaIoT\": {0:0, 1:2, 2:1, 3:1, 4:1, 5:1, 6:2, 7:0, 8:0},\n }\n\n OUT = \"tch_net_results\"\n\nos.makedirs(Config.OUT, exist_ok=True)\nprint(\"Dataset availability:\")\nfor n, f in [(\"CICIDS-2017\", Config.USE_CICIDS), (\"CIC-IoT-2023\", Config.USE_CICIOT),\n (\"Bot-IoT\", Config.USE_BOTIOT), (\"Edge-IIoTset\", Config.USE_EDGE),\n (\"N-BaIoT\", Config.USE_NBAIOT)]:\n tag = \"PRIMARY\" if n in [\"CICIDS-2017\",\"CIC-IoT-2023\",\"Bot-IoT\"] else \"SUPPLEMENTARY\"\n print(f\" {n:<15}: {'FOUND' if f else 'NOT FOUND'} [{tag}]\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell4_md","cell_type":"markdown","source":"## Cell 4 — Dataset Loaders\n\nEach loader reads CSV/Parquet files, maps to the 46-feature canonical space, and returns `(X, y, device_ids)`.\nNo fabricated mappings — missing features are zero-filled.","metadata":{}},{"id":"cell4_code","cell_type":"code","source":"ATTACK_TYPE_STORE = {}\n\ndef _subsample(X, y, max_n, extras=None, seed=42):\n if max_n is None or len(X) <= max_n:\n return (X, y) + (tuple(extras) if extras else ())\n rng = np.random.RandomState(seed)\n ben = np.where(y==0)[0]; atk = np.where(y==1)[0]\n r = len(ben) / len(y)\n nb = min(len(ben), max(5_000, int(max_n * r)))\n na = min(len(atk), max(5_000, max_n - nb))\n if nb + na > max_n:\n scale = max_n / (nb + na)\n nb = max(5_000, int(nb * scale))\n na = max(5_000, min(max_n - nb, int(na * scale)))\n nb = min(nb, len(ben)); na = min(na, len(atk))\n keep = np.concatenate([rng.choice(ben, nb, replace=False),\n rng.choice(atk, na, replace=False)])\n rng.shuffle(keep)\n result = [X[keep], y[keep]]\n if extras:\n for a in extras: result.append(a[keep])\n print(f\" Subsampled {len(X):,} -> {len(keep):,} \"\n f\"(ben={nb:,} atk={na:,} atk%={na/(nb+na)*100:.1f}%)\")\n return tuple(result)\n\n\ndef load_cicids2017(data_dir, max_samples=None):\n print(\"\\n\"+\"=\"*70+\"\\nLOADING CICIDS-2017\\n\"+\"=\"*70)\n files = sorted(glob.glob(os.path.join(data_dir,\"**\",\"*.csv\"), recursive=True) +\n glob.glob(os.path.join(data_dir,\"**\",\"*.parquet\"), recursive=True))\n files = [f for f in files if not any(kw in os.path.basename(f).lower()\n for kw in ['readme','feature','description'])]\n if not files: print(f\" No files in {data_dir}\"); return None, None, None\n X_all, y_all = [], []\n for fp in files:\n print(f\" Loading {os.path.basename(fp)}...\")\n try:\n df = pd.read_parquet(fp) if fp.endswith('.parquet') else pd.read_csv(fp, low_memory=False)\n df.columns = df.columns.str.strip()\n lc = next((c for c in df.columns if c.lower().strip() == 'label'), None)\n if lc is None: print(f\" No label column\"); continue\n y = (~df[lc].astype(str).str.strip().str.upper().isin(['BENIGN','NORMAL','0'])).values.astype(np.int32)\n X_sem, matched = build_semantic_vector(df, \"CICIDS-2017\")\n np.nan_to_num(X_sem, copy=False, nan=0., posinf=1e9, neginf=-1e9)\n X_all.append(X_sem); y_all.append(y)\n print(f\" {len(y):,} rows | coverage={len(matched)}/{N_SEM}\")\n except Exception as e: print(f\" Skip: {e}\")\n if not X_all: return None, None, None\n X = np.vstack(X_all); y = np.hstack(y_all)\n dev = np.full(len(y), 12, dtype=np.int32)\n print(f\" Total: {len(y):,} | Ben={int((y==0).sum()):,} Atk={int((y==1).sum()):,}\")\n X, y, dev = _subsample(X, y, max_samples, [dev])\n return X, y, dev\n\n\ndef load_ciciot2023(data_dir, max_samples=None):\n print(\"\\n\"+\"=\"*70+\"\\nLOADING CIC-IoT-2023\\n\"+\"=\"*70)\n files = sorted(glob.glob(os.path.join(data_dir,\"**\",\"*.csv\"), recursive=True))\n files = [f for f in files if not any(kw in os.path.basename(f).lower()\n for kw in ['readme','feature','summary'])]\n if not files: print(f\" No files in {data_dir}\"); return None, None, None\n X_all, y_all = [], []; _acc_rows = 0; _printed_cols = False\n for fp in files:\n try:\n df = pd.read_csv(fp, low_memory=False)\n df.columns = df.columns.str.strip()\n if not _printed_cols:\n print(f\" Columns ({len(df.columns)}): {df.columns.tolist()[:15]}\")\n _printed_cols = True\n lc = next((c for c in df.columns if c.lower().strip() == 'label'), None)\n if lc is None: continue\n lbl = df[lc]\n if lbl.dtype == object:\n y = (~lbl.str.strip().str.lower().isin(['benign','benigntraffic','normal','0'])).values.astype(np.int32)\n else:\n y = (lbl.values != 0).astype(np.int32)\n X_sem, matched = build_semantic_vector(df, \"CIC-IoT-2023\")\n np.nan_to_num(X_sem, copy=False, nan=0., posinf=1e9, neginf=-1e9)\n X_all.append(X_sem); y_all.append(y); _acc_rows += len(y)\n print(f\" {os.path.basename(fp)}: {len(y):,} rows | cov={len(matched)}/{N_SEM}\")\n if max_samples is not None and _acc_rows >= max_samples * 2:\n print(f\" Early exit at {_acc_rows:,} rows\")\n break\n except Exception as e: print(f\" Skip {os.path.basename(fp)}: {e}\")\n if not X_all: return None, None, None\n X = np.vstack(X_all); y = np.hstack(y_all)\n dev = np.full(len(y), 11, dtype=np.int32)\n print(f\" Total: {len(y):,} | Ben={int((y==0).sum()):,} Atk={int((y==1).sum()):,}\")\n rng = np.random.RandomState(42)\n perm = rng.permutation(len(X)); X, y, dev = X[perm], y[perm], dev[perm]\n X, y, dev = _subsample(X, y, max_samples, [dev])\n return X, y, dev\n\n\ndef load_botiot(data_dir, max_samples=None):\n print(\"\\n\"+\"=\"*70+\"\\nLOADING Bot-IoT\\n\"+\"=\"*70)\n files = sorted(glob.glob(os.path.join(data_dir,\"**\",\"*.csv\"), recursive=True))\n files = [f for f in files if 'names' not in os.path.basename(f).lower()]\n if not files: print(f\" No files in {data_dir}\"); return None, None, None\n X_all, y_all = [], []; _acc_rows = 0\n rng = np.random.RandomState(42)\n for fp in files:\n try:\n df = pd.read_csv(fp, low_memory=False)\n df.columns = df.columns.str.strip()\n lc = next((c for c in df.columns if c.lower().strip() in\n {'label','attack','category','type','subcategory'}), None)\n if lc is None: continue\n lbl = df[lc]\n n = pd.to_numeric(lbl, errors='coerce')\n if n.notna().mean() > 0.5:\n y = (n.fillna(0) != 0).values.astype(np.int32)\n else:\n y = (~lbl.astype(str).str.strip().str.lower().isin(['0','normal','benign'])).values.astype(np.int32)\n X_sem, matched = build_semantic_vector(df, \"Bot-IoT\")\n np.nan_to_num(X_sem, copy=False, nan=0., posinf=1e9, neginf=-1e9)\n X_all.append(X_sem); y_all.append(y); _acc_rows += len(y)\n print(f\" {os.path.basename(fp)}: {len(y):,} rows | cov={len(matched)}/{N_SEM}\")\n if max_samples is not None and _acc_rows >= max_samples * 2: break\n except Exception as e: print(f\" Skip {os.path.basename(fp)}: {e}\")\n if not X_all: return None, None, None\n X = np.vstack(X_all); y = np.hstack(y_all)\n dev = np.full(len(y), 9, dtype=np.int32)\n print(f\" Total: {len(y):,} | Ben={int((y==0).sum()):,} Atk={int((y==1).sum()):,}\")\n perm = rng.permutation(len(X)); X, y, dev = X[perm], y[perm], dev[perm]\n X, y, dev = _subsample(X, y, max_samples, [dev])\n return X, y, dev\n\n\ndef load_edgeiiot(data_dir, max_samples=None):\n print(\"\\n\"+\"=\"*70+\"\\nLOADING Edge-IIoTset (supplementary)\\n\"+\"=\"*70)\n csv_files = glob.glob(os.path.join(data_dir,\"**\",\"*.csv\"), recursive=True)\n dnn_files = [f for f in csv_files if 'DNN' in os.path.basename(f)]\n targets = dnn_files if dnn_files else csv_files\n targets = [f for f in targets if not any(kw in os.path.basename(f).lower()\n for kw in ['readme','feature','metadata','description'])]\n if not targets: print(\" Not found\"); return None, None, None\n print(f\" Files: {[os.path.basename(t) for t in targets]}\")\n X_all, y_all = [], []\n for target in targets:\n try:\n df = pd.read_csv(target, low_memory=False)\n df.columns = df.columns.str.strip()\n lc = next((c for c in df.columns if c.lower().strip().replace(' ','_') in\n {'attack_type','attack_label','label','attack','class','type','category'}), None)\n if lc is None: continue\n y_f = (~df[lc].astype(str).str.strip().str.lower().isin(['normal','benign','0'])).values.astype(np.int32)\n X_f, matched = build_semantic_vector(df, \"Edge-IIoTset\")\n np.nan_to_num(X_f, copy=False, nan=0., posinf=1e9, neginf=-1e9)\n X_all.append(X_f); y_all.append(y_f)\n print(f\" {os.path.basename(target)}: {len(y_f):,} rows | cov={len(matched)}/{N_SEM}\")\n except Exception as e: print(f\" Skip {os.path.basename(target)}: {e}\")\n if not X_all: return None, None, None\n X_sem = np.vstack(X_all); y = np.hstack(y_all)\n dev = np.full(len(y), 10, dtype=np.int32)\n X_sem, y, dev = _subsample(X_sem, y, max_samples, [dev])\n return X_sem, y, dev\n\n\ndef load_nbaiot(data_dir, max_samples=None):\n print(\"\\n\"+\"=\"*70+\"\\nLOADING N-BaIoT (supplementary)\\n\"+\"=\"*70)\n csv_files = sorted(set(\n glob.glob(os.path.join(data_dir,\"*.csv\")) +\n glob.glob(os.path.join(data_dir,\"**\",\"*.csv\"), recursive=True)))\n csv_files = [f for f in csv_files if not any(\n e in os.path.basename(f).lower() for e in ['summary','features','readme'])]\n print(f\" Found {len(csv_files)} CSV files\")\n if not csv_files: return None, None, None\n dev_kw = {'danmini':0,'ecobee':1,'ennio':2,'philips':3,\n 'pt_737':4,'pt_838':5,'samsung':6,'xcs7_1002':7,'xcs7_1003':8}\n X_all, y_all, dev_all = [], [], []\n for fp in csv_files:\n fname = os.path.basename(fp).lower()\n search = (os.path.basename(os.path.dirname(fp))+' '+fname).lower()\n did = next((d for kw, d in dev_kw.items() if kw in search), 0)\n is_ben = 'benign' in fname\n try:\n df = pd.read_csv(fp, header=0)\n df.columns = df.columns.str.strip()\n for col in df.columns: df[col] = pd.to_numeric(df[col], errors='coerce')\n df.dropna(how='all', inplace=True)\n if df.empty: continue\n X_sem, _ = build_semantic_vector(df, \"N-BaIoT\")\n np.nan_to_num(X_sem, copy=False, nan=0., posinf=1e9, neginf=-1e9)\n X_all.append(X_sem)\n y_all.append(np.full(len(df), 0 if is_ben else 1, dtype=np.int32))\n dev_all.append(np.full(len(df), did, dtype=np.int32))\n except Exception as _e: print(f\" Skip {os.path.basename(fp)}: {_e}\")\n if not X_all: return None, None, None\n X = np.vstack(X_all); y = np.hstack(y_all); dev = np.hstack(dev_all)\n print(f\" Total: {len(y):,} | Ben={int((y==0).sum()):,} Atk={int((y==1).sum()):,}\")\n X, y, dev = _subsample(X, y, max_samples, [dev])\n return X, y, dev\n\nprint(\"All 5 dataset loaders defined\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell5_md","cell_type":"markdown","source":"## Cell 5 — Multi-Dataset Loader (Leakage-Free)\n\nPipeline: balance each dataset → stratified 80/20 split → fit RobustScaler on train → clip [-10, 10]","metadata":{}},{"id":"cell5_code","cell_type":"code","source":"class MultiDatasetLoader:\n MIN_FLOOR = 5_000\n\n def __init__(self):\n self.datasets = {}\n self.scaler = None\n\n def add(self, name, X, y, dev, ds_src_id):\n if X is None: print(f\" Skip {name}: no data\"); return\n cat_map = Config.DEVICE_CAT_MAP.get(name, {})\n dev_cats = np.array([cat_map.get(int(d), 5) for d in dev], dtype=np.int32)\n ds_ids = np.full(len(y), ds_src_id, dtype=np.int32)\n ctx = np.stack([ds_ids, dev_cats], axis=1)\n n_matched = int((X != 0).any(axis=0).sum())\n self.datasets[name] = {'X': X, 'y': y, 'dev': dev, 'ctx': ctx,\n 'ds_src': ds_src_id, 'n': X.shape[0]}\n print(f\" Added {name}: {X.shape[0]:,}x{X.shape[1]} \"\n f\"(non-zero dims: {n_matched}/{N_SEM} = {n_matched/N_SEM*100:.0f}%)\")\n\n def _balance(self, X, y, ctx, ratio, rng, name=\"\"):\n ben = np.where(y == 0)[0]; atk = np.where(y == 1)[0]\n if len(ben) == 0 or len(atk) == 0:\n print(f\" WARNING {name}: single-class data — skipping balance\"); return X, y, ctx\n if len(atk) > ratio * len(ben):\n atk = rng.choice(atk, int(ratio * len(ben)), replace=False)\n elif len(ben) > ratio * len(atk):\n ben = rng.choice(ben, int(ratio * len(atk)), replace=False)\n n_total = len(ben) + len(atk)\n if n_total < 2 * self.MIN_FLOOR:\n print(f\" ADVISORY {name}: only {n_total:,} samples after balancing \"\n f\"(minority: {min(len(ben),len(atk)):,}). FocalLoss compensates.\")\n keep = np.concatenate([ben, atk]); rng.shuffle(keep)\n return X[keep], y[keep], ctx[keep]\n\n def combine(self, ratio=None, test_size=0.2, seed=42):\n if ratio is None: ratio = Config.TARGET_ATK_BEN_RATIO\n rng = np.random.RandomState(seed)\n print(\"\\n\"+\"=\"*70+\"\\nCOMBINE & PREPROCESS\\n\"+\"=\"*70)\n Xa, ya, ca, da = [], [], [], []\n for di, (name, d) in enumerate(self.datasets.items()):\n X_r = d['X'].copy(); y_r = d['y'].copy(); c_r = d['ctx'].copy()\n np.nan_to_num(X_r, copy=False, nan=0., posinf=1e6, neginf=-1e6)\n n_nonzero = int((X_r != 0).any(axis=0).sum())\n if n_nonzero == 0:\n print(f\" SKIP {name}: 0% coverage\"); continue\n # derive pkt_count_total from fwd+bwd BEFORE scaling\n _fwd = SEM_IDX[\"pkt_count_fwd\"]; _bwd = SEM_IDX[\"pkt_count_bwd\"]\n _tot = SEM_IDX[\"pkt_count_total\"]\n zero_tot = (X_r[:, _tot] == 0)\n X_r[zero_tot, _tot] = X_r[zero_tot, _fwd] + X_r[zero_tot, _bwd]\n X_r, y_r, c_r = self._balance(X_r, y_r, c_r, ratio, rng, name)\n ben_r = int((y_r==0).sum()); atk_r = int((y_r==1).sum())\n print(f\" {name}: {len(y_r):,} rows (ben={ben_r:,} atk={atk_r:,})\")\n Xa.append(X_r); ya.append(y_r); ca.append(c_r)\n da.append(np.full(len(y_r), di, dtype=np.int32))\n if not Xa: raise RuntimeError(\"No datasets survived combine()\")\n X_c = np.vstack(Xa).astype(np.float32)\n y_c = np.hstack(ya).astype(np.int32)\n ctx_c = np.vstack(ca).astype(np.int32)\n ds_c = np.hstack(da).astype(np.int32)\n X_tr, X_te, y_tr, y_te, c_tr, c_te, d_tr, d_te = train_test_split(\n X_c, y_c, ctx_c, ds_c, test_size=test_size, stratify=y_c, random_state=seed)\n self.scaler = RobustScaler(quantile_range=(5, 95))\n X_tr = np.clip(self.scaler.fit_transform(X_tr), -10, 10).astype(np.float32)\n X_te = np.clip(self.scaler.transform(X_te), -10, 10).astype(np.float32)\n tr_ben = int((y_tr==0).sum()); tr_atk = int((y_tr==1).sum())\n te_ben = int((y_te==0).sum()); te_atk = int((y_te==1).sum())\n print(f\" Train: {len(y_tr):,} | ben={tr_ben:,} ({tr_ben/len(y_tr)*100:.1f}%) \"\n f\"| atk={tr_atk:,} ({tr_atk/len(y_tr)*100:.1f}%)\")\n print(f\" Test: {len(y_te):,} | ben={te_ben:,} ({te_ben/len(y_te)*100:.1f}%) \"\n f\"| atk={te_atk:,} ({te_atk/len(y_te)*100:.1f}%)\")\n for tag_, y_ in [('Train', y_tr), ('Test', y_te)]:\n r_ = float((y_==1).mean())\n if r_ < 0.05 or r_ > 0.95:\n raise RuntimeError(f'{tag_} severely imbalanced: {r_*100:.1f}%')\n return X_tr, y_tr, c_tr, d_tr, X_te, y_te, c_te, d_te\n\nloader = MultiDatasetLoader()\nDS_SRC = {\"CICIDS-2017\": 0, \"CIC-IoT-2023\": 1, \"Bot-IoT\": 2,\n \"Edge-IIoTset\": 3, \"N-BaIoT\": 4}\nprint(\"MultiDatasetLoader ready\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell6_md","cell_type":"markdown","source":"## Cell 6 — Load All Datasets","metadata":{}},{"id":"cell6_code","cell_type":"code","source":"print(\"=\"*70+\"\\nLOADING DATASETS\\n\"+\"=\"*70)\n\nif Config.USE_CICIDS:\n try:\n X,y,d = load_cicids2017(Config.CICIDS_PATH, Config.CICIDS_MAX)\n if X is not None: loader.add('CICIDS-2017', X, y, d, DS_SRC['CICIDS-2017'])\n del X,y,d\n except Exception as e: print(f\"CICIDS-2017 failed: {e}\")\n\nif Config.USE_CICIOT:\n try:\n X,y,d = load_ciciot2023(Config.CICIOT_PATH, Config.CICIOT_MAX)\n if X is not None: loader.add('CIC-IoT-2023', X, y, d, DS_SRC['CIC-IoT-2023'])\n del X,y,d\n except Exception as e: print(f\"CIC-IoT-2023 failed: {e}\")\n\nif Config.USE_BOTIOT:\n try:\n X,y,d = load_botiot(Config.BOTIOT_PATH, Config.BOTIOT_MAX)\n if X is not None: loader.add('Bot-IoT', X, y, d, DS_SRC['Bot-IoT'])\n del X,y,d\n except Exception as e: print(f\"Bot-IoT failed: {e}\")\n\nif Config.USE_EDGE:\n try:\n X,y,d = load_edgeiiot(Config.EDGE_PATH, Config.EDGE_MAX)\n if X is not None: loader.add('Edge-IIoTset', X, y, d, DS_SRC['Edge-IIoTset'])\n del X,y,d\n except Exception as e: print(f\"Edge-IIoTset failed: {e}\")\n\nif Config.USE_NBAIOT:\n try:\n X,y,d = load_nbaiot(Config.NBAIOT_PATH, Config.NBAIOT_MAX)\n if X is not None: loader.add('N-BaIoT', X, y, d, DS_SRC['N-BaIoT'])\n del X,y,d\n except Exception as e: print(f\"N-BaIoT failed: {e}\")\n\ngc.collect()\nDS_NAMES = list(loader.datasets.keys())\nprimary = [n for n in DS_NAMES if n in ['CICIDS-2017','CIC-IoT-2023','Bot-IoT']]\nsupplementary = [n for n in DS_NAMES if n in ['Edge-IIoTset','N-BaIoT']]\nprint(f\"\\n{len(DS_NAMES)} datasets loaded: {DS_NAMES}\")\nprint(f\" Primary ({len(primary)}): {primary}\")\nprint(f\" Supplementary ({len(supplementary)}): {supplementary}\")\nassert len(primary) >= 1, \"No primary datasets loaded\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell7_md","cell_type":"markdown","source":"## Cell 7 — Combine, Scale & Create Sequences","metadata":{}},{"id":"cell7_code","cell_type":"code","source":"(X_train_raw, y_train_raw, ctx_train_raw, ds_train_raw,\n X_test_raw, y_test_raw, ctx_test_raw, ds_test_raw) = loader.combine()\n\nN_FEATURES = N_SEM\nN_DATASETS = len(loader.datasets)\nConfig.N_DS_SRC = N_DATASETS\nprint(f\"N_FEATURES={N_FEATURES} N_DATASETS={N_DATASETS} DS_NAMES={DS_NAMES}\")\n\n\ndef create_sequences(X, y, ctx, ds, window, stride):\n \"\"\"Sliding-window sequencing — vectorised via fancy-index gather.\n Dataset-boundary-safe. Label = majority vote.\"\"\"\n n_feat = X.shape[1]; n_ctx = ctx.shape[1]\n d_ids = np.unique(ds)\n counts = {d: max(0, (int((ds==d).sum()) - window) // stride + 1)\n if int((ds==d).sum()) >= window else 0 for d in d_ids}\n total = sum(counts.values())\n if total == 0: raise ValueError(\"No sequences — dataset too small for given window/stride\")\n Xs = np.empty((total, window, n_feat), dtype=np.float32)\n ys = np.empty(total, dtype=np.int32)\n cs = np.empty((total, n_ctx), dtype=np.int32)\n dss = np.empty(total, dtype=np.int32)\n ptr = 0\n for d_id in d_ids:\n n_seq = counts[d_id]\n if n_seq == 0: continue\n m = (ds == d_id); Xd = X[m]; yd = y[m]; cd = ctx[m]\n start_idx = np.arange(n_seq, dtype=np.int32) * stride\n col_idx = np.arange(window, dtype=np.int32)\n idx2d = start_idx[:, None] + col_idx[None, :]\n Xs[ptr:ptr+n_seq] = Xd[idx2d]\n ys[ptr:ptr+n_seq] = (yd[idx2d].mean(axis=1) > 0.5).astype(np.int32)\n cs[ptr:ptr+n_seq] = cd[start_idx]\n dss[ptr:ptr+n_seq] = d_id\n ptr += n_seq\n return Xs, ys, cs, dss\n\n\ndef _cap(X, y, c, d, mx, tag):\n \"\"\"Cap to mx sequences enforcing 1:1 class balance.\"\"\"\n if len(y) <= mx:\n ben = int((y==0).sum()); atk = int((y==1).sum())\n print(f\" [{tag}] {len(y):,} seqs (ben={ben:,} {ben/len(y)*100:.1f}% | atk={atk:,} {atk/len(y)*100:.1f}%)\")\n return X, y, c, d\n rng = np.random.RandomState(42)\n ben_idx = np.where(y==0)[0]; atk_idx = np.where(y==1)[0]\n n_each = mx // 2\n nb = min(len(ben_idx), n_each); na = min(len(atk_idx), mx - nb)\n if nb < n_each: na = min(len(atk_idx), mx - nb)\n if na < n_each: nb = min(len(ben_idx), mx - na)\n keep = np.concatenate([rng.choice(ben_idx, nb, replace=False),\n rng.choice(atk_idx, na, replace=False)])\n rng.shuffle(keep)\n print(f\" [{tag}] Capped {len(y):,} -> {len(keep):,} \"\n f\"(ben={nb:,} {nb/len(keep)*100:.1f}% | atk={na:,} {na/len(keep)*100:.1f}%)\")\n return X[keep], y[keep], c[keep], d[keep]\n\n\nprint(\"\\nCreating sliding-window sequences...\")\nX_train_s, y_train_s, ctx_train_s, ds_train_s = create_sequences(\n X_train_raw, y_train_raw, ctx_train_raw, ds_train_raw, Config.WINDOW_SIZE, Config.STRIDE)\nX_test_s, y_test_s, ctx_test_s, ds_test_s = create_sequences(\n X_test_raw, y_test_raw, ctx_test_raw, ds_test_raw, Config.WINDOW_SIZE, Config.STRIDE)\n\nprint(f\" Raw sequences — Train: {len(y_train_s):,} \"\n f\"(atk={int((y_train_s==1).sum()):,} {(y_train_s==1).mean()*100:.1f}%) | \"\n f\"Test: {len(y_test_s):,} \"\n f\"(atk={int((y_test_s==1).sum()):,} {(y_test_s==1).mean()*100:.1f}%)\")\n\nX_train, y_train, ctx_train, ds_train = _cap(\n X_train_s, y_train_s, ctx_train_s, ds_train_s, Config.MAX_TRAIN_SEQ, \"train\")\nX_test, y_test, ctx_test, ds_test = _cap(\n X_test_s, y_test_s, ctx_test_s, ds_test_s, Config.MAX_TEST_SEQ, \"test\")\n\nfor _v in ['X_train_raw','X_test_raw','X_train_s','X_test_s']:\n try: exec(f'del {_v}')\n except: pass\ngc.collect()\n\nprint(f\"\\nFinal shapes — Train: {X_train.shape} | Test: {X_test.shape}\")\nassert X_train.dtype == np.float32\nassert X_train.shape[1] == Config.WINDOW_SIZE\nassert X_train.shape[2] == N_FEATURES\nassert (y_train==1).mean() > 0.10\nassert not np.isnan(X_train).any()\nprint(\"All assertions passed — data is clean and balanced\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell8_md","cell_type":"markdown","source":"## Cell 8 — Dataset & DataLoader","metadata":{}},{"id":"cell8_code","cell_type":"code","source":"_NW = 2 if os.path.exists('/kaggle') else 0\n\nclass BotnetDataset(Dataset):\n def __init__(self, X, y, ctx):\n self.X = torch.FloatTensor(X)\n self.y = torch.LongTensor(y)\n self.ctx = torch.LongTensor(ctx)\n def __len__(self): return len(self.X)\n def __getitem__(self, i):\n return {'sequence': self.X[i], 'label': self.y[i], 'context': self.ctx[i]}\n\ndef make_loaders(Xtr, ytr, ctr, Xte, yte, cte, bs=None):\n if bs is None: bs = Config.BATCH_SIZE\n pw = (_NW > 0)\n tr = DataLoader(BotnetDataset(Xtr, ytr, ctr), batch_size=bs,\n shuffle=True, num_workers=_NW,\n pin_memory=torch.cuda.is_available(), drop_last=True, persistent_workers=pw)\n te = DataLoader(BotnetDataset(Xte, yte, cte), batch_size=bs, shuffle=False,\n num_workers=_NW, pin_memory=torch.cuda.is_available(), persistent_workers=pw)\n return tr, te\n\ntrain_loader, test_loader = make_loaders(X_train, y_train, ctx_train, X_test, y_test, ctx_test)\nprint(f\"DataLoaders: {len(train_loader)} train | {len(test_loader)} test batches\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell9_md","cell_type":"markdown","source":"## Cell 9 — TCH-Net Architecture\n\n**T-Branch** — MSTE (Multi-Scale Temporal Encoding):\n- Path 1: ResConv blocks → BiGRU(128, 2-layer) — local patterns\n- Path 2: Stride-2 Conv → BiGRU(64) — medium-scale patterns\n- Path 3: Full-resolution Transformer (Pre-LN) — global context\n\n**H-Branch** — time-averaged features → MLP → 64d statistical summary\n\n**C-Branch** — dataset source + device category embeddings → 64d\n\n**CB-GAF** — Cross-Branch Gated Attention fuses all three → 384d → classifier","metadata":{}},{"id":"cell9_code","cell_type":"code","source":"class DepthwiseSepConv1d(nn.Module):\n def __init__(self, inc, outc, ks=3, pad=1):\n super().__init__()\n self.dw = nn.Conv1d(inc, inc, ks, padding=pad, groups=inc, bias=False)\n self.pw = nn.Conv1d(inc, outc, 1, bias=False)\n self.bn = nn.BatchNorm1d(outc)\n def forward(self, x): return F.relu(self.bn(self.pw(self.dw(x))))\n\nclass SEBlock1d(nn.Module):\n def __init__(self, ch, r=8):\n super().__init__()\n self.fc = nn.Sequential(nn.AdaptiveAvgPool1d(1), nn.Flatten(),\n nn.Linear(ch, max(ch//r, 4)), nn.ReLU(),\n nn.Linear(max(ch//r, 4), ch), nn.Sigmoid())\n def forward(self, x): return x * self.fc(x).unsqueeze(-1)\n\nclass ResConvBlock(nn.Module):\n def __init__(self, inc, outc):\n super().__init__()\n self.c1 = DepthwiseSepConv1d(inc, outc); self.c2 = DepthwiseSepConv1d(outc, outc)\n self.se = SEBlock1d(outc)\n self.skip = nn.Conv1d(inc, outc, 1, bias=False) if inc != outc else nn.Identity()\n self.bn = nn.BatchNorm1d(outc) if inc != outc else nn.Identity()\n def forward(self, x):\n out = self.se(self.c2(self.c1(x))); sk = self.bn(self.skip(x))\n if out.shape[-1] != sk.shape[-1]: sk = F.adaptive_avg_pool1d(sk, out.shape[-1])\n return F.relu(out + sk)\n\nclass CrossBranchGatedAttention(nn.Module):\n def __init__(self, dt, dc, dh, out):\n super().__init__()\n self.out = out; self.scale = out**0.5\n self.pt = nn.Linear(dt, out); self.pc = nn.Linear(dc, out); self.ph = nn.Linear(dh, out)\n for b in 'tch':\n for m in 'qkv': setattr(self, f'{m}_{b}', nn.Linear(out, out))\n self.gt = nn.Sequential(nn.Linear(out*2, out), nn.Sigmoid())\n self.gc = nn.Sequential(nn.Linear(out*2, out), nn.Sigmoid())\n self.gh = nn.Sequential(nn.Linear(out*2, out), nn.Sigmoid())\n self.ln = nn.LayerNorm(out*3)\n\n def _ca(self, q, K, V):\n q = q.unsqueeze(1)\n a = F.softmax(torch.bmm(q, K.transpose(1,2)) / self.scale, dim=-1)\n return torch.bmm(a, V).squeeze(1)\n\n def forward(self, ht, hc, hh):\n t = self.pt(ht); c = self.pc(hc); h = self.ph(hh)\n qt,kt,vt = self.q_t(t),self.k_t(t),self.v_t(t)\n qc,kc,vc = self.q_c(c),self.k_c(c),self.v_c(c)\n qh,kh,vh = self.q_h(h),self.k_h(h),self.v_h(h)\n ct_ = self._ca(qt, torch.stack([kc,kh],1), torch.stack([vc,vh],1))\n cc_ = self._ca(qc, torch.stack([kt,kh],1), torch.stack([vt,vh],1))\n ch_ = self._ca(qh, torch.stack([kt,kc],1), torch.stack([vt,vc],1))\n gt = self.gt(torch.cat([t,ct_],-1)); gc = self.gc(torch.cat([c,cc_],-1))\n gh = self.gh(torch.cat([h,ch_],-1))\n return self.ln(torch.cat([gt*t+(1-gt)*ct_, gc*c+(1-gc)*cc_, gh*h+(1-gh)*ch_], -1))\n\n\nclass TCHNet(nn.Module):\n \"\"\"TCH-Net — T(Temporal) + C(Context) + H(Statistical) branches with CB-GAF.\"\"\"\n def __init__(self, nf, ws, n_ds=None, n_dc=None, ed=None, cc=None, nc=None,\n gh=None, gl=None, ah=None, do=None, cd=None):\n super().__init__()\n n_ds=n_ds or Config.N_DS_SRC; n_dc=n_dc or Config.N_DEV_CATS\n nc=nc or Config.N_CLASSES; self.nc=nc\n ed=ed or Config.EMBED_DIM; cc=cc or Config.CONV_CH\n gh=gh or Config.GRU_HIDDEN; gl=gl or Config.GRU_LAYERS\n ah=ah or Config.ATTN_HEADS; do=do if do is not None else Config.DROPOUT\n cd=cd or Config.CBGAF_DIM\n\n # ── T-branch Path 1: Multi-Scale Conv-GRU ────────────────────────\n ls, ic = [], nf\n for i, oc in enumerate(cc):\n ls.append(ResConvBlock(ic, oc))\n ls.append(nn.MaxPool1d(2,2) if i < len(cc)-1 else nn.AdaptiveAvgPool1d(8))\n ic = oc\n self.t_conv = nn.Sequential(*ls)\n self.t_gru1 = nn.GRU(cc[-1], gh, gl, batch_first=True, bidirectional=True,\n dropout=do if gl>1 else 0)\n s1 = gh * 2\n\n # ── T-branch Path 2: Stride-Conv GRU ─────────────────────────────\n self.t_down = nn.Sequential(nn.Conv1d(nf, cc[0], 3, stride=2, padding=1, bias=False),\n nn.BatchNorm1d(cc[0]), nn.ReLU())\n gh2 = gh // 2\n self.t_gru2 = nn.GRU(cc[0], gh2, 1, batch_first=True, bidirectional=True)\n s2 = gh2 * 2\n\n # ── T-branch Path 3: Full-Resolution Transformer ──────────────────\n t_dm = 128\n self.t_proj = nn.Linear(nf, t_dm)\n self.t_pos = nn.Parameter(torch.randn(1, ws, t_dm) * 0.02)\n _tel = nn.TransformerEncoderLayer(t_dm, 8, t_dm*4, do, batch_first=True, norm_first=True)\n self.t_cls = nn.Parameter(torch.zeros(1, 1, t_dm))\n self.t_enc = nn.TransformerEncoder(_tel, 2)\n st = t_dm\n\n # ── T-branch merge ────────────────────────────────────────────────\n td = s1 + s2 + st\n h_ = ah\n while td % h_ != 0 and h_ > 1: h_ -= 1\n self.feat_proj = nn.Sequential(\n nn.Linear(nf, nf*2), nn.LayerNorm(nf*2), nn.GELU(),\n nn.Dropout(do*0.5), nn.Linear(nf*2, nf), nn.LayerNorm(nf))\n self.t_ln = nn.LayerNorm(td)\n self.t_mha = nn.MultiheadAttention(td, h_, batch_first=True, dropout=do)\n\n # ── H-branch ──────────────────────────────────────────────────────\n self.h_mlp = nn.Sequential(nn.Linear(nf,128), nn.BatchNorm1d(128), nn.GELU(),\n nn.Dropout(do), nn.Linear(128,64), nn.BatchNorm1d(64),\n nn.GELU(), nn.Dropout(do))\n hd = 64\n\n # ── C-branch ──────────────────────────────────────────────────────\n self.c_ds = nn.Embedding(max(n_ds,1), ed)\n self.c_cat = nn.Embedding(max(n_dc,1), ed)\n cod = ed * 2\n\n # ── CB-GAF ────────────────────────────────────────────────────────\n self.cbgaf = CrossBranchGatedAttention(td, cod, hd, cd)\n fd = cd * 3\n\n # ── Classifier head ───────────────────────────────────────────────\n self.raw_proj = nn.Sequential(nn.Linear(nf, 64), nn.BatchNorm1d(64), nn.GELU())\n _clf_in = fd + 64\n self.clf1 = nn.Sequential(nn.Linear(_clf_in, 256), nn.BatchNorm1d(256),\n nn.GELU(), nn.Dropout(do))\n self.clf2 = nn.Sequential(nn.Linear(256, 128), nn.BatchNorm1d(128),\n nn.GELU(), nn.Dropout(do))\n self.clf_res = nn.Linear(_clf_in, 128)\n self.clf_out = nn.Linear(128, nc)\n\n # ── Aux reconstruction ────────────────────────────────────────────\n self.aux = nn.Sequential(nn.Linear(fd, 64), nn.GELU(), nn.Linear(64, nf))\n\n self._td=td; self._cd=cod; self._hd=hd; self._fd=fd\n\n def forward(self, x, ctx=None, return_features=False):\n B = x.shape[0]\n x = x + self.feat_proj(x)\n xt = x.transpose(1, 2)\n\n # T-branch\n c1 = self.t_conv(xt).transpose(1, 2)\n g1, _ = self.t_gru1(c1)\n c2 = self.t_down(xt).transpose(1, 2)\n g2, _ = self.t_gru2(c2)\n g2e = F.adaptive_avg_pool1d(g2.transpose(1,2), g1.size(1)).transpose(1,2)\n t_tok = self.t_proj(x) + self.t_pos[:, :x.size(1), :]\n cls_tok = self.t_cls.expand(B, -1, -1)\n t_in = torch.cat([cls_tok, t_tok], dim=1)\n t_out = self.t_enc(t_in)[:, 1:, :]\n t_e = F.adaptive_avg_pool1d(t_out.transpose(1,2), g1.size(1)).transpose(1,2)\n gc = torch.cat([g1, g2e, t_e], dim=-1)\n ao, _ = self.t_mha(self.t_ln(gc), self.t_ln(gc), self.t_ln(gc))\n ht = ao.mean(1)\n\n # H-branch\n hh = self.h_mlp(x.mean(1))\n\n # C-branch\n if ctx is not None:\n hc = torch.cat([self.c_ds(ctx[:,0]), self.c_cat(ctx[:,1])], -1)\n else:\n hc = torch.zeros(B, self._cd, device=x.device)\n\n # CB-GAF\n fused = self.cbgaf(ht, hc, hh)\n raw = self.raw_proj(x.mean(1))\n combined = torch.cat([fused, raw], dim=-1)\n\n h1 = self.clf1(combined)\n h2 = self.clf2(h1) + self.clf_res(combined)\n logits = self.clf_out(h2)\n recon = self.aux(fused)\n\n if return_features: return logits, recon, fused\n return logits, recon\n\n\ndef make_tch_net(): return TCHNet(N_FEATURES, Config.WINDOW_SIZE, nc=Config.N_CLASSES)\n\nm = make_tch_net().to(device)\nnp_ = sum(p.numel() for p in m.parameters() if p.requires_grad)\nprint(f\"TCH-Net: {np_:,} params | T={m._td}d H={m._hd}d C={m._cd}d -> CB-GAF -> {m._fd}d\")\ndel m","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell10_md","cell_type":"markdown","source":"## Cell 10 — Training Infrastructure","metadata":{}},{"id":"cell10_code","cell_type":"code","source":"_amp_en = torch.cuda.is_available()\n_amp_dev = 'cuda' if _amp_en else 'cpu'\ndef _make_scaler(): return torch.amp.GradScaler(enabled=_amp_en)\n\nclass FocalLoss(nn.Module):\n def __init__(self, gamma=2., alpha=None, ls=0.):\n super().__init__(); self.gamma=gamma; self.alpha=alpha; self.ls=ls\n def forward(self, lg, tgt):\n lf = lg.float(); af = self.alpha.float() if self.alpha is not None else None\n ce = F.cross_entropy(lf, tgt, weight=af, reduction='none', label_smoothing=self.ls)\n pt = torch.exp(-ce.clamp(max=80)).clamp(1e-7, 1-1e-7)\n return ((1-pt)**self.gamma * ce).mean()\n\ndef get_sched(opt, wu, total):\n def fn(ep):\n if ep < wu: return (ep+1)/wu\n prog = (ep-wu)/max(total-wu,1)\n return max(0.05, 0.5*(1+np.cos(np.pi*prog)))\n return optim.lr_scheduler.LambdaLR(opt, fn)\n\ndef make_criterion(y=None):\n if y is None: y = y_train\n cc = np.bincount(y)\n w = torch.FloatTensor([1., cc[0]/max(cc[1],1)]).to(device)\n w = w / w.sum() * 2\n return FocalLoss(Config.FOCAL_GAMMA, w, Config.LABEL_SMOOTH)\n\ndef train_ep_tch(mdl, dl, crit, opt, aw, scaler):\n mdl.train(); tl=c=t=0; mse=nn.MSELoss()\n for b in tqdm(dl, desc=\" batches\", leave=False):\n s,l,ctx = b['sequence'].to(device), b['label'].to(device), b['context'].to(device)\n opt.zero_grad()\n with torch.amp.autocast(_amp_dev, enabled=_amp_en):\n lg, rc = mdl(s, ctx); loss = crit(lg,l) + aw*mse(rc, s.mean(1))\n scaler.scale(loss).backward(); scaler.unscale_(opt)\n torch.nn.utils.clip_grad_norm_(mdl.parameters(), 1.0)\n scaler.step(opt); scaler.update()\n tl += loss.item(); _, p = lg.max(1); c += (p==l).sum().item(); t += l.size(0)\n return tl/len(dl), c/t\n\n@torch.no_grad()\ndef evaluate(mdl, dl, crit, verbose=True):\n mdl.eval(); ps=[]; ls=[]; pbs=[]; tl=0\n _dl = tqdm(dl, desc=\" eval\", leave=False) if verbose else dl\n for b in _dl:\n s,l,ctx = b['sequence'].to(device), b['label'].to(device), b['context'].to(device)\n with torch.amp.autocast(_amp_dev, enabled=_amp_en):\n lg, _ = mdl(s, ctx)\n tl += crit(lg.float(), l).item()\n pr = F.softmax(lg.float(),-1); _, pd = lg.max(1)\n ps.extend(pd.cpu().numpy()); ls.extend(l.cpu().numpy()); pbs.extend(pr[:,1].cpu().numpy())\n yt=np.array(ls); yp=np.array(ps); ypr=np.array(pbs)\n try: roc=roc_auc_score(yt,ypr)\n except: roc=0.\n try: pp,pr_,_=precision_recall_curve(yt,ypr); prauc=auc(pr_,pp)\n except: prauc=0.\n try: fa,ta,_=roc_curve(yt,ypr); fpr99=fa[min(np.searchsorted(ta,0.99),len(fa)-1)]\n except: fpr99=1.\n cr = classification_report(yt,yp,target_names=['Benign','Attack'],output_dict=True,zero_division=0)\n return {'accuracy':accuracy_score(yt,yp),'precision':precision_score(yt,yp,zero_division=0),\n 'recall':recall_score(yt,yp,zero_division=0),'f1':f1_score(yt,yp,zero_division=0),\n 'roc_auc':roc,'mcc':matthews_corrcoef(yt,yp),'pr_auc':prauc,'fpr_at_tpr99':fpr99,\n 'benign_f1':cr['Benign']['f1-score'],'attack_f1':cr['Attack']['f1-score'],\n 'loss':tl/max(len(dl),1),'predictions':yp,'labels':yt,'probabilities':ypr}\n\ndef train_full(mdl, trdl, tedl, crit, epochs, lr, wd, es, wu=5, aw=0.1, verbose=True):\n opt = optim.AdamW(mdl.parameters(), lr=lr, weight_decay=wd)\n sch = get_sched(opt, wu, epochs)\n scaler = _make_scaler()\n bf=0; bs_=None; pat=0; hist=[]\n ep_iter = tqdm(range(epochs), desc=\"Epochs\", leave=True) if verbose else range(epochs)\n for ep in ep_iter:\n tl,_ = train_ep_tch(mdl,trdl,crit,opt,aw,scaler)\n sch.step(); m = evaluate(mdl,tedl,crit, verbose=verbose)\n m['train_loss']=tl\n _m_hist = {k:v for k,v in m.items() if k not in ('predictions','labels','probabilities')}\n hist.append(_m_hist)\n is_best = m['f1'] > bf\n if verbose:\n _f1=m['f1']; _auc=m['roc_auc']; _mcc=m['mcc']\n _bfv=m['benign_f1']; _afv=m['attack_f1']\n btag = \" \\u2605\" if is_best else \"\"\n if hasattr(ep_iter,'set_postfix'):\n ep_iter.set_postfix(loss=\"{:.4f}\".format(tl), f1=\"{:.4f}\".format(_f1),\n auc=\"{:.4f}\".format(_auc), best=\"{:.4f}\".format(bf))\n tqdm.write(\" Ep {:3d}/{:d} Loss={:.4f} F1={:.4f} AUC={:.4f} MCC={:.4f}\"\n \" BenF1={:.4f} AtkF1={:.4f}{}\".format(\n ep+1, epochs, tl, _f1, _auc, _mcc, _bfv, _afv, btag))\n if is_best: bf=m['f1']; bs_=copy.deepcopy(mdl.state_dict()); pat=0\n else:\n pat += 1\n if pat >= es:\n if verbose: tqdm.write(\" \\u23f9 Early stop ep {:d} Best F1={:.4f}\".format(ep+1,bf))\n break\n if bs_: mdl.load_state_dict(bs_)\n return evaluate(mdl,tedl,crit), hist, bs_\n\ndef set_seed(s):\n np.random.seed(s); torch.manual_seed(s)\n if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)\n\ndef multi_seed_tch(seeds, epochs,\n Xtr=None, ytr=None, ctr=None, Xte=None, yte=None, cte=None):\n \"\"\"Train TCH-Net across multiple seeds, returning metrics and best state dicts.\"\"\"\n if Xtr is None: Xtr,ytr,ctr = X_train,y_train,ctx_train\n if Xte is None: Xte,yte,cte = X_test,y_test,ctx_test\n all_m = []; all_h = []; all_sd = {}\n for i, s in enumerate(seeds):\n print(\" -- Seed {} [{}/{}] --\".format(s, i+1, len(seeds)))\n set_seed(s)\n mdl = make_tch_net().to(device)\n tl, tel = make_loaders(Xtr, ytr, ctr, Xte, yte, cte)\n cr = make_criterion(ytr)\n m, h_, best_sd = train_full(mdl, tl, tel, cr, epochs,\n Config.LR, Config.WD, Config.EARLY_STOP,\n Config.WARMUP, Config.AUX_WT)\n print(\" | F1={:.4f} AUC={:.4f} MCC={:.4f} PR-AUC={:.4f}\".format(\n m['f1'], m['roc_auc'], m['mcc'], m['pr_auc']))\n print(\" | Acc={:.4f} Prec={:.4f} Rec={:.4f}\".format(\n m['accuracy'], m['precision'], m['recall']))\n print(\" | BenF1={:.4f} AtkF1={:.4f} FPR@TPR99={:.4f}\".format(\n m['benign_f1'], m['attack_f1'], m['fpr_at_tpr99']))\n all_m.append(m); all_h.append(h_)\n if best_sd is not None: all_sd[s] = best_sd\n del mdl; gc.collect()\n if torch.cuda.is_available(): torch.cuda.empty_cache()\n return all_m, all_h, all_sd\n\ndef summarise(ml):\n ks = ['accuracy','precision','recall','f1','roc_auc','mcc','pr_auc',\n 'benign_f1','attack_f1','fpr_at_tpr99']\n s = {}\n for k in ks:\n v = [m[k] for m in ml if k in m]\n if v:\n s[k+'_mean'] = float(np.mean(v))\n s[k+'_std'] = float(np.std(v))\n s[k+'_vals'] = v\n return s\n\nprint(\"Training infrastructure ready\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell11_md","cell_type":"markdown","source":"## Cell 11 — Train TCH-Net (5 Seeds)","metadata":{}},{"id":"cell11_code","cell_type":"code","source":"print(\"=\"*70+\"\\nTRAINING TCH-Net (5 seeds)\\n\"+\"=\"*70)\n\ntch_metrics, tch_history, tch_state_dicts = multi_seed_tch(\n seeds=Config.EVAL_SEEDS,\n epochs=Config.EPOCHS\n)\n\ntch_summary = summarise(tch_metrics)\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"TCH-Net — 5-Seed Summary\")\nprint(\"=\"*70)\nmetrics_to_show = [\n ('F1 Score', 'f1'),\n ('ROC-AUC', 'roc_auc'),\n ('MCC', 'mcc'),\n ('PR-AUC', 'pr_auc'),\n ('Accuracy', 'accuracy'),\n ('Precision', 'precision'),\n ('Recall', 'recall'),\n ('Benign F1', 'benign_f1'),\n ('Attack F1', 'attack_f1'),\n ('FPR@TPR99', 'fpr_at_tpr99'),\n]\nfor label, key in metrics_to_show:\n mean = tch_summary.get(f'{key}_mean', 0)\n std = tch_summary.get(f'{key}_std', 0)\n vals = tch_summary.get(f'{key}_vals', [])\n vals_str = ' '.join(f'{v:.4f}' for v in vals)\n print(f\" {label:<16} {mean:.4f} ± {std:.4f} [{vals_str}]\")\n\nprint(\"\\nPer-seed best F1s:\", [f\"{m['f1']:.4f}\" for m in tch_metrics])\n\n# Save summary to JSON\n_summary_out = {k: v for k, v in tch_summary.items() if not isinstance(v, np.ndarray)}\nwith open(os.path.join(Config.OUT, 'tch_net_summary.json'), 'w') as f:\n json.dump(_summary_out, f, indent=2, default=str)\nprint(f\"\\nSummary saved to {Config.OUT}/tch_net_summary.json\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell12_md","cell_type":"markdown","source":"## Cell 12 — Save Best Checkpoints\n\nSaves one `.pth` file per seed plus a combined metadata file.\nThese are the files to upload to Hugging Face Hub.","metadata":{}},{"id":"cell12_code","cell_type":"code","source":"import json\n\nckpt_dir = os.path.join(Config.OUT, 'checkpoints')\nos.makedirs(ckpt_dir, exist_ok=True)\n\n# ── Save one checkpoint per seed ─────────────────────────────────────────\nsaved_paths = {}\nfor seed, state_dict in tch_state_dicts.items():\n path = os.path.join(ckpt_dir, f'tch_net_seed_{seed}.pth')\n torch.save({\n 'seed': seed,\n 'state_dict': state_dict,\n 'config': {\n 'n_features': N_FEATURES,\n 'window_size': Config.WINDOW_SIZE,\n 'n_classes': Config.N_CLASSES,\n 'n_ds_src': Config.N_DS_SRC,\n 'n_dev_cats': Config.N_DEV_CATS,\n 'embed_dim': Config.EMBED_DIM,\n 'conv_ch': Config.CONV_CH,\n 'gru_hidden': Config.GRU_HIDDEN,\n 'gru_layers': Config.GRU_LAYERS,\n 'attn_heads': Config.ATTN_HEADS,\n 'dropout': Config.DROPOUT,\n 'cbgaf_dim': Config.CBGAF_DIM,\n },\n 'metrics': {k: v for k, v in tch_metrics[Config.EVAL_SEEDS.index(seed)].items()\n if not isinstance(v, np.ndarray)},\n }, path)\n saved_paths[seed] = path\n f1 = tch_metrics[Config.EVAL_SEEDS.index(seed)]['f1']\n print(f\" Saved seed {seed} checkpoint -> {path} (F1={f1:.4f})\")\n\n# ── Save the best single checkpoint (highest F1) ─────────────────────────\nbest_idx = int(np.argmax([m['f1'] for m in tch_metrics]))\nbest_seed = Config.EVAL_SEEDS[best_idx]\nbest_sd = tch_state_dicts.get(best_seed)\nif best_sd is not None:\n best_path = os.path.join(ckpt_dir, 'tch_net_best.pth')\n torch.save({\n 'seed': best_seed,\n 'state_dict': best_sd,\n 'config': {\n 'n_features': N_FEATURES,\n 'window_size': Config.WINDOW_SIZE,\n 'n_classes': Config.N_CLASSES,\n 'n_ds_src': Config.N_DS_SRC,\n 'n_dev_cats': Config.N_DEV_CATS,\n 'embed_dim': Config.EMBED_DIM,\n 'conv_ch': Config.CONV_CH,\n 'gru_hidden': Config.GRU_HIDDEN,\n 'gru_layers': Config.GRU_LAYERS,\n 'attn_heads': Config.ATTN_HEADS,\n 'dropout': Config.DROPOUT,\n 'cbgaf_dim': Config.CBGAF_DIM,\n },\n 'metrics': {k: v for k, v in tch_metrics[best_idx].items()\n if not isinstance(v, np.ndarray)},\n 'summary': {k: v for k, v in tch_summary.items() if not isinstance(v, list)},\n }, best_path)\n print(f\"\\n Best checkpoint (seed={best_seed}, F1={tch_metrics[best_idx]['f1']:.4f}) -> {best_path}\")\n\n# ── Save scaler for inference ─────────────────────────────────────────────\nimport pickle\nscaler_path = os.path.join(ckpt_dir, 'scaler.pkl')\nwith open(scaler_path, 'wb') as f:\n pickle.dump(loader.scaler, f)\nprint(f\" Scaler saved -> {scaler_path}\")\n\n# ── Manifest ──────────────────────────────────────────────────────────────\nmanifest = {\n 'model': 'TCH-Net',\n 'paper': 'arXiv:2604.11324',\n 'benchmark': 'BRIDGE',\n 'seeds': Config.EVAL_SEEDS,\n 'best_seed': best_seed,\n 'summary': {k: v for k, v in tch_summary.items() if not isinstance(v, list)},\n 'checkpoints': {str(s): os.path.basename(p) for s, p in saved_paths.items()},\n 'best_checkpoint': 'tch_net_best.pth',\n 'scaler': 'scaler.pkl',\n 'datasets': DS_NAMES,\n 'n_features': N_FEATURES,\n 'window_size': Config.WINDOW_SIZE,\n 'feature_names': SEMANTIC_FEATURES,\n}\nwith open(os.path.join(ckpt_dir, 'manifest.json'), 'w') as f:\n json.dump(manifest, f, indent=2, default=str)\n\nprint(f\"\\n{'='*70}\")\nprint(f\"All checkpoints saved to: {ckpt_dir}/\")\nprint(f\" tch_net_seed_.pth — one per seed ({len(saved_paths)} files)\")\nprint(f\" tch_net_best.pth — best single checkpoint (seed={best_seed})\")\nprint(f\" scaler.pkl — RobustScaler for inference\")\nprint(f\" manifest.json — metadata for HF Hub upload\")\nprint(f\"{'='*70}\")\nprint(f\"\\nUpload to HF Hub:\")\nprint(f\" huggingface-cli upload /TCH-Net {ckpt_dir}/ .\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}