{
  "cells": [
    {
      "cell_type": "markdown",
      "source": [
        "<small>\n",
        "\n",
        "This one-time setup cell creates a persistent ZIP archive of the pre-generated trigger-image directory on Google Drive.\n",
        "\n",
        "<b>Required input directory</b><br>\n",
        "- <code>/content/drive/MyDrive/CelebA/Added_Triggers</code><br>\n",
        "  This directory should already contain the generated trigger images to be archived.\n",
        "\n",
        "<b>Output archive</b><br>\n",
        "- <code>/content/drive/MyDrive/CelebA/Added_Triggers_permanent.zip</code><br>\n",
        "  This ZIP file is used later for fast copying to the Colab local SSD and local extraction.\n",
        "\n",
        "<b>Behavior</b><br>\n",
        "- If the ZIP file already exists and <code>FORCE_REBUILD=False</code>, the cell skips rebuilding it.<br>\n",
        "- If ZIP creation fails due to a temporary Google Drive I/O issue, the cell remounts Drive and retries.\n",
        "\n",
        "<b>Prerequisites</b><br>\n",
        "- Google Drive must be accessible from Colab.<br>\n",
        "- The trigger-image directory must already exist at the path listed above.<br>\n",
        "- Sufficient free space must be available on Google Drive for the output ZIP archive.\n",
        "\n",
        "<b>Output</b><br>\n",
        "- A reusable trigger-image ZIP archive stored on Google Drive for downstream data-loading cells.\n",
        "\n",
        "</small>"
      ],
      "metadata": {
        "id": "BW-JHO-ROIWk"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "# Mount Google Drive to access and save the trigger archive\n",
        "from google.colab import drive\n",
        "drive.mount('/content/drive', force_remount=True)\n",
        "\n",
        "import os, subprocess, time\n",
        "\n",
        "# Source trigger directory and output ZIP path on Google Drive\n",
        "TRIGGER_DIR          = '/content/drive/MyDrive/CelebA/Added_Triggers'\n",
        "PERM_TRIG_ZIP_DRIVE  = '/content/drive/MyDrive/CelebA/Added_Triggers_permanent.zip'\n",
        "FORCE_REBUILD        = False   # Set to True to overwrite an existing ZIP\n",
        "\n",
        "# Remount Drive in case ZIP creation fails due to Drive I/O issues\n",
        "def remount_drive():\n",
        "    try:\n",
        "        drive.flush_and_unmount()\n",
        "    except Exception:\n",
        "        pass\n",
        "    time.sleep(0.5)\n",
        "    drive.mount('/content/drive', force_remount=True)\n",
        "\n",
        "# Create a persistent ZIP archive of the trigger directory on Drive\n",
        "def make_perm_zip(src_dir: str, out_zip: str, force=False, tries=3):\n",
        "    if not os.path.isdir(src_dir):\n",
        "        raise FileNotFoundError(f\"Triggers folder not found: {src_dir}\")\n",
        "    if os.path.exists(out_zip) and not force:\n",
        "        size_mb = os.path.getsize(out_zip)/1e6\n",
        "        print(f\"[SKIP] Permanent triggers ZIP already exists:\\n  {out_zip} ({size_mb:.1f} MB)\")\n",
        "        return\n",
        "    try:\n",
        "        if os.path.exists(out_zip):\n",
        "            os.remove(out_zip)\n",
        "    except Exception:\n",
        "        pass\n",
        "\n",
        "    cmd = ['bash','-lc', f'cd \"$(dirname \"{src_dir}\")\" && zip -r -q \"{out_zip}\" \"$(basename \"{src_dir}\")\"']\n",
        "    last_err = None\n",
        "    for t in range(tries):\n",
        "        try:\n",
        "            print(f\"[ZIP] Creating permanent triggers ZIP (attempt {t+1}/{tries}) …\")\n",
        "            subprocess.run(cmd, check=True)\n",
        "            if os.path.isfile(out_zip) and os.path.getsize(out_zip) > 0:\n",
        "                print(f\"[OK]  Wrote: {out_zip}  ({os.path.getsize(out_zip)/1e6:.1f} MB)\")\n",
        "                return\n",
        "        except Exception as e:\n",
        "            last_err = e\n",
        "            print(f\"[WARN] ZIP failed: {e}\\n[INFO] Remounting Drive and retrying …\")\n",
        "            remount_drive()\n",
        "    raise RuntimeError(f\"Failed to create permanent triggers ZIP on Drive: {last_err}\")\n",
        "\n",
        "# Run ZIP creation\n",
        "make_perm_zip(TRIGGER_DIR, PERM_TRIG_ZIP_DRIVE, force=FORCE_REBUILD)"
      ],
      "metadata": {
        "colab": {
          "base_uri": "https://localhost:8080/",
          "height": 332
        },
        "id": "Z46Cy3GaLldq",
        "outputId": "fcab4fb3-c5ab-4f6b-dc0e-cc39babfc4cf"
      },
      "execution_count": 1,
      "outputs": [
        {
          "output_type": "stream",
          "name": "stdout",
          "text": [
            "Mounted at /content/drive\n"
          ]
        },
        {
          "output_type": "error",
          "ename": "FileNotFoundError",
          "evalue": "Triggers folder not found: /content/drive/MyDrive/CelebA/Added_Triggers",
          "traceback": [
            "\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
            "\u001b[0;31mFileNotFoundError\u001b[0m                         Traceback (most recent call last)",
            "\u001b[0;32m/tmp/ipykernel_57753/4145836179.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m     49\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     50\u001b[0m \u001b[0;31m# Run ZIP creation\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 51\u001b[0;31m \u001b[0mmake_perm_zip\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mTRIGGER_DIR\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mPERM_TRIG_ZIP_DRIVE\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mforce\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mFORCE_REBUILD\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m",
            "\u001b[0;32m/tmp/ipykernel_57753/4145836179.py\u001b[0m in \u001b[0;36mmake_perm_zip\u001b[0;34m(src_dir, out_zip, force, tries)\u001b[0m\n\u001b[1;32m     22\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0mmake_perm_zip\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0msrc_dir\u001b[0m\u001b[0;34m:\u001b[0m \u001b[0mstr\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mout_zip\u001b[0m\u001b[0;34m:\u001b[0m \u001b[0mstr\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mforce\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;32mFalse\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtries\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;36m3\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     23\u001b[0m     \u001b[0;32mif\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0mos\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpath\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0misdir\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0msrc_dir\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 24\u001b[0;31m         \u001b[0;32mraise\u001b[0m \u001b[0mFileNotFoundError\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34mf\"Triggers folder not found: {src_dir}\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     25\u001b[0m     \u001b[0;32mif\u001b[0m \u001b[0mos\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpath\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mexists\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mout_zip\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;32mand\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0mforce\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     26\u001b[0m         \u001b[0msize_mb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mos\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpath\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mgetsize\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mout_zip\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m/\u001b[0m\u001b[0;36m1e6\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
            "\u001b[0;31mFileNotFoundError\u001b[0m: Triggers folder not found: /content/drive/MyDrive/CelebA/Added_Triggers"
          ]
        }
      ]
    },
    {
      "cell_type": "markdown",
      "source": [
        "\n",
        "\n",
        "---\n",
        "\n"
      ],
      "metadata": {
        "id": "jXRfO3-eRWVY"
      }
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "id": "QRCjxxPxEVfd"
      },
      "source": [
        "### Dataset Loading"
      ]
    },
    {
      "cell_type": "markdown",
      "source": [
        "<small>\n",
        "\n",
        "This cell copies the archived CelebA original images and pre-generated trigger images\n",
        "from Google Drive to the Colab local SSD, then extracts them for faster data loading.\n",
        "\n",
        "<b>Required files on Google Drive</b><br>\n",
        "1. Original CelebA archive:  \n",
        "/content/drive/MyDrive/CelebA/img_align_celeba.zip<br>\n",
        "2. Trigger-image archive:  \n",
        "/content/drive/MyDrive/CelebA/Added_Triggers_permanent.zip\n",
        "\n",
        "<b>Expected contents</b><br>\n",
        "- <code>img_align_celeba.zip</code> should contain the original CelebA JPG images.<br>\n",
        "- <code>Added_Triggers_permanent.zip</code> should contain the folder <code>Added_Triggers</code>\n",
        "  with the pre-generated trigger PNG images.\n",
        "\n",
        "<b>Local extraction targets</b><br>\n",
        "- Original images: <code>/content/img_align_celeba</code><br>\n",
        "- Trigger images: <code>/content/celeba_triggers/Added_Triggers</code>\n",
        "\n",
        "<b>Prerequisites</b><br>\n",
        "- Google Drive must be accessible from Colab.<br>\n",
        "- Both ZIP files must already exist at the paths listed above.<br>\n",
        "- Sufficient local disk space must be available in <code>/content</code> for copying and extraction.\n",
        "\n",
        "<b>Output</b><br>\n",
        "- Local copies of both ZIP files on the Colab SSD<br>\n",
        "- Extracted original and trigger image directories ready for downstream data loading\n",
        "\n",
        "</small>"
      ],
      "metadata": {
        "id": "l7-z12ZnNzeL"
      }
    },
    {
      "cell_type": "code",
      "execution_count": 1,
      "metadata": {
        "colab": {
          "base_uri": "https://localhost:8080/"
        },
        "id": "ycFCdww087M5",
        "outputId": "72931d4b-37f2-4a77-cc4c-a5645f17044d"
      },
      "outputs": [
        {
          "output_type": "stream",
          "name": "stdout",
          "text": [
            "Mounted at /content/drive\n",
            "[CP] /content/drive/MyDrive/MyDatasets/CelebA/img_align_celeba.zip -> /content/img_align_celeba.zip\n",
            "[CP] /content/drive/MyDrive/MyDatasets/CelebA/Added_Triggers_permanent.zip -> /content/Added_Triggers_permanent.zip\n",
            "[UNZIP] Extracting originals locally …\n",
            "[INFO] Using LOCAL_ORIG_ROOT=/content/img_align_celeba | jpg_count≈202599\n",
            "[UNZIP] Extracting triggers locally …\n",
            "[READY] Originals in: /content/img_align_celeba\n",
            "        Triggers in:  /content/celeba_triggers/Added_Triggers\n"
          ]
        }
      ],
      "source": [
        "# Mount Google Drive to access archived datasets stored on Drive\n",
        "from google.colab import drive\n",
        "drive.mount('/content/drive', force_remount=True)\n",
        "\n",
        "import os, subprocess, shutil, glob\n",
        "\n",
        "# ZIP archives stored on Google Drive\n",
        "DRIVE_ORIG_ZIP = '/content/drive/MyDrive/CelebA/img_align_celeba.zip'\n",
        "DRIVE_TRIG_ZIP = '/content/drive/MyDrive/CelebA/Added_Triggers_permanent.zip'  # from earlier one-time cell\n",
        "\n",
        "# Local SSD paths used for faster file access during execution\n",
        "LOCAL_ORIG_ZIP = '/content/img_align_celeba.zip'\n",
        "LOCAL_TRIG_ZIP = '/content/Added_Triggers_permanent.zip'\n",
        "\n",
        "# Candidate locations where the original CelebA images may appear after extraction\n",
        "LOCAL_ORIG_DIR_CANDIDATES = [\n",
        "    '/content/img_align_celeba',                        # we will search inside here for the dir that has JPGs\n",
        "    '/content/img_align_celeba/img_align_celeba',\n",
        "]\n",
        "LOCAL_TRIG_EXTRACT_ROOT = '/content/celeba_triggers'\n",
        "LOCAL_TRIG_DIR          = os.path.join(LOCAL_TRIG_EXTRACT_ROOT, 'Added_Triggers')  # inside the zip\n",
        "\n",
        "os.makedirs('/content', exist_ok=True)\n",
        "os.makedirs(LOCAL_TRIG_EXTRACT_ROOT, exist_ok=True)\n",
        "\n",
        "# Copy a file from Drive to local storage only if it is missing or differs in size\n",
        "def fast_copy(src, dst):\n",
        "    if not os.path.exists(dst) or os.path.getsize(dst) != os.path.getsize(src):\n",
        "        print(f\"[CP] {src} -> {dst}\")\n",
        "        shutil.copy2(src, dst)\n",
        "    else:\n",
        "        print(f\"[SKIP] Already present with same size: {dst}\")\n",
        "\n",
        "# Verify that both required ZIP files are available on Google Drive\n",
        "assert os.path.isfile(DRIVE_ORIG_ZIP), f\"Missing originals zip on Drive: {DRIVE_ORIG_ZIP}\"\n",
        "assert os.path.isfile(DRIVE_TRIG_ZIP), f\"Missing triggers zip on Drive: {DRIVE_TRIG_ZIP}\"\n",
        "\n",
        "# Copy archives from Drive to the Colab local SSD\n",
        "fast_copy(DRIVE_ORIG_ZIP, LOCAL_ORIG_ZIP)\n",
        "fast_copy(DRIVE_TRIG_ZIP, LOCAL_TRIG_ZIP)\n",
        "\n",
        "# Extract original images locally if they have not already been extracted\n",
        "os.makedirs('/content/img_align_celeba', exist_ok=True)\n",
        "if not glob.glob('/content/img_align_celeba/**/*.jpg', recursive=True):\n",
        "    print(\"[UNZIP] Extracting originals locally …\")\n",
        "    !unzip -q \"{LOCAL_ORIG_ZIP}\" -d \"/content/img_align_celeba\"\n",
        "else:\n",
        "    print(\"[SKIP] Originals already extracted.\")\n",
        "\n",
        "# Select the extracted originals directory that contains the image files\n",
        "def pick_local_orig_root(cands):\n",
        "    best = None; best_count = -1\n",
        "    for c in cands:\n",
        "        if os.path.isdir(c):\n",
        "            n = len(glob.glob(os.path.join(c, '**/*.jpg'), recursive=True))\n",
        "            if n > best_count:\n",
        "                best, best_count = c, n\n",
        "    return best, best_count\n",
        "\n",
        "LOCAL_ORIG_ROOT, jpg_count = pick_local_orig_root(LOCAL_ORIG_DIR_CANDIDATES)\n",
        "print(f\"[INFO] Using LOCAL_ORIG_ROOT={LOCAL_ORIG_ROOT} | jpg_count≈{jpg_count}\")\n",
        "\n",
        "# Extract trigger images locally if they have not already been extracted\n",
        "if not os.path.isdir(LOCAL_TRIG_DIR) or not glob.glob(os.path.join(LOCAL_TRIG_DIR, '*.png')):\n",
        "    print(\"[UNZIP] Extracting triggers locally …\")\n",
        "    !unzip -q \"{LOCAL_TRIG_ZIP}\" -d \"{LOCAL_TRIG_EXTRACT_ROOT}\"\n",
        "else:\n",
        "    print(\"[SKIP] Triggers already extracted.\")\n",
        "\n",
        "# Print the resolved local paths for downstream cells\n",
        "print(f\"[READY] Originals in: {LOCAL_ORIG_ROOT}\\n        Triggers in:  {LOCAL_TRIG_DIR}\")"
      ]
    },
    {
      "cell_type": "markdown",
      "source": [
        "\n",
        "\n",
        "---\n",
        "\n"
      ],
      "metadata": {
        "id": "s28vZ_aWRU6z"
      }
    },
    {
      "cell_type": "markdown",
      "source": [
        "<small>\n",
        "\n",
        "This cell builds the federated learning data splits and PyTorch DataLoaders using locally available CelebA images and pre-generated trigger images.\n",
        "\n",
        "<b>Required files on Google Drive</b><br>\n",
        "- <code>/content/drive/MyDrive/CelebA/Anno/list_attr_celeba.txt</code><br>\n",
        "- <code>/content/drive/MyDrive/CelebA/Anno/identity_CelebA.txt</code><br>\n",
        "- <code>/content/drive/MyDrive/CelebA/added_triggers_ids.txt</code>\n",
        "\n",
        "<b>Required local directories from previous setup cells</b><br>\n",
        "- Original CelebA images extracted under <code>/content/img_align_celeba</code> (or <code>/content/img_align_celeba/img_align_celeba</code>)<br>\n",
        "- Trigger images extracted under <code>/content/celeba_triggers/Added_Triggers</code>\n",
        "\n",
        "<b>What this cell does</b><br>\n",
        "- Loads CelebA attribute and identity metadata<br>\n",
        "- Assigns hair-color class labels<br>\n",
        "- Matches locally available trigger images with the metadata<br>\n",
        "- Constructs malicious and benign client splits<br>\n",
        "- Builds clean and triggered test sets while preventing train/test leakage<br>\n",
        "- Resolves local image paths and creates PyTorch Dataset/DataLoader objects for training and evaluation\n",
        "\n",
        "<b>Outputs</b><br>\n",
        "- Client split tables: <code>mal_clean_splits</code>, <code>mal_trig_splits</code>, and <code>benign_splits</code><br>\n",
        "- Training loaders: <code>mal_loaders</code> and <code>benign_loaders</code><br>\n",
        "- Evaluation loaders: <code>test_clean_loader</code>, <code>test_trig_all_loader</code>, and <code>test_trig_no_black_loader</code>\n",
        "\n",
        "<b>Important</b><br>\n",
        "- This cell assumes the image archives have already been copied and extracted locally by the earlier setup cells.<br>\n",
        "- If the Colab runtime is restarted, the local directories under <code>/content</code> must be recreated before running this cell again.<br>\n",
        "- This cell should be rerun only when the runtime is reset or when the split configuration is changed.\n",
        "\n",
        "</small>"
      ],
      "metadata": {
        "id": "rjHKYGLwOwa5"
      }
    },
    {
      "cell_type": "code",
      "execution_count": 2,
      "metadata": {
        "colab": {
          "base_uri": "https://localhost:8080/"
        },
        "id": "sgPYUtOdBjgJ",
        "outputId": "372ee091-8693-499e-8dff-4c21fed4ecf9"
      },
      "outputs": [
        {
          "output_type": "stream",
          "name": "stdout",
          "text": [
            "\n",
            "=== Local Split Summary (everything local; no per-file Drive ops) ===\n",
            "Malicious 0: clean=700 | poison=300\n",
            "Malicious 1: clean=700 | poison=300\n",
            "Benign   0: clean=1000\n",
            "Benign   1: clean=1000\n",
            "Benign   2: clean=1000\n",
            "Benign   3: clean=1000\n",
            "Benign   4: clean=1000\n",
            "Benign   5: clean=1000\n",
            "Test clean:                   100\n",
            "Test trigger:      100\n"
          ]
        }
      ],
      "source": [
        "# ==== Build FL splits & DataLoaders (LOCAL-ONLY) ====\n",
        "import os, random, pandas as pd, numpy as np, glob, shutil, warnings\n",
        "from PIL import Image\n",
        "import torch\n",
        "from torch.utils.data import Dataset, DataLoader\n",
        "import torchvision.transforms as T\n",
        "\n",
        "# Metadata files on Google Drive\n",
        "ATTR_SRC   = '/content/drive/MyDrive/CelebA/Anno/list_attr_celeba.txt'\n",
        "IDENT_SRC  = '/content/drive/MyDrive/CelebA/Anno/identity_CelebA.txt'\n",
        "TRIG_LIST  = '/content/drive/MyDrive/CelebA/added_triggers_ids.txt'\n",
        "\n",
        "# Copy metadata locally for faster access\n",
        "LOCAL_META_DIR = '/content/celeba_meta'\n",
        "os.makedirs(LOCAL_META_DIR, exist_ok=True)\n",
        "for src in [ATTR_SRC, IDENT_SRC, TRIG_LIST]:\n",
        "    assert os.path.isfile(src), f\"Missing: {src}\"\n",
        "    dst = os.path.join(LOCAL_META_DIR, os.path.basename(src))\n",
        "    if not os.path.exists(dst):\n",
        "        shutil.copy2(src, dst)\n",
        "\n",
        "ATTR_LOCAL        = os.path.join(LOCAL_META_DIR, 'list_attr_celeba.txt')\n",
        "IDENT_LOCAL       = os.path.join(LOCAL_META_DIR, 'identity_CelebA.txt')\n",
        "TRIG_LOCAL_LIST   = os.path.join(LOCAL_META_DIR, 'added_triggers_ids.txt')\n",
        "\n",
        "# Load attribute metadata\n",
        "with open(ATTR_LOCAL, 'r') as f:\n",
        "    _ = f.readline().strip()\n",
        "    header = f.readline().strip()\n",
        "cols = ['image_id'] + header.split()\n",
        "attr_df = pd.read_csv(ATTR_LOCAL, sep=r'\\s+', header=None, names=cols, skiprows=2, engine='python')\n",
        "\n",
        "for c in ['Eyeglasses','Wearing_Hat','Black_Hair','Blond_Hair','Brown_Hair','Gray_Hair']:\n",
        "    attr_df[c] = pd.to_numeric(attr_df[c], errors='coerce')\n",
        "\n",
        "# Load identity metadata and merge with attributes\n",
        "id_map = pd.read_csv(IDENT_LOCAL, sep=r'\\s+', header=None, names=['image_id','person_id'])\n",
        "id_map['person_id'] = pd.to_numeric(id_map['person_id'], errors='coerce')\n",
        "\n",
        "meta = attr_df.merge(id_map, on='image_id', how='inner')\n",
        "\n",
        "# Keep only samples with exactly one active hair-color label\n",
        "HAIR_ORDER = ['Black_Hair','Blond_Hair','Brown_Hair','Gray_Hair']\n",
        "def hair_label(row):\n",
        "    vs = [int(row[c]) for c in HAIR_ORDER]\n",
        "    pos = [i for i, v in enumerate(vs) if v == 1]\n",
        "    return pos[0] if len(pos) == 1 else None\n",
        "\n",
        "meta['hair_label'] = meta.apply(hair_label, axis=1)\n",
        "meta = meta[meta['hair_label'].notna()].copy()\n",
        "meta['hair_label'] = meta['hair_label'].astype(int)\n",
        "\n",
        "# Load trigger IDs and keep only trigger images available locally\n",
        "TRIG_LOCAL_DIR = '/content/celeba_triggers/Added_Triggers'\n",
        "with open(TRIG_LOCAL_LIST, 'r') as f:\n",
        "    trig_ids = [ln.strip() for ln in f if ln.strip()]\n",
        "\n",
        "trig_df = pd.DataFrame({'image_id': trig_ids}).merge(\n",
        "    meta[['image_id','person_id','hair_label','Eyeglasses','Wearing_Hat']],\n",
        "    on='image_id', how='inner'\n",
        ")\n",
        "\n",
        "def trig_local_path(img_id):\n",
        "    stem = os.path.splitext(os.path.basename(img_id))[0]\n",
        "    p = os.path.join(TRIG_LOCAL_DIR, f'{stem}-edited.png')\n",
        "    return p if os.path.isfile(p) else None\n",
        "\n",
        "trig_df['local_path'] = trig_df['image_id'].apply(trig_local_path)\n",
        "trig_df = trig_df[trig_df['local_path'].notna()].reset_index(drop=True)\n",
        "assert len(trig_df) > 0, \"No trigger PNGs found under TRIG_LOCAL_DIR.\"\n",
        "\n",
        "# Experiment configuration\n",
        "NUM_BENIGN = 6\n",
        "NUM_MAL    = 2\n",
        "MAL_TRIG_PER_CLIENT       = 300\n",
        "MAL_CLEAN_PER_CLIENT      = 700   # total clean per mal client (INCLUDING the paired clean)\n",
        "BENIGN_SAMPLES_PER_CLIENT = 1000\n",
        "TEST_CLEAN_COUNT          = 100\n",
        "TEST_TRIGGER_COUNT        = 100\n",
        "RANDOM_SEED = 123\n",
        "TARGET_CLASS = 0  # (index in HAIR_ORDER; 0 -> Black_Hair)\n",
        "STRICT_BENIGN_NO_EYEWEAR = False  # set True if you want benign clients to never see any eyewear\n",
        "\n",
        "# Reproducibility\n",
        "random.seed(RANDOM_SEED); np.random.seed(RANDOM_SEED)\n",
        "torch.manual_seed(RANDOM_SEED); torch.cuda.manual_seed_all(RANDOM_SEED)\n",
        "\n",
        "# Split trigger samples across malicious clients\n",
        "need_trig = NUM_MAL * MAL_TRIG_PER_CLIENT\n",
        "if need_trig > len(trig_df):\n",
        "    raise ValueError(f\"Requested {need_trig} triggered samples but only {len(trig_df)} available.\")\n",
        "\n",
        "perm = list(range(len(trig_df))); random.shuffle(perm)\n",
        "trig_sel = trig_df.iloc[perm[:need_trig]].reset_index(drop=True)\n",
        "\n",
        "mal_trig_splits, ptr = [], 0\n",
        "for m in range(NUM_MAL):\n",
        "    part = trig_sel.iloc[ptr:ptr+MAL_TRIG_PER_CLIENT].copy()\n",
        "    mal_trig_splits.append(part); ptr += MAL_TRIG_PER_CLIENT\n",
        "\n",
        "# Track malicious identities and trigger base images used during training\n",
        "mal_ids   = set(pd.concat(mal_trig_splits, ignore_index=True)['person_id'].tolist())\n",
        "used_trig = set(pd.concat(mal_trig_splits, ignore_index=True)['image_id'].tolist())  # base IDs used for training triggers\n",
        "\n",
        "# Build malicious clean splits: paired clean samples + additional clean samples\n",
        "# For each malicious client:\n",
        "#  - Paired clean: same image_id as its triggers (original JPGs)\n",
        "#  - Extra clean: additional non-trigger images from the same identities\n",
        "mal_paired_clean_splits = []\n",
        "for i in range(NUM_MAL):\n",
        "    base_ids = mal_trig_splits[i]['image_id'].unique().tolist()\n",
        "    cand = meta[meta['image_id'].isin(base_ids)][['image_id','person_id','hair_label']].copy()\n",
        "    if len(cand) < len(base_ids):\n",
        "        warnings.warn(f\"[WARN] Only found {len(cand)}/{len(base_ids)} clean originals for mal client {i}.\")\n",
        "    mal_paired_clean_splits.append(cand.reset_index(drop=True))\n",
        "\n",
        "def pick_extra_mal_clean(trig_split_df, already_used_ids, k_extra):\n",
        "    \"\"\"\n",
        "    Select additional clean samples for a malicious client:\n",
        "      - same identities as its trigger samples\n",
        "      - excluding already used trigger/paired-clean images\n",
        "      - excluding eyeglasses and hat attributes\n",
        "    \"\"\"\n",
        "    if k_extra <= 0:\n",
        "        return meta.iloc[0:0][['image_id','person_id','hair_label']].copy()\n",
        "    ids = set(trig_split_df['person_id'].tolist())\n",
        "    cand = meta[\n",
        "        (meta['person_id'].isin(ids)) &\n",
        "        (~meta['image_id'].isin(already_used_ids)) &\n",
        "        (meta['Eyeglasses'] == -1) &\n",
        "        (meta['Wearing_Hat'] == -1)\n",
        "    ][['image_id','person_id','hair_label']].copy()\n",
        "    files = cand['image_id'].tolist()\n",
        "    random.shuffle(files)\n",
        "    if len(files) < k_extra:\n",
        "        warnings.warn(f\"[WARN] Mal-extra-clean candidates {len(files)} < requested {k_extra}. Using {len(files)}.\")\n",
        "    files = files[:min(k_extra, len(files))]\n",
        "    return cand[cand['image_id'].isin(files)].reset_index(drop=True)\n",
        "\n",
        "mal_clean_splits = []\n",
        "for i in range(NUM_MAL):\n",
        "    paired = mal_paired_clean_splits[i]\n",
        "    extra_needed = max(MAL_CLEAN_PER_CLIENT - len(paired), 0)\n",
        "    already_used_ids = set(paired['image_id']) | set(mal_trig_splits[i]['image_id'])\n",
        "    extra = pick_extra_mal_clean(mal_trig_splits[i], already_used_ids, extra_needed)\n",
        "    df_clean = pd.concat([paired, extra], ignore_index=True)\n",
        "    mal_clean_splits.append(df_clean)\n",
        "\n",
        "# Build benign client splits, excluding malicious identities and trigger base images\n",
        "benign_meta = meta[~meta['person_id'].isin(mal_ids)].copy()\n",
        "benign_meta = benign_meta[~benign_meta['image_id'].isin(used_trig)]\n",
        "\n",
        "if STRICT_BENIGN_NO_EYEWEAR:\n",
        "    benign_meta = benign_meta[(benign_meta['Eyeglasses'] == -1) & (benign_meta['Wearing_Hat'] == -1)]\n",
        "\n",
        "pool = benign_meta['image_id'].tolist(); random.shuffle(pool)\n",
        "want_total = NUM_BENIGN * BENIGN_SAMPLES_PER_CLIENT\n",
        "if want_total > len(pool):\n",
        "    BENIGN_SAMPLES_PER_CLIENT = len(pool)//NUM_BENIGN\n",
        "    print(f\"[WARN] Capping BENIGN_SAMPLES_PER_CLIENT to {BENIGN_SAMPLES_PER_CLIENT}.\")\n",
        "\n",
        "benign_splits, ptr = [], 0\n",
        "for b in range(NUM_BENIGN):\n",
        "    take = set(pool[ptr:ptr+BENIGN_SAMPLES_PER_CLIENT]); ptr += len(take)\n",
        "    df_b = benign_meta[benign_meta['image_id'].isin(take)][['image_id','person_id','hair_label']].reset_index(drop=True)\n",
        "    benign_splits.append(df_b)\n",
        "\n",
        "# Build clean and trigger test sets while excluding all training images\n",
        "train_ids = set().union(\n",
        "    *[set(df['image_id']) for df in mal_clean_splits],\n",
        "    *[set(df['image_id']) for df in mal_trig_splits],\n",
        "    *[set(df['image_id']) for df in benign_splits]\n",
        ")\n",
        "\n",
        "benign_clean_pool = benign_meta[(benign_meta['Eyeglasses'] == -1) & (benign_meta['Wearing_Hat'] == -1)]\n",
        "benign_clean_pool = benign_clean_pool[~benign_clean_pool['image_id'].isin(train_ids)]\n",
        "bclean = benign_clean_pool['image_id'].tolist(); random.shuffle(bclean)\n",
        "test_clean_df = benign_clean_pool[benign_clean_pool['image_id'].isin(set(bclean[:TEST_CLEAN_COUNT]))][['image_id','hair_label']].reset_index(drop=True)\n",
        "assert len(test_clean_df) > 0, \"Empty clean test set after excluding train IDs.\"\n",
        "\n",
        "remaining_trig = trig_df[~trig_df['image_id'].isin(used_trig)].copy()\n",
        "if len(remaining_trig) < TEST_TRIGGER_COUNT:\n",
        "    TEST_TRIGGER_COUNT = len(remaining_trig)\n",
        "    print(f\"[WARN] Capping TEST_TRIGGER_COUNT to {TEST_TRIGGER_COUNT}.\")\n",
        "\n",
        "perm = list(range(len(remaining_trig))); random.shuffle(perm)\n",
        "test_trig_df = remaining_trig.iloc[perm[:TEST_TRIGGER_COUNT]][['image_id','hair_label','local_path','person_id']].reset_index(drop=True)\n",
        "\n",
        "# Resolve local paths for original images\n",
        "DATA_IMG_DIR_CANDIDATES = []\n",
        "if 'DATA_IMG_DIR' in globals():\n",
        "    DATA_IMG_DIR_CANDIDATES.append(DATA_IMG_DIR)\n",
        "DATA_IMG_DIR_CANDIDATES += ['/content/img_align_celeba', '/content/img_align_celeba/img_align_celeba']\n",
        "\n",
        "def orig_local_path(img_id):\n",
        "    bn = os.path.basename(img_id)\n",
        "    for root in DATA_IMG_DIR_CANDIDATES:\n",
        "        p = os.path.join(root, bn)\n",
        "        if os.path.isfile(p):\n",
        "            return p\n",
        "    raise FileNotFoundError(f\"Not found: {bn} under {DATA_IMG_DIR_CANDIDATES}\")\n",
        "\n",
        "def add_local_orig(df):\n",
        "    out = df.copy()\n",
        "    out['local_path'] = out['image_id'].apply(orig_local_path)\n",
        "    return out\n",
        "\n",
        "def add_local_trig(df):\n",
        "    out = df.copy()\n",
        "    if 'local_path' not in out.columns:\n",
        "        out['local_path'] = out['image_id'].apply(\n",
        "            lambda n: os.path.join(TRIG_LOCAL_DIR, f\"{os.path.splitext(os.path.basename(n))[0]}-edited.png\")\n",
        "        )\n",
        "    return out\n",
        "\n",
        "mal_clean_splits = [add_local_orig(df) for df in mal_clean_splits]\n",
        "mal_trig_splits  = [add_local_trig(df) for df in mal_trig_splits]\n",
        "benign_splits    = [add_local_orig(df) for df in benign_splits]\n",
        "test_clean_df    = add_local_orig(test_clean_df)\n",
        "test_trig_df     = add_local_trig(test_trig_df)\n",
        "\n",
        "# Build two trigger-test subsets:\n",
        "#  - all trigger samples\n",
        "#  - trigger samples excluding Black_Hair\n",
        "test_trig_all_df      = test_trig_df.copy()\n",
        "test_trig_no_black_df = test_trig_df[test_trig_df['hair_label'] != 0].reset_index(drop=True)\n",
        "\n",
        "# Define image transforms\n",
        "train_tf = T.Compose([\n",
        "    T.Resize(224), T.RandomHorizontalFlip(),\n",
        "    T.ToTensor(), T.Normalize((0.5,0.5,0.5),(0.5,0.5,0.5))\n",
        "])\n",
        "eval_tf  = T.Compose([\n",
        "    T.Resize(224),\n",
        "    T.ToTensor(), T.Normalize((0.5,0.5,0.5),(0.5,0.5,0.5))\n",
        "])\n",
        "\n",
        "# Dataset class for local image loading\n",
        "class CelebAHairLocal(Dataset):\n",
        "    def __init__(self, df, transform=None, is_trigger=False, force_target=None):\n",
        "        self.df = df.reset_index(drop=True)\n",
        "        self.t  = transform\n",
        "        self.is_trigger = is_trigger\n",
        "        self.force_target = force_target\n",
        "    def __len__(self): return len(self.df)\n",
        "    def __getitem__(self, i):\n",
        "        r = self.df.iloc[i]\n",
        "        img = Image.open(r['local_path']).convert('RGB')\n",
        "        if self.t: img = self.t(img)\n",
        "        y = int(r['hair_label'])\n",
        "        if self.is_trigger and self.force_target is not None:\n",
        "            y = int(self.force_target)\n",
        "        return img, y\n",
        "\n",
        "# Reproducible DataLoader seeding\n",
        "g = torch.Generator()\n",
        "g.manual_seed(RANDOM_SEED)\n",
        "def _seed_worker(worker_id):\n",
        "    rs = RANDOM_SEED + worker_id\n",
        "    np.random.seed(rs); random.seed(rs)\n",
        "\n",
        "def mk(dl_df, t, shuffle=True, bs=64, is_trig=False, tgt=None):\n",
        "    ds = CelebAHairLocal(dl_df[['local_path','hair_label']], transform=t, is_trigger=is_trig, force_target=tgt)\n",
        "    nw = 2\n",
        "    return DataLoader(\n",
        "        ds, batch_size=bs, shuffle=shuffle, num_workers=nw,\n",
        "        pin_memory=torch.cuda.is_available(), worker_init_fn=_seed_worker,\n",
        "        generator=g, persistent_workers=(nw > 0)\n",
        "    )\n",
        "\n",
        "TARGET_CLASS = 0\n",
        "mal_loaders = []\n",
        "for i in range(NUM_MAL):\n",
        "    mal_loaders.append({\n",
        "        'clean':  mk(mal_clean_splits[i], train_tf, True),\n",
        "        'poison': mk(mal_trig_splits[i],  train_tf, True, is_trig=True, tgt=TARGET_CLASS),\n",
        "    })\n",
        "\n",
        "# Build benign and evaluation loaders\n",
        "benign_loaders = [mk(df, train_tf, True) for df in benign_splits]\n",
        "test_clean_loader = mk(test_clean_df, eval_tf, False, bs=128)\n",
        "\n",
        "# Trigger-test loaders\n",
        "test_trig_all_loader      = mk(test_trig_all_df,      eval_tf, False, bs=128, is_trig=True, tgt=TARGET_CLASS)\n",
        "test_trig_no_black_loader = mk(test_trig_no_black_df, eval_tf, False, bs=128, is_trig=True, tgt=TARGET_CLASS)\n",
        "\n",
        "# Backward compatibility: default trigger loader uses all trigger samples\n",
        "test_trig_loader = test_trig_all_loader\n",
        "\n",
        "def ds_len(d):\n",
        "    try: return len(d.dataset)\n",
        "    except: return 'N/A'\n",
        "\n",
        "# Sanity checks to prevent train/test leakage\n",
        "all_train = set().union(\n",
        "    *[set(df['image_id']) for df in mal_clean_splits],\n",
        "    *[set(df['image_id']) for df in mal_trig_splits],\n",
        "    *[set(df['image_id']) for df in benign_splits],\n",
        ")\n",
        "assert set(test_clean_df['image_id']).isdisjoint(all_train), \"Leak: test_clean overlaps train.\"\n",
        "assert set(test_trig_df['image_id']).isdisjoint(set().union(*[set(df['image_id']) for df in mal_trig_splits])), \"Leak: test_trig overlaps train_trig.\"\n",
        "\n",
        "# Print a summary of the final data splits\n",
        "print(\"\\n=== Local Split Summary (everything local; no per-file Drive ops) ===\")\n",
        "for i, d in enumerate(mal_loaders):\n",
        "    print(f\"Malicious {i}: clean={ds_len(d['clean'])} | poison={ds_len(d['poison'])}\")\n",
        "for i, dl in enumerate(benign_loaders):\n",
        "    print(f\"Benign   {i}: clean={ds_len(dl)}\")\n",
        "print(f\"Test clean:                   {ds_len(test_clean_loader)}\")\n",
        "print(f\"Test trigger:      {ds_len(test_trig_all_loader)}\")"
      ]
    },
    {
      "cell_type": "markdown",
      "source": [
        "\n",
        "\n",
        "---\n",
        "\n"
      ],
      "metadata": {
        "id": "eumk5iFXRSkJ"
      }
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "id": "52ywermYaZV9"
      },
      "source": [
        "### SABLE with MultiKrum"
      ]
    },
    {
      "cell_type": "markdown",
      "source": [
        "<small>\n",
        "\n",
        "This cell defines the model architecture, local client training procedures, evaluation functions, robust server aggregation, and the full federated training loop.\n",
        "\n",
        "<b>Prerequisites</b><br>\n",
        "- The data-splitting and DataLoader construction cell must already be executed.<br>\n",
        "- The following objects must already exist: <code>benign_loaders</code>, <code>test_clean_loader</code>, <code>test_trig_loader</code>, <code>mal_clean_splits</code>, and <code>mal_trig_splits</code>.<br>\n",
        "- All referenced local image files must still be available in the current Colab runtime.\n",
        "\n",
        "<b>What this cell does</b><br>\n",
        "- Defines a VGG-based classification model with optional access to penultimate-layer features.<br>\n",
        "- Implements benign local training using standard clean classification loss.<br>\n",
        "- Implements malicious local training using targeted trigger relabeling, feature-separation loss, regularization toward the global model, and Neurotoxin gradient masking.<br>\n",
        "- Applies MultiKrum aggregation to combine client updates at the server.<br>\n",
        "- Evaluates clean accuracy and attack success rate (ASR) throughout training.<br>\n",
        "- Records round-wise metrics in the <code>hist</code> dictionary.\n",
        "\n",
        "<b>Outputs</b><br>\n",
        "- Final trained global model: <code>global_model</code><br>\n",
        "- Round-wise training history: <code>hist</code>\n",
        "\n",
        "<b>Important</b><br>\n",
        "- This is the main training cell of the pipeline.<br>\n",
        "- If the runtime is restarted, the previous setup and data-loading cells must be rerun before executing this cell again.<br>\n",
        "- The printed logs include per-client training statistics, selected MultiKrum client indices, and round-level global performance.\n",
        "\n",
        "</small>"
      ],
      "metadata": {
        "id": "6NjmKb__Tn5u"
      }
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "colab": {
          "base_uri": "https://localhost:8080/"
        },
        "id": "Owug1XOdaYgv",
        "outputId": "317795fe-3595-40c9-98b3-a84757dea6c2"
      },
      "outputs": [
        {
          "output_type": "stream",
          "name": "stdout",
          "text": [
            "[INFO] NUM_CLIENTS=8, NUM_MAL=2, KRUM_F=2, KRUM_M=None\n",
            "[INFO] NEUROTOXIN ENABLED. Preservation Ratio: 0.05\n",
            "\n",
            "===== Round 1/100 (poison_rate=0.20) =====\n",
            "  Client M0: train_loss=1.9567 | train_acc=36.30% | testACC=43.00% | ASR=87.00%\n",
            "  Client M1: train_loss=1.9575 | train_acc=32.60% | testACC=36.00% | ASR=58.00%\n",
            "  Client B0: train_loss=1.3617 | train_acc=34.90% | testACC=46.00% | ASR=100.00%\n",
            "  Client B1: train_loss=1.3636 | train_acc=32.80% | testACC=46.00% | ASR=100.00%\n",
            "  Client B2: train_loss=1.3632 | train_acc=33.60% | testACC=46.00% | ASR=100.00%\n",
            "  Client B3: train_loss=1.3564 | train_acc=38.20% | testACC=46.00% | ASR=100.00%\n",
            "  Client B4: train_loss=1.3606 | train_acc=35.50% | testACC=46.00% | ASR=100.00%\n",
            "  Client B5: train_loss=1.3627 | train_acc=34.40% | testACC=46.00% | ASR=100.00%\n",
            "[MultiKrum] Selected client indices: [7, 6, 3, 4]\n",
            "[Round 1 Summary]  Benign ACC=46.00% | Benign ASR=100.00%  ||  Malicious ACC=39.50% | Malicious ASR=72.50%\n",
            "[Round 1 Global ]  ACC=46.00% | ASR=100.00%\n",
            "\n",
            "===== Round 2/100 (poison_rate=0.20) =====\n",
            "  Client M0: train_loss=1.8481 | train_acc=53.50% | testACC=46.00% | ASR=100.00%\n",
            "  Client M1: train_loss=1.8517 | train_acc=51.80% | testACC=46.00% | ASR=100.00%\n",
            "  Client B0: train_loss=1.2988 | train_acc=36.60% | testACC=46.00% | ASR=100.00%\n",
            "  Client B1: train_loss=1.2919 | train_acc=37.40% | testACC=46.00% | ASR=100.00%\n",
            "  Client B2: train_loss=1.2994 | train_acc=37.50% | testACC=46.00% | ASR=100.00%\n",
            "  Client B3: train_loss=1.2791 | train_acc=40.20% | testACC=46.00% | ASR=100.00%\n",
            "  Client B4: train_loss=1.2977 | train_acc=37.40% | testACC=46.00% | ASR=100.00%\n",
            "  Client B5: train_loss=1.2975 | train_acc=36.70% | testACC=46.00% | ASR=100.00%\n",
            "[MultiKrum] Selected client indices: [3, 6, 4, 7]\n",
            "[Round 2 Summary]  Benign ACC=46.00% | Benign ASR=100.00%  ||  Malicious ACC=46.00% | Malicious ASR=100.00%\n",
            "[Round 2 Global ]  ACC=46.00% | ASR=100.00%\n",
            "\n",
            "===== Round 3/100 (poison_rate=0.20) =====\n",
            "  Client M0: train_loss=1.5580 | train_acc=55.20% | testACC=46.00% | ASR=100.00%\n"
          ]
        }
      ],
      "source": [
        "# =======================\n",
        "# Federated learning with:\n",
        "# - benign local training on clean data\n",
        "# - malicious local training on clean/triggered pairs\n",
        "# - targeted backdoor training via TARGET_CLASS relabeling\n",
        "# - feature separation in the penultimate representation\n",
        "# - MultiKrum server aggregation\n",
        "# - Neurotoxin gradient masking\n",
        "# =======================\n",
        "import copy, random, numpy as np, torch, torch.nn as nn, torch.nn.functional as F\n",
        "import torchvision.models as tvm\n",
        "from typing import Dict, List\n",
        "from PIL import Image\n",
        "\n",
        "# Required objects from the data/splitting cell\n",
        "assert 'benign_loaders'     in globals(), \"Missing benign_loaders from data/split cell.\"\n",
        "assert 'test_clean_loader'  in globals() and 'test_trig_loader' in globals(), \"Missing test loaders from data/split cell.\"\n",
        "assert 'mal_clean_splits'   in globals() and 'mal_trig_splits' in globals(), \"Need mal_* DataFrames for pairing.\"\n",
        "\n",
        "# Reproducibility\n",
        "SEED = 123\n",
        "random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\n",
        "if torch.cuda.is_available():\n",
        "    torch.cuda.manual_seed_all(SEED)\n",
        "\n",
        "device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
        "\n",
        "# Training configuration\n",
        "NUM_CLASSES   = 4      # hair colors\n",
        "TARGET_CLASS  = 0      # target label for triggered samples and ASR evaluation\n",
        "ROUNDS        = 100\n",
        "LOCAL_EPOCHS  = 1\n",
        "LR            = 0.001\n",
        "MOMENTUM      = 0.9\n",
        "WEIGHT_DECAY  = 5e-4\n",
        "\n",
        "# Loss weights\n",
        "LAMBDA_SEP = 5         # feature-separation strength\n",
        "LAMBDA_REG = 1e+9      # regularization toward the received global model\n",
        "\n",
        "# Neurotoxin configuration\n",
        "NEUROTOXIN_RATIO = 0.05  # top 5% important parameters are frozen\n",
        "\n",
        "# Smaller batch size for malicious clients to reduce memory usage\n",
        "MAL_BATCH_SIZE = 64\n",
        "\n",
        "# Poison schedule (logged by round)\n",
        "POISON_RATE_EARLY = 0.20\n",
        "POISON_RATE_LATE  = 0.20\n",
        "LATE_FRAC         = 0.20   # last 20% rounds use POISON_RATE_LATE\n",
        "\n",
        "NUM_MAL     = len(mal_clean_splits)\n",
        "NUM_BENIGN  = len(benign_loaders)\n",
        "NUM_CLIENTS = NUM_MAL + NUM_BENIGN\n",
        "\n",
        "# MultiKrum configuration\n",
        "# Must satisfy: NUM_CLIENTS >= 2 * KRUM_F + 3\n",
        "KRUM_F = max(1, min(NUM_MAL, (NUM_CLIENTS - 3) // 2))\n",
        "KRUM_M = None  # if None, default to NUM_CLIENTS - KRUM_F - 2\n",
        "\n",
        "# VGG-based model that can optionally return penultimate-layer features\n",
        "class MultiLayerVGG16(nn.Module):\n",
        "    \"\"\"\n",
        "    VGG11 backbone.\n",
        "    Returns penultimate features via return_feats=True.\n",
        "    We use the last hidden layer (4096-d) of the classifier as 'penult'.\n",
        "    \"\"\"\n",
        "    def __init__(self, num_classes=NUM_CLASSES):\n",
        "        super().__init__()\n",
        "        base = tvm.vgg16(weights=None)  # offline-friendly initialization\n",
        "\n",
        "        # Convolutional feature extractor and pooling\n",
        "        self.features = base.features\n",
        "        self.avgpool  = base.avgpool\n",
        "\n",
        "        # Classifier up to the final output layer\n",
        "        self.feat_extractor = nn.Sequential(*list(base.classifier.children())[:-1])\n",
        "\n",
        "        # Final classification layer\n",
        "        in_feats = base.classifier[-1].in_features\n",
        "        self.fc = nn.Linear(in_feats, num_classes)\n",
        "        self.feat_dim = in_feats\n",
        "\n",
        "    def forward(self, x, return_feats=False):\n",
        "        # Forward pass through the convolutional backbone\n",
        "        x = self.features(x)\n",
        "        x = self.avgpool(x)\n",
        "        x = torch.flatten(x, 1)\n",
        "\n",
        "        # Penultimate representation\n",
        "        pen = self.feat_extractor(x)\n",
        "\n",
        "        # Final logits\n",
        "        logits = self.fc(pen)\n",
        "\n",
        "        if not return_feats:\n",
        "            return logits\n",
        "        feats = {'penult': pen}\n",
        "        return logits, feats\n",
        "\n",
        "def build_model():\n",
        "    return MultiLayerVGG16(NUM_CLASSES).to(device)\n",
        "\n",
        "# Clean-accuracy evaluation\n",
        "@torch.no_grad()\n",
        "def eval_clean_acc(model, loader):\n",
        "    model.eval()\n",
        "    n, corr = 0, 0\n",
        "    for xb, yb in loader:\n",
        "        xb, yb = xb.to(device), yb.to(device)\n",
        "        pred = model(xb).argmax(1)\n",
        "        n += yb.numel(); corr += (pred == yb).sum().item()\n",
        "    return 100.0 * corr / max(1, n)\n",
        "\n",
        "# Attack success rate (ASR) evaluation on triggered samples\n",
        "@torch.no_grad()\n",
        "def eval_asr_simple(model, trigger_loader, target_class=TARGET_CLASS):\n",
        "    \"\"\"\n",
        "    ASR = P(model(x_trig) == target_class) over ALL triggered samples.\n",
        "    (Ignores dataset labels.)\n",
        "    \"\"\"\n",
        "    model.eval()\n",
        "    n, hit = 0, 0\n",
        "    for xb, _ in trigger_loader:\n",
        "        xb = xb.to(device)\n",
        "        pred = model(xb).argmax(1)\n",
        "        n += pred.numel(); hit += (pred == target_class).sum().item()\n",
        "    return 100.0 * hit / max(1, n)\n",
        "\n",
        "# Snapshot current trainable parameters for regularization\n",
        "def snapshot_params(model):\n",
        "    return {\n",
        "        n: p.detach().clone()\n",
        "        for n, p in model.named_parameters()\n",
        "        if p.requires_grad and p.is_floating_point()\n",
        "    }\n",
        "\n",
        "# Normalized L2 distance between current parameters and a reference snapshot\n",
        "# BatchNorm and bias terms are excluded\n",
        "def l2_params_normalized(model, prev):\n",
        "    s = torch.tensor(0.0, device=device); n_elem = 0\n",
        "    for n, p in model.named_parameters():\n",
        "        if not (p.requires_grad and p.is_floating_point()):\n",
        "            continue\n",
        "        if 'bn' in n.lower() or n.endswith('.bias'):\n",
        "            continue\n",
        "        d = (p - prev[n]).float()\n",
        "        s += (d * d).sum()\n",
        "        n_elem += d.numel()\n",
        "    return s / max(1, n_elem)\n",
        "\n",
        "# Aggregation option; BatchNorm parameters/statistics can be excluded if needed\n",
        "EXCLUDE_BN_IN_AGG = False  # set True to emulate FedBN\n",
        "\n",
        "# Build flattening metadata for state_dict aggregation\n",
        "def _get_flatten_info_from_state_dict(sd):\n",
        "    \"\"\"\n",
        "    Returns a list of (key, numel, shape, dtype) for float tensors to be aggregated.\n",
        "    Respects EXCLUDE_BN_IN_AGG.\n",
        "    \"\"\"\n",
        "    info = []\n",
        "    for k, v in sd.items():\n",
        "        if not (isinstance(v, torch.Tensor) and torch.is_floating_point(v)):\n",
        "            continue\n",
        "        if EXCLUDE_BN_IN_AGG and ('running_mean' in k or 'running_var' in k or '.bn' in k):\n",
        "            # Skip BatchNorm statistics/parameters if emulating FedBN\n",
        "            continue\n",
        "        info.append((k, v.numel(), v.shape, v.dtype))\n",
        "    return info\n",
        "\n",
        "# Flatten selected tensors into a single vector\n",
        "def _flatten_with_info(sd, info):\n",
        "    vecs = []\n",
        "    for (k, numel, shape, dtype) in info:\n",
        "        vecs.append(sd[k].detach().cpu().float().view(-1))\n",
        "    return torch.cat(vecs, dim=0)\n",
        "\n",
        "# Reconstruct a state_dict from an aggregated vector\n",
        "def _vector_to_state_dict(vec, template_sd, info):\n",
        "    out = {}\n",
        "    offset = 0\n",
        "    for (k, numel, shape, dtype) in info:\n",
        "        chunk = vec[offset:offset+numel].view(shape)\n",
        "        offset += numel\n",
        "        out[k] = chunk.to(dtype)\n",
        "    # Preserve untouched entries from the template\n",
        "    for k, v in template_sd.items():\n",
        "        if k not in out:\n",
        "            out[k] = v.detach().clone()\n",
        "    return out\n",
        "\n",
        "# MultiKrum on flattened client model vectors\n",
        "def multi_krum_vectors(grad_list: List[torch.Tensor], f: int, m: int = None):\n",
        "    \"\"\"\n",
        "    MultiKrum on list of flattened model vectors.\n",
        "    grad_list: list of 1D tensors, length N\n",
        "    f: max number of Byzantine clients\n",
        "    m: number of selected good clients to average; if None, use N - f - 2\n",
        "    Returns:\n",
        "       agg_vec: aggregated 1D tensor\n",
        "       selected_idx: indices of selected clients\n",
        "    \"\"\"\n",
        "    N = len(grad_list)\n",
        "    assert N >= 2 * f + 3, f\"MultiKrum requires N >= 2f+3, got N={N}, f={f}\"\n",
        "    grads = torch.stack(grad_list, dim=0)  # [N, D]\n",
        "\n",
        "    if m is None:\n",
        "        m = N - f - 2\n",
        "    m = max(1, min(m, N - f - 2))\n",
        "\n",
        "    # Pairwise squared distances between client updates\n",
        "    with torch.no_grad():\n",
        "        dists = torch.cdist(grads, grads, p=2) ** 2  # [N, N]\n",
        "        scores = torch.zeros(N)\n",
        "        nbhd = N - f - 2  # number of nearest neighbors used per client\n",
        "        for i in range(N):\n",
        "            row = dists[i].clone()\n",
        "            row[i] = float('inf')          # ignore self-distance\n",
        "            vals, _ = torch.sort(row)      # ascending\n",
        "            scores[i] = vals[:nbhd].sum()  # sum closest nbhd distances\n",
        "\n",
        "        # Select the lowest-score updates and average them\n",
        "        _, idx = torch.topk(scores, k=m, largest=False)\n",
        "        agg = grads[idx].mean(dim=0)\n",
        "\n",
        "    return agg, idx.cpu().tolist()\n",
        "\n",
        "# MultiKrum directly on client state_dicts\n",
        "def multi_krum_state_dicts(dicts: List[Dict[str, torch.Tensor]], f: int, m: int = None):\n",
        "    \"\"\"\n",
        "    MultiKrum aggregator on list of state_dicts.\n",
        "    Returns aggregated state_dict and list of selected client indices.\n",
        "    \"\"\"\n",
        "    template_sd = dicts[0]\n",
        "    info = _get_flatten_info_from_state_dict(template_sd)\n",
        "    grad_list = [_flatten_with_info(sd, info) for sd in dicts]\n",
        "    agg_vec, selected_idx = multi_krum_vectors(grad_list, f=f, m=m)\n",
        "    agg_sd = _vector_to_state_dict(agg_vec, template_sd, info)\n",
        "    return agg_sd, selected_idx\n",
        "\n",
        "# Standard benign local training\n",
        "def train_benign(model, loader, epochs=LOCAL_EPOCHS):\n",
        "    model.train()\n",
        "    opt = torch.optim.SGD(model.parameters(), lr=LR, momentum=MOMENTUM, weight_decay=WEIGHT_DECAY)\n",
        "    tot_loss, tot_n, tot_corr = 0.0, 0, 0\n",
        "    for _ in range(epochs):\n",
        "        for xb, yb in loader:\n",
        "            xb, yb = xb.to(device), yb.to(device)\n",
        "            logits = model(xb)\n",
        "            loss = F.cross_entropy(logits, yb)\n",
        "            opt.zero_grad(set_to_none=True); loss.backward(); opt.step()\n",
        "            tot_loss += loss.item() * yb.size(0)\n",
        "            tot_n    += yb.size(0)\n",
        "            tot_corr += (logits.argmax(1) == yb).sum().item()\n",
        "    return (tot_loss / max(1, tot_n), 100.0 * tot_corr / max(1, tot_n))\n",
        "\n",
        "# Transform used for malicious local image loading\n",
        "if 'train_tf' in globals():\n",
        "    _mal_transform = train_tf\n",
        "else:\n",
        "    import torchvision.transforms as T\n",
        "    _mal_transform = T.Compose([\n",
        "        T.Resize(224),\n",
        "        T.RandomHorizontalFlip(),\n",
        "        T.ToTensor(),\n",
        "        T.Normalize((0.5,0.5,0.5),(0.5,0.5,0.5))\n",
        "    ])\n",
        "\n",
        "# Load a list of image paths into a batch tensor\n",
        "def _load_batch(paths: List[str]):\n",
        "    ims = []\n",
        "    for p in paths:\n",
        "        img = Image.open(p).convert('RGB')\n",
        "        ims.append(_mal_transform(img))\n",
        "    xb = torch.stack(ims, dim=0).to(device)\n",
        "    return xb\n",
        "\n",
        "# Margin-based separation loss between clean and triggered features\n",
        "def separation_margin_loss(f_clean, f_trig, margin=0.02):\n",
        "    dist = F.mse_loss(f_clean, f_trig)          # >= 0\n",
        "    loss_sep = F.relu(margin - dist)            # max(0, margin - dist)\n",
        "    return loss_sep, dist\n",
        "\n",
        "# Malicious local training with:\n",
        "# - clean-only samples\n",
        "# - trigger-only samples\n",
        "# - paired clean/trigger samples\n",
        "# - feature separation\n",
        "# - regularization to the global model\n",
        "# - Neurotoxin gradient masking\n",
        "def train_malicious_separation(\n",
        "    model,\n",
        "    client_idx: int,\n",
        "    epochs: int = LOCAL_EPOCHS,\n",
        "    lambda_sep: float = LAMBDA_SEP,\n",
        "    lambda_reg: float = LAMBDA_REG,\n",
        "    batch_size: int = MAL_BATCH_SIZE,\n",
        "):\n",
        "    \"\"\"\n",
        "    For malicious client `client_idx`:\n",
        "    1. Generate Neurotoxin Mask based on clean data gradients.\n",
        "    2. Train with Separation + Backdoor + Regularizer.\n",
        "    3. Apply Mask before opt.step() to preserve clean accuracy.\n",
        "    \"\"\"\n",
        "    clean_df = mal_clean_splits[client_idx]\n",
        "    trig_df  = mal_trig_splits[client_idx]\n",
        "\n",
        "    # Partition this malicious client's data into paired, clean-only, and trigger-only subsets\n",
        "    clean_ids = set(clean_df['image_id'].tolist())\n",
        "    trig_ids  = set(trig_df['image_id'].tolist())\n",
        "\n",
        "    paired_ids      = clean_ids & trig_ids\n",
        "    clean_only_ids  = clean_ids - paired_ids\n",
        "    trig_only_ids   = trig_ids  - paired_ids\n",
        "\n",
        "    clean_map = {\n",
        "        r['image_id']: (r['local_path'], int(r['hair_label']))\n",
        "        for _, r in clean_df.iterrows()\n",
        "    }\n",
        "    trig_map = {\n",
        "        r['image_id']: (r['local_path'], TARGET_CLASS)\n",
        "        for _, r in trig_df.iterrows()\n",
        "    }\n",
        "\n",
        "    # Build record lists used in malicious local training\n",
        "    clean_only_records = [(clean_map[i][0], clean_map[i][1]) for i in sorted(clean_only_ids)]\n",
        "    trig_only_records  = []\n",
        "    for i in sorted(trig_only_ids):\n",
        "        if i in trig_map:\n",
        "            trig_only_records.append((trig_map[i][0], trig_map[i][1]))  # label = TARGET_CLASS\n",
        "\n",
        "    paired_records = []\n",
        "    for i in sorted(paired_ids):\n",
        "        if (i in clean_map) and (i in trig_map):\n",
        "            p_clean, y_clean = clean_map[i]\n",
        "            p_trig,  y_trig  = trig_map[i]\n",
        "            paired_records.append((p_clean, y_clean, p_trig, y_trig))\n",
        "\n",
        "    # Generate Neurotoxin gradient mask from clean-only samples\n",
        "    model.eval()\n",
        "    model.zero_grad()\n",
        "\n",
        "    neuro_sample_size = min(len(clean_only_records), 256)\n",
        "    if neuro_sample_size > 0:\n",
        "        neuro_idxs = np.random.choice(len(clean_only_records), neuro_sample_size, replace=False)\n",
        "        neuro_paths = [clean_only_records[k][0] for k in neuro_idxs]\n",
        "        neuro_labels = torch.tensor([clean_only_records[k][1] for k in neuro_idxs],\n",
        "                                    dtype=torch.long, device=device)\n",
        "\n",
        "        neuro_xb = _load_batch(neuro_paths)\n",
        "        neuro_out = model(neuro_xb)\n",
        "        neuro_loss = F.cross_entropy(neuro_out, neuro_labels)\n",
        "        neuro_loss.backward()\n",
        "\n",
        "    # Collect absolute gradients for thresholding\n",
        "    grad_abs_list = []\n",
        "    for n, p in model.named_parameters():\n",
        "        if p.grad is not None:\n",
        "            grad_abs_list.append(p.grad.detach().abs().view(-1).cpu())\n",
        "\n",
        "    if grad_abs_list:\n",
        "        all_grads = torch.cat(grad_abs_list)\n",
        "\n",
        "        # Subsample gradients before quantile computation if needed\n",
        "        max_samples = 1_000_000  # 1e6 entries to estimate the quantile\n",
        "        if all_grads.numel() > max_samples:\n",
        "            idx = torch.randperm(all_grads.numel())[:max_samples]\n",
        "            sample = all_grads[idx]\n",
        "        else:\n",
        "            sample = all_grads\n",
        "\n",
        "        # Compute Neurotoxin threshold\n",
        "        threshold = torch.quantile(sample, 1.0 - NEUROTOXIN_RATIO).to(device)\n",
        "\n",
        "        neurotoxin_mask = {}\n",
        "        for n, p in model.named_parameters():\n",
        "            if p.grad is not None:\n",
        "                g_abs = p.grad.detach().abs()\n",
        "                # 1 = updatable, 0 = frozen\n",
        "                neurotoxin_mask[n] = (g_abs < threshold).float()\n",
        "    else:\n",
        "        neurotoxin_mask = {n: torch.ones_like(p) for n, p in model.named_parameters()}\n",
        "\n",
        "    model.zero_grad()\n",
        "\n",
        "    model.train()\n",
        "    opt = torch.optim.SGD(model.parameters(), lr=LR, momentum=MOMENTUM, weight_decay=WEIGHT_DECAY)\n",
        "    theta_global = snapshot_params(model)\n",
        "    tot_loss, tot_n, tot_corr = 0.0, 0, 0\n",
        "\n",
        "    # Apply the Neurotoxin mask before each optimizer step\n",
        "    def _apply_neurotoxin_mask():\n",
        "        for n, p in model.named_parameters():\n",
        "            if p.grad is not None and n in neurotoxin_mask:\n",
        "                p.grad.data *= neurotoxin_mask[n]\n",
        "\n",
        "    # Train on clean-only samples\n",
        "    def _run_clean_only_epoch():\n",
        "        nonlocal tot_loss, tot_n, tot_corr\n",
        "        if not clean_only_records: return\n",
        "        idxs = np.random.permutation(len(clean_only_records))\n",
        "        for start in range(0, len(idxs), batch_size):\n",
        "            batch_idx = idxs[start:start+batch_size]\n",
        "            if len(batch_idx) == 0: continue\n",
        "            paths = [clean_only_records[k][0] for k in batch_idx]\n",
        "            labels = torch.tensor([clean_only_records[k][1] for k in batch_idx],\n",
        "                                  dtype=torch.long, device=device)\n",
        "\n",
        "            xb = _load_batch(paths)\n",
        "            logits = model(xb)\n",
        "            loss_ce = F.cross_entropy(logits, labels)\n",
        "            loss_reg = l2_params_normalized(model, theta_global) if lambda_reg > 0 else torch.tensor(0.0, device=device)\n",
        "            loss = loss_ce + lambda_reg * loss_reg\n",
        "\n",
        "            opt.zero_grad(set_to_none=True)\n",
        "            loss.backward()\n",
        "            _apply_neurotoxin_mask()\n",
        "            opt.step()\n",
        "\n",
        "            n = labels.size(0)\n",
        "\n",
        "            tot_loss += loss.item() * n\n",
        "            tot_n    += n\n",
        "            tot_corr += (logits.argmax(1) == labels).sum().item()\n",
        "\n",
        "    # Train on trigger-only samples\n",
        "    def _run_trig_only_epoch():\n",
        "        nonlocal tot_loss, tot_n, tot_corr\n",
        "        if not trig_only_records: return\n",
        "        idxs = np.random.permutation(len(trig_only_records))\n",
        "        for start in range(0, len(idxs), batch_size):\n",
        "            batch_idx = idxs[start:start+batch_size]\n",
        "            if len(batch_idx) == 0: continue\n",
        "            paths = [trig_only_records[k][0] for k in batch_idx]\n",
        "            labels = torch.tensor([trig_only_records[k][1] for k in batch_idx],\n",
        "                                  dtype=torch.long, device=device)\n",
        "\n",
        "            xb = _load_batch(paths)\n",
        "            logits = model(xb)\n",
        "            loss_ce = F.cross_entropy(logits, labels)\n",
        "            loss_reg = l2_params_normalized(model, theta_global) if lambda_reg > 0 else torch.tensor(0.0, device=device)\n",
        "            loss = loss_ce + lambda_reg * loss_reg\n",
        "\n",
        "            opt.zero_grad(set_to_none=True)\n",
        "            loss.backward()\n",
        "            _apply_neurotoxin_mask()\n",
        "            opt.step()\n",
        "\n",
        "            n = labels.size(0)\n",
        "\n",
        "            tot_loss += loss.item() * n\n",
        "            tot_n    += n\n",
        "            tot_corr += (logits.argmax(1) == labels).sum().item()\n",
        "\n",
        "    # Train on paired clean/trigger samples with feature separation\n",
        "    def _run_paired_epoch():\n",
        "        nonlocal tot_loss, tot_n, tot_corr\n",
        "        if not paired_records: return\n",
        "        idxs = np.random.permutation(len(paired_records))\n",
        "        for start in range(0, len(idxs), batch_size):\n",
        "            batch_idx = idxs[start:start+batch_size]\n",
        "            if len(batch_idx) == 0: continue\n",
        "\n",
        "            clean_paths = [paired_records[k][0] for k in batch_idx]\n",
        "            trig_paths  = [paired_records[k][2] for k in batch_idx]\n",
        "            y_clean = torch.tensor([paired_records[k][1] for k in batch_idx],\n",
        "                                   dtype=torch.long, device=device)\n",
        "            y_trig  = torch.tensor([paired_records[k][3] for k in batch_idx],\n",
        "                                   dtype=torch.long, device=device)\n",
        "\n",
        "            xb_clean = _load_batch(clean_paths)\n",
        "            xb_trig  = _load_batch(trig_paths)\n",
        "\n",
        "            logits_clean, feats_clean = model(xb_clean, return_feats=True)\n",
        "            logits_trig,  feats_trig  = model(xb_trig,  return_feats=True)\n",
        "\n",
        "            loss_ce = 0.5 * (F.cross_entropy(logits_clean, y_clean) +\n",
        "                             F.cross_entropy(logits_trig,  y_trig))\n",
        "            loss_sep, dist = separation_margin_loss(feats_clean['penult'], feats_trig['penult'], margin=0.2)\n",
        "            loss_reg = l2_params_normalized(model, theta_global) if lambda_reg > 0 else torch.tensor(0.0, device=device)\n",
        "\n",
        "            loss = loss_ce + lambda_sep * loss_sep + lambda_reg * loss_reg\n",
        "\n",
        "            opt.zero_grad(set_to_none=True)\n",
        "            loss.backward()\n",
        "            _apply_neurotoxin_mask()\n",
        "            opt.step()\n",
        "\n",
        "            pred_clean = logits_clean.argmax(1)\n",
        "            pred_trig  = logits_trig.argmax(1)\n",
        "            n = y_clean.size(0) + y_trig.size(0)\n",
        "\n",
        "            tot_loss += loss.item() * n\n",
        "            tot_n    += n\n",
        "            tot_corr += (pred_clean == y_clean).sum().item()\n",
        "            tot_corr += (pred_trig  == y_trig).sum().item()\n",
        "\n",
        "    for _ in range(epochs):\n",
        "        _run_clean_only_epoch()\n",
        "        _run_trig_only_epoch()\n",
        "        _run_paired_epoch()\n",
        "\n",
        "    return (tot_loss / max(1, tot_n), 100.0 * tot_corr / max(1, tot_n))\n",
        "\n",
        "# Client sample counts (not used by MultiKrum; retained for reference)\n",
        "def client_weight(idx):\n",
        "    if idx < NUM_MAL:\n",
        "        return len(mal_clean_splits[idx]) + len(mal_trig_splits[idx])\n",
        "    else:\n",
        "        b = idx - NUM_MAL\n",
        "        return len(benign_loaders[b].dataset)\n",
        "\n",
        "total_w = float(sum(client_weight(i) for i in range(NUM_CLIENTS)))\n",
        "weights = [client_weight(i)/max(1.0, total_w) for i in range(NUM_CLIENTS)]  # unused\n",
        "\n",
        "# Federated training loop\n",
        "global_model = build_model()\n",
        "\n",
        "# Round-wise history\n",
        "hist = {\n",
        "    'round': [],\n",
        "    'benign_acc': [], 'benign_asr': [],\n",
        "    'mal_acc': [],    'mal_asr': [],\n",
        "    'global_acc': [], 'global_asr': [],\n",
        "}\n",
        "\n",
        "print(f\"[INFO] NUM_CLIENTS={NUM_CLIENTS}, NUM_MAL={NUM_MAL}, KRUM_F={KRUM_F}, KRUM_M={KRUM_M}\")\n",
        "print(f\"[INFO] NEUROTOXIN ENABLED. Preservation Ratio: {NEUROTOXIN_RATIO}\")\n",
        "\n",
        "for rnd in range(1, ROUNDS+1):\n",
        "\n",
        "    late_phase  = (rnd > ROUNDS * (1.0 - LATE_FRAC))\n",
        "    poison_rate = POISON_RATE_LATE if late_phase else POISON_RATE_EARLY\n",
        "\n",
        "    local_states = []\n",
        "    per_client_clean_acc, per_client_asr = [], []\n",
        "\n",
        "    print(f\"\\n===== Round {rnd}/{ROUNDS} (poison_rate={poison_rate:.2f}) =====\")\n",
        "    for i in range(NUM_CLIENTS):\n",
        "        # Each client starts from the current global model\n",
        "        local = build_model()\n",
        "        local.load_state_dict(copy.deepcopy(global_model.state_dict()))\n",
        "\n",
        "        if i < NUM_MAL:\n",
        "            # Malicious local training\n",
        "            loss_tr, acc_tr = train_malicious_separation(\n",
        "                local,\n",
        "                client_idx=i,\n",
        "                epochs=LOCAL_EPOCHS,\n",
        "                lambda_sep=LAMBDA_SEP,\n",
        "                lambda_reg=LAMBDA_REG,\n",
        "                batch_size=MAL_BATCH_SIZE,\n",
        "            )\n",
        "            tag = f\"Client M{i}\"\n",
        "        else:\n",
        "            # Benign local training\n",
        "            b = i - NUM_MAL\n",
        "            loss_tr, acc_tr = train_benign(local, benign_loaders[b], epochs=LOCAL_EPOCHS)\n",
        "            tag = f\"Client B{b}\"\n",
        "\n",
        "        # Evaluate each local model before aggregation\n",
        "        lc = eval_clean_acc(local, test_clean_loader)\n",
        "        la = eval_asr_simple(local, test_trig_loader, TARGET_CLASS)\n",
        "        per_client_clean_acc.append(lc); per_client_asr.append(la)\n",
        "\n",
        "        print(f\"  {tag}: train_loss={loss_tr:.4f} | train_acc={acc_tr:5.2f}% | \"\n",
        "              f\"testACC={lc:5.2f}% | ASR={la:5.2f}%\")\n",
        "\n",
        "        # Move local weights to CPU before server aggregation\n",
        "        local_states.append({k: v.detach().cpu() for k, v in local.state_dict().items()})\n",
        "        del local\n",
        "        if torch.cuda.is_available():\n",
        "            torch.cuda.empty_cache()\n",
        "\n",
        "    # Robust server aggregation using MultiKrum\n",
        "    agg_sd, selected_idx = multi_krum_state_dicts(local_states, f=KRUM_F, m=KRUM_M)\n",
        "    print(f\"[MultiKrum] Selected client indices: {selected_idx}\")\n",
        "\n",
        "    # Load aggregated weights into the global model\n",
        "    gsd = global_model.state_dict()\n",
        "    for k, v in agg_sd.items():\n",
        "        if torch.is_tensor(v):\n",
        "            gsd[k] = v.to(dtype=gsd[k].dtype, device=gsd[k].device)\n",
        "        else:\n",
        "            gsd[k] = v\n",
        "    global_model.load_state_dict(gsd)\n",
        "\n",
        "    # Group-level summary metrics\n",
        "    benign_clean = float(np.mean(per_client_clean_acc[NUM_MAL:])) if NUM_BENIGN > 0 else float('nan')\n",
        "    benign_asr   = float(np.mean(per_client_asr[NUM_MAL:]))       if NUM_BENIGN > 0 else float('nan')\n",
        "    mal_clean    = float(np.mean(per_client_clean_acc[:NUM_MAL])) if NUM_MAL > 0 else float('nan')\n",
        "    mal_asr      = float(np.mean(per_client_asr[:NUM_MAL]))       if NUM_MAL > 0 else float('nan')\n",
        "\n",
        "    # Evaluate the aggregated global model\n",
        "    g_acc = eval_clean_acc(global_model, test_clean_loader)\n",
        "    g_asr = eval_asr_simple(global_model, test_trig_loader, TARGET_CLASS)\n",
        "\n",
        "    print(f\"[Round {rnd} Summary]  \"\n",
        "          f\"Benign ACC={benign_clean:5.2f}% | Benign ASR={benign_asr:5.2f}%  ||  \"\n",
        "          f\"Malicious ACC={mal_clean:5.2f}% | Malicious ASR={mal_asr:5.2f}%\")\n",
        "    print(f\"[Round {rnd} Global ]  ACC={g_acc:5.2f}% | ASR={g_asr:5.2f}%\")\n",
        "\n",
        "    # Record round-level metrics\n",
        "    hist['round'].append(rnd)\n",
        "    hist['benign_acc'].append(benign_clean); hist['benign_asr'].append(benign_asr)\n",
        "    hist['mal_acc'].append(mal_clean);       hist['mal_asr'].append(mal_asr)\n",
        "    hist['global_acc'].append(g_acc);        hist['global_asr'].append(g_asr)\n",
        "\n",
        "print(\"\\n[DONE] FL with feature-separation + MultiKrum aggregation + Neurotoxin complete. Hist keys:\", list(hist.keys()))"
      ]
    },
    {
      "cell_type": "code",
      "source": [
        "# Mount Google Drive\n",
        "from google.colab import drive\n",
        "drive.mount('/content/drive', force_remount=True)\n",
        "\n",
        "import os\n",
        "import pandas as pd\n",
        "import numpy as np\n",
        "\n",
        "# Path to your CSV\n",
        "save_dir = \"/content/drive/MyDrive/CelebA/Models/SABLE_MultiKrum\"\n",
        "csv_path = os.path.join(save_dir, \"hist_firstcode.csv\")\n",
        "\n",
        "# Load history\n",
        "hist_df = pd.read_csv(csv_path)\n",
        "\n",
        "# Make sure rounds are in order\n",
        "hist_df = hist_df.sort_values(\"round\")\n",
        "\n",
        "# Use last 10 rounds (or fewer if you don't have 10)\n",
        "K = 10\n",
        "if len(hist_df) < K:\n",
        "    print(f\"Warning: only {len(hist_df)} rounds found, using all of them.\")\n",
        "    last_df = hist_df\n",
        "else:\n",
        "    last_df = hist_df.tail(K)\n",
        "\n",
        "# Compute mean and std for global ACC and ASR\n",
        "acc_mean = last_df[\"global_acc\"].mean()\n",
        "acc_std  = last_df[\"global_acc\"].std(ddof=1)  # sample std\n",
        "asr_mean = last_df[\"global_asr\"].mean()\n",
        "asr_std  = last_df[\"global_asr\"].std(ddof=1)\n",
        "\n",
        "print(f\"Global ACC over last {len(last_df)} rounds: mean = {acc_mean:.2f}%, std = {acc_std:.2f}\")\n",
        "print(f\"Global ASR over last {len(last_df)} rounds: mean = {asr_mean:.2f}%, std = {asr_std:.2f}\")"
      ],
      "metadata": {
        "id": "YaW18E09TtK7"
      },
      "execution_count": null,
      "outputs": []
    }
  ],
  "metadata": {
    "accelerator": "GPU",
    "colab": {
      "collapsed_sections": [
        "f3gWM1nyENxL",
        "U9CoxRsd60ve",
        "nRXPydFTG5Xc",
        "4X0MU-SPF2Kn",
        "i8lZKoF1XO4I",
        "52ywermYaZV9",
        "5ossow-neGOC"
      ],
      "gpuType": "A100",
      "machine_shape": "hm",
      "provenance": []
    },
    "kernelspec": {
      "display_name": "Python 3",
      "name": "python3"
    },
    "language_info": {
      "name": "python"
    }
  },
  "nbformat": 4,
  "nbformat_minor": 0
}