{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"0bfc0593-b08b-4f77-af67-2001c90bde39","cell_type":"markdown","source":"# Task 2 - Channel Occlusion Audit\n\nThis notebook implements Task 2 in\n`yeu_cau_dieu_chinh_thuc_nghiem_camera_ready.md`.\n\nIt does **not** train or fine-tune a model. For every test sample it evaluates\nthe original PPS input and three channel-occluded variants, saves all true-class\nprobabilities, and derives every summary and figure from the same per-sample CSV.\n\nRequired outputs:\n\n- `channel_occlusion_per_sample.csv`\n- `channel_occlusion_family_summary.csv`\n- `channel_occlusion_dataset_summary.csv`\n- `channel_occlusion_consistency_report.md`\n- regenerated `channel_attribution_heatmap.png`\n\n","metadata":{}},{"id":"40dff105-7417-4649-acad-bbaf1e1d731d","cell_type":"code","source":"# ============================================================\n# CELL 1: Imports and configuration\n# ============================================================\nimport csv\nimport gc\nimport glob\nimport hashlib\nimport json\nimport logging\nimport os\nimport sys\nfrom pathlib import Path\n\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\nos.environ[\"OMP_NUM_THREADS\"] = \"2\"\nos.environ[\"MKL_NUM_THREADS\"] = \"2\"\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nfrom PIL import Image\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import datasets, models, transforms\n\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nAMP_ENABLED = DEVICE.type == \"cuda\"\n\n# Run both datasets by default. Remove one key only for a deliberate partial audit.\nDATASETS_TO_RUN = [\"microsoft_rgb\", \"malimg_rgb\"]\n\nMICROSOFT_CKPT_OVERRIDE = None\nMALIMG_CKPT_OVERRIDE = None\nMICROSOFT_SPLIT_OVERRIDE = None\nMALIMG_SPLIT_OVERRIDE = None\n\nDATASET_CONFIGS = {\n    \"microsoft_rgb\": {\n        \"dataset_label\": \"BIG-2015\",\n        \"description\": \"Microsoft BIG-2015\",\n        \"path\": \"/kaggle/input/datasets/vnhtbo/microsoft/train_rgb/train_rgb\",\n        \"csv_path\": \"/kaggle/input/competitions/malware-classification/trainLabels.csv\",\n        \"num_classes\": 9,\n        \"checkpoint_override\": MICROSOFT_CKPT_OVERRIDE,\n        \"checkpoint_candidates\": [\n            \"/kaggle/input/datasets/vnhtbo/eb3-train/best_eb3_teacher_run1.pt\",\n            \"/kaggle/working/best_eb3_teacher_run1.pt\",\n        ],\n        \"checkpoint_patterns\": [\n            \"*/eb3-train/best_eb3_teacher_run1.pt\",\n            \"best_eb3_teacher_run1.pt\",\n        ],\n        \"split_override\": MICROSOFT_SPLIT_OVERRIDE,\n        \"split_candidates\": [\n            \"/kaggle/working/checkpoints/split_microsoft_rgb.json\",\n            \"/kaggle/input/datasets/vnhtbo/eb3-train/split_microsoft_rgb.json\",\n        ],\n        \"split_pattern\": \"split_microsoft_rgb.json\",\n    },\n    \"malimg_rgb\": {\n        \"dataset_label\": \"Malimg\",\n        \"description\": \"Malimg\",\n        \"path\": \"/kaggle/input/datasets/dongquan/malimg-rgb/kaggle/working/malimg_rgb\",\n        \"csv_path\": None,\n        \"num_classes\": 25,\n        \"checkpoint_override\": MALIMG_CKPT_OVERRIDE,\n        \"checkpoint_candidates\": [\n            \"/kaggle/input/datasets/vnhtbo/eb3-train-malimg/best_eb3_teacher_run1 (1).pt\",\n            \"/kaggle/input/datasets/vnhtbo/eb3-train-malimg/best_eb3_teacher_run1.pt\",\n        ],\n        \"checkpoint_patterns\": [\n            \"*/eb3-train-malimg/best_eb3_teacher_run1*.pt\",\n        ],\n        \"split_override\": MALIMG_SPLIT_OVERRIDE,\n        \"split_candidates\": [\n            \"/kaggle/working/checkpoints/split_malimg_rgb.json\",\n            \"/kaggle/input/datasets/vnhtbo/eb3-train-malimg/split_malimg_rgb.json\",\n        ],\n        \"split_pattern\": \"split_malimg_rgb.json\",\n    },\n}\n\nIMG_SIZE = 300\nBATCH_SIZE = 32\nNUM_WORKERS = 2\nVAL_RATIO = 0.15\nTEST_RATIO = 0.15\nSPLIT_SEED = 42\nTRAINING_RUN = 1\nTRAINING_SEED = SPLIT_SEED + TRAINING_RUN  # eb3-training.ipynb uses SEED + run.\nFORCE_RERUN = False\n\nNORM_MEAN = [0.485, 0.456, 0.406]\nNORM_STD = [0.229, 0.224, 0.225]\nOCCLUDE_VALUES = [\n    (0.0 - NORM_MEAN[index]) / NORM_STD[index]\n    for index in range(3)\n]\n\nMS_FAMILY_NAMES = {\n    \"1\": \"Ramnit\",\n    \"2\": \"Lollipop\",\n    \"3\": \"Kelihos_v3\",\n    \"4\": \"Vundo\",\n    \"5\": \"Simda\",\n    \"6\": \"Tracur\",\n    \"7\": \"Kelihos_v1\",\n    \"8\": \"Obfuscator.ACY\",\n    \"9\": \"Gatak\",\n}\n\n# Rounded values currently printed in the paper. They are used only for the\n# consistency report, never as inputs to a calculation or figure.\nPAPER_DATASET_VALUES = {\n    \"Malimg\": {\"delta_R\": 0.70, \"delta_G\": 0.55, \"delta_B\": 0.47},\n    \"BIG-2015\": {\"delta_R\": 0.21, \"delta_G\": 0.25, \"delta_B\": 0.09},\n}\nPAPER_TEST_METADATA = {\n    \"Malimg\": {\"num_samples\": 1400, \"accuracy\": 0.9943},\n    \"BIG-2015\": {\"num_samples\": 1630, \"accuracy\": 0.9908},\n}\nPAPER_FAMILY_CLAIMS = [\n    (\"Malimg\", \"Wintrim.BX\", \"delta_G\", 0.89),\n    (\"Malimg\", \"Wintrim.BX\", \"delta_R\", 0.20),\n    (\"Malimg\", \"Lolyda.AT\", \"delta_G\", 0.83),\n    (\"Malimg\", \"Lolyda.AT\", \"delta_B\", 0.01),\n    (\"Malimg\", \"Rbot!gen\", \"delta_G\", 0.79),\n    (\"Malimg\", \"Allaple.A\", \"delta_R\", 0.98),\n    (\"Malimg\", \"Malex.gen!J\", \"delta_R\", 0.96),\n    (\"Malimg\", \"Lolyda.AA3\", \"delta_B\", 0.31),\n    (\"BIG-2015\", \"Gatak\", \"delta_G\", 0.64),\n    (\"BIG-2015\", \"Tracur\", \"delta_G\", 0.47),\n    (\"BIG-2015\", \"Kelihos_v3\", \"delta_R\", 0.40),\n    (\"BIG-2015\", \"Kelihos_v3\", \"delta_G\", 0.18),\n]\n\nAUDIT_ROOT = Path(\"/kaggle/working/camera_ready_experiment_audit\")\nATTR_DIR = AUDIT_ROOT / \"attribution\"\nLOG_DIR = AUDIT_ROOT / \"logs\"\nATTR_DIR.mkdir(parents=True, exist_ok=True)\nLOG_DIR.mkdir(parents=True, exist_ok=True)\n\nPER_SAMPLE_CSV = ATTR_DIR / \"channel_occlusion_per_sample.csv\"\nFAMILY_SUMMARY_CSV = ATTR_DIR / \"channel_occlusion_family_summary.csv\"\nDATASET_SUMMARY_CSV = ATTR_DIR / \"channel_occlusion_dataset_summary.csv\"\nCONSISTENCY_REPORT = ATTR_DIR / \"channel_occlusion_consistency_report.md\"\nHEATMAP_PATH = ATTR_DIR / \"channel_attribution_heatmap.png\"\nRUN_METADATA_JSON = ATTR_DIR / \"channel_occlusion_run_metadata.json\"\nATTR_README = ATTR_DIR / \"README.md\"\nATTR_LOG = LOG_DIR / \"attribution.log\"\n\n\ndef configure_logger():\n    logger = logging.getLogger(\"channel_occlusion_audit\")\n    logger.setLevel(logging.INFO)\n    for handler in list(logger.handlers):\n        handler.close()\n        logger.removeHandler(handler)\n    formatter = logging.Formatter(\"%(asctime)s | %(levelname)s | %(message)s\")\n    file_handler = logging.FileHandler(ATTR_LOG, mode=\"a\", encoding=\"utf-8\")\n    file_handler.setFormatter(formatter)\n    stream_handler = logging.StreamHandler(sys.stdout)\n    stream_handler.setFormatter(formatter)\n    logger.addHandler(file_handler)\n    logger.addHandler(stream_handler)\n    return logger\n\n\nLOGGER = configure_logger()\nLOGGER.info(\"Task 2 audit configured on device=%s, AMP=%s\", DEVICE, AMP_ENABLED)\n\n","metadata":{},"outputs":[],"execution_count":null},{"id":"a6cef8de-86f1-40a4-a0e4-e2b937bcc0f7","cell_type":"markdown","source":"## Input and provenance preflight\n\nThe audit stops before inference if a required dataset or checkpoint is absent.\nSaved split JSON is preferred. If it is unavailable, the notebook reproduces\nthe original 70/15/15 split deterministically with seed 42 and records that fact.\n\n","metadata":{}},{"id":"68e7238f-2985-4e6a-9404-e6376fb5b835","cell_type":"code","source":"# ============================================================\n# CELL 2: Path, hash, dataset, and split helpers\n# ============================================================\ndef sha256_file(path, chunk_size=1024 * 1024):\n    digest = hashlib.sha256()\n    with open(path, \"rb\") as handle:\n        while True:\n            chunk = handle.read(chunk_size)\n            if not chunk:\n                break\n            digest.update(chunk)\n    return digest.hexdigest()\n\n\ndef sha256_json(value):\n    payload = json.dumps(value, sort_keys=True, separators=(\",\", \":\")).encode(\"utf-8\")\n    return hashlib.sha256(payload).hexdigest()\n\n\ndef checkpoint_num_classes(checkpoint_path):\n    checkpoint = torch.load(checkpoint_path, map_location=\"cpu\")\n    state = checkpoint\n    if isinstance(state, dict) and \"model\" in state:\n        state = state[\"model\"]\n    if isinstance(state, dict) and \"state_dict\" in state:\n        state = state[\"state_dict\"]\n    if not isinstance(state, dict):\n        raise TypeError(f\"Checkpoint has no state dictionary: {checkpoint_path}\")\n    for key in (\"classifier.1.weight\", \"module.classifier.1.weight\"):\n        weight = state.get(key)\n        if weight is not None and hasattr(weight, \"shape\") and len(weight.shape) == 2:\n            return int(weight.shape[0])\n    raise ValueError(f\"Cannot infer EfficientNet-B3 class count from {checkpoint_path}\")\n\n\ndef resolve_file(override, candidates, patterns, label):\n    if override:\n        if not os.path.isfile(override):\n            raise FileNotFoundError(f\"{label} override does not exist: {override}\")\n        return override\n    for path in candidates:\n        if path and os.path.isfile(path):\n            return path\n    for root in (\"/kaggle/working\", \"/kaggle/input\"):\n        if not os.path.isdir(root):\n            continue\n        for pattern in patterns:\n            matches = sorted(glob.glob(os.path.join(root, \"**\", pattern), recursive=True))\n            if matches:\n                return matches[0]\n    raise FileNotFoundError(\n        f\"Could not resolve {label}. Attach the original artifact or set its override.\"\n    )\n\n\ndef build_dataset_records(dataset_key, config):\n    root = config[\"path\"]\n    if not os.path.isdir(root):\n        raise FileNotFoundError(f\"Dataset directory not found: {root}\")\n\n    records = []\n    if config[\"csv_path\"]:\n        if not os.path.isfile(config[\"csv_path\"]):\n            raise FileNotFoundError(f\"Label CSV not found: {config['csv_path']}\")\n        labels = pd.read_csv(config[\"csv_path\"])\n        classes = [str(value) for value in sorted(labels[\"Class\"].unique())]\n        class_to_index = {int(value): index for index, value in enumerate(map(int, classes))}\n        existing = set(os.listdir(root))\n        for row in labels.itertuples(index=False):\n            sample_id = str(row.Id)\n            filename = f\"{sample_id}.png\"\n            if filename not in existing:\n                continue\n            class_token = str(int(row.Class))\n            records.append({\n                \"path\": os.path.join(root, filename),\n                \"sample_id\": sample_id,\n                \"label_index\": class_to_index[int(row.Class)],\n                \"family\": MS_FAMILY_NAMES.get(class_token, class_token),\n            })\n    else:\n        base = datasets.ImageFolder(root)\n        classes = list(base.classes)\n        root_path = Path(root)\n        for path, label_index in base.samples:\n            relative = Path(path).relative_to(root_path)\n            sample_id = str(relative.with_suffix(\"\"))\n            records.append({\n                \"path\": path,\n                \"sample_id\": sample_id,\n                \"label_index\": int(label_index),\n                \"family\": classes[int(label_index)],\n            })\n\n    if len(records) == 0:\n        raise RuntimeError(f\"No samples found for {dataset_key}\")\n    if len({record[\"sample_id\"] for record in records}) != len(records):\n        raise ValueError(f\"Sample IDs are not unique for {dataset_key}\")\n    return records, classes\n\n\ndef resolve_split(dataset_key, config, num_samples):\n    split_path = None\n    if config[\"split_override\"]:\n        split_path = resolve_file(\n            config[\"split_override\"], [], [], f\"{dataset_key} split\"\n        )\n    else:\n        for candidate in config[\"split_candidates\"]:\n            if os.path.isfile(candidate):\n                split_path = candidate\n                break\n        if split_path is None:\n            for root in (\"/kaggle/working\", \"/kaggle/input\"):\n                if not os.path.isdir(root):\n                    continue\n                matches = sorted(glob.glob(\n                    os.path.join(root, \"**\", config[\"split_pattern\"]),\n                    recursive=True,\n                ))\n                if matches:\n                    split_path = matches[0]\n                    break\n\n    if split_path:\n        with open(split_path, \"r\", encoding=\"utf-8\") as handle:\n            split = json.load(handle)\n        source = f\"loaded:{split_path}\"\n    else:\n        n_test = int(num_samples * TEST_RATIO)\n        n_val = int(num_samples * VAL_RATIO)\n        n_train = num_samples - n_val - n_test\n        generator = torch.Generator().manual_seed(SPLIT_SEED)\n        permutation = torch.randperm(num_samples, generator=generator).tolist()\n        split = {\n            \"train\": permutation[:n_train],\n            \"val\": permutation[n_train:n_train + n_val],\n            \"test\": permutation[n_train + n_val:],\n            \"seed\": SPLIT_SEED,\n            \"total\": num_samples,\n        }\n        source = \"deterministically_recreated_from_seed_42_and_70_15_15_ratios\"\n        LOGGER.warning(\"%s split JSON unavailable; recreated deterministically\", dataset_key)\n\n    if int(split.get(\"total\", num_samples)) != num_samples:\n        raise ValueError(\n            f\"{dataset_key} split total={split.get('total')} does not match {num_samples}\"\n        )\n    test_indices = [int(index) for index in split[\"test\"]]\n    if not test_indices or min(test_indices) < 0 or max(test_indices) >= num_samples:\n        raise ValueError(f\"{dataset_key} split contains invalid test indices\")\n\n    snapshot = dict(split)\n    snapshot[\"audit_source\"] = source\n    snapshot_path = ATTR_DIR / f\"split_{dataset_key}_snapshot.json\"\n    snapshot_path.write_text(json.dumps(snapshot, indent=2) + \"\\n\", encoding=\"utf-8\")\n    return split, source, sha256_json(split), snapshot_path\n\n\nRUN_CONTEXTS = {}\nfor dataset_key in DATASETS_TO_RUN:\n    config = DATASET_CONFIGS[dataset_key]\n    records, classes = build_dataset_records(dataset_key, config)\n    split, split_source, split_hash, snapshot_path = resolve_split(\n        dataset_key, config, len(records)\n    )\n    checkpoint_path = resolve_file(\n        config[\"checkpoint_override\"],\n        config[\"checkpoint_candidates\"],\n        config[\"checkpoint_patterns\"],\n        f\"{dataset_key} EfficientNet-B3 checkpoint\",\n    )\n    checkpoint_classes = checkpoint_num_classes(checkpoint_path)\n    if checkpoint_classes != config[\"num_classes\"]:\n        raise ValueError(\n            f\"{dataset_key} checkpoint has {checkpoint_classes} output classes; \"\n            f\"expected {config['num_classes']}. The wrong run was resolved: {checkpoint_path}\"\n        )\n    checkpoint_hash = sha256_file(checkpoint_path)\n    test_records = [records[index] for index in split[\"test\"]]\n    RUN_CONTEXTS[dataset_key] = {\n        \"config\": config,\n        \"classes\": classes,\n        \"test_records\": test_records,\n        \"checkpoint_path\": checkpoint_path,\n        \"checkpoint_sha256\": checkpoint_hash,\n        \"checkpoint_num_classes\": checkpoint_classes,\n        \"split_source\": split_source,\n        \"split_sha256\": split_hash,\n        \"split_snapshot\": str(snapshot_path),\n    }\n    LOGGER.info(\n        \"%s: total=%d test=%d checkpoint=%s split=%s\",\n        dataset_key,\n        len(records),\n        len(test_records),\n        checkpoint_path,\n        split_source,\n    )\n\nprint(pd.DataFrame([\n    {\n        \"dataset\": context[\"config\"][\"dataset_label\"],\n        \"num_test_samples\": len(context[\"test_records\"]),\n        \"checkpoint\": context[\"checkpoint_path\"],\n        \"checkpoint_sha256\": context[\"checkpoint_sha256\"],\n        \"split_source\": context[\"split_source\"],\n    }\n    for context in RUN_CONTEXTS.values()\n]).to_string(index=False))\n\n","metadata":{},"outputs":[],"execution_count":null},{"id":"3611b4b6-928b-4e30-93a3-0631354c154c","cell_type":"markdown","source":"## Per-sample channel occlusion\n\nOcclusion sets one channel to raw pixel value zero. Because tensors are already\nnormalized, the replacement value is `(0 - mean) / std` for that channel. The\nscore for each channel is always based on the true class:\n\n`delta_C = p(true class | full PPS) - p(true class | channel C zeroed)`.\n\n","metadata":{}},{"id":"7bd022a6-3497-4b06-8208-ae632a8d24f2","cell_type":"code","source":"# ============================================================\n# CELL 3: Dataset, model, and inference helpers\n# ============================================================\nVAL_TRANSFORM = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(NORM_MEAN, NORM_STD),\n])\n\n\nclass OcclusionDataset(Dataset):\n    def __init__(self, records):\n        self.records = records\n\n    def __len__(self):\n        return len(self.records)\n\n    def __getitem__(self, index):\n        record = self.records[index]\n        with Image.open(record[\"path\"]) as image:\n            tensor = VAL_TRANSFORM(image.convert(\"RGB\"))\n        return (\n            tensor,\n            int(record[\"label_index\"]),\n            str(record[\"sample_id\"]),\n            str(record[\"family\"]),\n        )\n\n\ndef build_efficientnet_b3(num_classes):\n    model = models.efficientnet_b3(weights=None)\n    in_features = model.classifier[1].in_features\n    model.classifier = nn.Sequential(\n        nn.Dropout(p=0.3, inplace=True),\n        nn.Linear(in_features, num_classes),\n    )\n    return model.to(DEVICE)\n\n\ndef clean_state_dict(checkpoint):\n    state = checkpoint\n    if isinstance(state, dict) and \"model\" in state:\n        state = state[\"model\"]\n    if isinstance(state, dict) and \"state_dict\" in state:\n        state = state[\"state_dict\"]\n    if not isinstance(state, dict):\n        raise TypeError(\"Checkpoint does not contain a state dictionary\")\n    cleaned = {}\n    for key, value in state.items():\n        if key.startswith(\"module.\"):\n            key = key[len(\"module.\"):]\n        cleaned[key] = value\n    return cleaned\n\n\ndef load_model(checkpoint_path, num_classes):\n    model = build_efficientnet_b3(num_classes)\n    checkpoint = torch.load(checkpoint_path, map_location=DEVICE)\n    model.load_state_dict(clean_state_dict(checkpoint), strict=True)\n    model.eval()\n    return model\n\n\ndef true_class_probabilities(model, images, labels):\n    with torch.autocast(\n        device_type=\"cuda\",\n        dtype=torch.float16,\n        enabled=AMP_ENABLED,\n    ):\n        logits = model(images)\n    probabilities = torch.softmax(logits, dim=1)\n    indices = torch.arange(images.size(0), device=images.device)\n    return logits, probabilities[indices, labels]\n\n\nRAW_COLUMNS = [\n    \"dataset\",\n    \"sample_id\",\n    \"family\",\n    \"true_label\",\n    \"p_full\",\n    \"p_without_R\",\n    \"p_without_G\",\n    \"p_without_B\",\n    \"delta_R\",\n    \"delta_G\",\n    \"delta_B\",\n    \"predicted_label_full\",\n    \"correct_full\",\n    \"true_label_index\",\n    \"predicted_label_index\",\n    \"checkpoint_path\",\n    \"checkpoint_sha256\",\n    \"split_source\",\n    \"split_sha256\",\n]\n\n\ndef dataset_raw_path(dataset_key):\n    return ATTR_DIR / f\"channel_occlusion_per_sample_{dataset_key}.csv\"\n\n\ndef load_existing_dataset_rows(dataset_key, context):\n    path = dataset_raw_path(dataset_key)\n    if FORCE_RERUN and path.exists():\n        path.unlink()\n    if not path.exists():\n        return pd.DataFrame(columns=RAW_COLUMNS)\n    frame = pd.read_csv(path, dtype={\"sample_id\": str})\n    missing = set(RAW_COLUMNS) - set(frame.columns)\n    if missing:\n        raise ValueError(\n            f\"Existing {path} is from an incompatible format; missing {sorted(missing)}. \"\n            \"Keep it as provenance, rename it, and rerun this audit.\"\n        )\n    checkpoint_hashes = set(frame[\"checkpoint_sha256\"].dropna().astype(str))\n    split_hashes = set(frame[\"split_sha256\"].dropna().astype(str))\n    if checkpoint_hashes not in (set(), {context[\"checkpoint_sha256\"]}):\n        raise ValueError(f\"Existing {path} uses a different checkpoint\")\n    if split_hashes not in (set(), {context[\"split_sha256\"]}):\n        raise ValueError(f\"Existing {path} uses a different test split\")\n    return frame[RAW_COLUMNS]\n\n\n@torch.no_grad()\ndef run_dataset_occlusion(dataset_key, context):\n    config = context[\"config\"]\n    output_path = dataset_raw_path(dataset_key)\n    existing = load_existing_dataset_rows(dataset_key, context)\n    complete_ids = set(existing[\"sample_id\"].astype(str))\n    pending_records = [\n        record\n        for record in context[\"test_records\"]\n        if str(record[\"sample_id\"]) not in complete_ids\n    ]\n\n    if pending_records:\n        LOGGER.info(\n            \"%s: running occlusion for %d/%d pending samples\",\n            dataset_key,\n            len(pending_records),\n            len(context[\"test_records\"]),\n        )\n        model = load_model(context[\"checkpoint_path\"], config[\"num_classes\"])\n        loader = DataLoader(\n            OcclusionDataset(pending_records),\n            batch_size=BATCH_SIZE,\n            shuffle=False,\n            num_workers=NUM_WORKERS,\n            pin_memory=DEVICE.type == \"cuda\",\n        )\n        mode = \"a\" if output_path.exists() else \"w\"\n        with open(output_path, mode, newline=\"\", encoding=\"utf-8\") as handle:\n            writer = csv.DictWriter(handle, fieldnames=RAW_COLUMNS)\n            if mode == \"w\":\n                writer.writeheader()\n\n            for images, labels, sample_ids, families in tqdm(\n                loader,\n                desc=f\"Occlusion {config['dataset_label']}\",\n            ):\n                images = images.to(DEVICE, non_blocking=True)\n                labels = labels.to(DEVICE, non_blocking=True)\n                logits_full, p_full = true_class_probabilities(model, images, labels)\n                predictions = logits_full.argmax(dim=1)\n\n                p_without = []\n                for channel_index, occlude_value in enumerate(OCCLUDE_VALUES):\n                    occluded = images.clone()\n                    occluded[:, channel_index, :, :] = float(occlude_value)\n                    _, probabilities = true_class_probabilities(model, occluded, labels)\n                    p_without.append(probabilities)\n                    del occluded\n\n                p_full_np = p_full.float().cpu().numpy()\n                p_without_np = [value.float().cpu().numpy() for value in p_without]\n                labels_np = labels.cpu().numpy()\n                predictions_np = predictions.cpu().numpy()\n\n                for index, sample_id in enumerate(sample_ids):\n                    true_index = int(labels_np[index])\n                    predicted_index = int(predictions_np[index])\n                    true_family = str(families[index])\n                    predicted_family = str(context[\"classes\"][predicted_index])\n                    if dataset_key == \"microsoft_rgb\":\n                        predicted_family = MS_FAMILY_NAMES.get(\n                            predicted_family, predicted_family\n                        )\n                    without_r = float(p_without_np[0][index])\n                    without_g = float(p_without_np[1][index])\n                    without_b = float(p_without_np[2][index])\n                    full = float(p_full_np[index])\n                    writer.writerow({\n                        \"dataset\": config[\"dataset_label\"],\n                        \"sample_id\": str(sample_id),\n                        \"family\": true_family,\n                        \"true_label\": true_family,\n                        \"p_full\": full,\n                        \"p_without_R\": without_r,\n                        \"p_without_G\": without_g,\n                        \"p_without_B\": without_b,\n                        \"delta_R\": full - without_r,\n                        \"delta_G\": full - without_g,\n                        \"delta_B\": full - without_b,\n                        \"predicted_label_full\": predicted_family,\n                        \"correct_full\": bool(true_index == predicted_index),\n                        \"true_label_index\": true_index,\n                        \"predicted_label_index\": predicted_index,\n                        \"checkpoint_path\": context[\"checkpoint_path\"],\n                        \"checkpoint_sha256\": context[\"checkpoint_sha256\"],\n                        \"split_source\": context[\"split_source\"],\n                        \"split_sha256\": context[\"split_sha256\"],\n                    })\n                handle.flush()\n                del images, labels, logits_full, p_full, p_without\n\n        del model, loader\n        gc.collect()\n        if DEVICE.type == \"cuda\":\n            torch.cuda.empty_cache()\n\n    frame = pd.read_csv(output_path, dtype={\"sample_id\": str})\n    expected_ids = {str(record[\"sample_id\"]) for record in context[\"test_records\"]}\n    observed_ids = set(frame[\"sample_id\"].astype(str))\n    if len(frame) != len(expected_ids) or observed_ids != expected_ids:\n        missing = sorted(expected_ids - observed_ids)[:5]\n        extra = sorted(observed_ids - expected_ids)[:5]\n        raise RuntimeError(\n            f\"{dataset_key} raw output is incomplete or duplicated: \"\n            f\"rows={len(frame)}, expected={len(expected_ids)}, \"\n            f\"missing={missing}, extra={extra}\"\n        )\n    return frame[RAW_COLUMNS]\n\n\nPER_DATASET_FRAMES = []\nfor dataset_key, context in RUN_CONTEXTS.items():\n    PER_DATASET_FRAMES.append(run_dataset_occlusion(dataset_key, context))\n\nPER_SAMPLE = pd.concat(PER_DATASET_FRAMES, ignore_index=True)\nPER_SAMPLE.to_csv(PER_SAMPLE_CSV, index=False)\nprint(PER_SAMPLE.groupby(\"dataset\").agg(\n    num_samples=(\"sample_id\", \"count\"),\n    accuracy=(\"correct_full\", \"mean\"),\n))\n\n","metadata":{},"outputs":[],"execution_count":null},{"id":"33f470b1-6c9e-4e19-a1b9-5dc2b955940c","cell_type":"markdown","source":"## Aggregation and figure regeneration\n\n`sample_weighted_mean` gives every test sample equal weight. `macro_family_mean`\nfirst averages within each family and then gives every family equal weight.\nBoth are exported so the paper text and heatmap can be reconciled explicitly.\n\n","metadata":{}},{"id":"1cc39601-4583-4ff5-9548-783678bd5310","cell_type":"code","source":"# ============================================================\n# CELL 4: Family and dataset summaries\n# ============================================================\nDELTA_COLUMNS = [\"delta_R\", \"delta_G\", \"delta_B\"]\n\n\ndef build_family_summary(per_sample):\n    grouped = per_sample.groupby([\"dataset\", \"family\"], sort=False)\n    rows = []\n    for (dataset_label, family), group in grouped:\n        rows.append({\n            \"dataset\": dataset_label,\n            \"family\": family,\n            \"num_samples\": int(len(group)),\n            \"delta_R_mean\": float(group[\"delta_R\"].mean()),\n            \"delta_G_mean\": float(group[\"delta_G\"].mean()),\n            \"delta_B_mean\": float(group[\"delta_B\"].mean()),\n            \"delta_R_std\": float(group[\"delta_R\"].std(ddof=1)),\n            \"delta_G_std\": float(group[\"delta_G\"].std(ddof=1)),\n            \"delta_B_std\": float(group[\"delta_B\"].std(ddof=1)),\n        })\n    return pd.DataFrame(rows)\n\n\ndef build_dataset_summary(per_sample, family_summary):\n    rows = []\n    for dataset_label, group in per_sample.groupby(\"dataset\", sort=False):\n        num_families = int(group[\"family\"].nunique())\n        rows.append({\n            \"dataset\": dataset_label,\n            \"aggregation\": \"sample_weighted_mean\",\n            \"delta_R\": float(group[\"delta_R\"].mean()),\n            \"delta_G\": float(group[\"delta_G\"].mean()),\n            \"delta_B\": float(group[\"delta_B\"].mean()),\n            \"num_samples\": int(len(group)),\n            \"num_families\": num_families,\n        })\n        family_group = family_summary[family_summary[\"dataset\"] == dataset_label]\n        rows.append({\n            \"dataset\": dataset_label,\n            \"aggregation\": \"macro_family_mean\",\n            \"delta_R\": float(family_group[\"delta_R_mean\"].mean()),\n            \"delta_G\": float(family_group[\"delta_G_mean\"].mean()),\n            \"delta_B\": float(family_group[\"delta_B_mean\"].mean()),\n            \"num_samples\": int(len(group)),\n            \"num_families\": num_families,\n        })\n    return pd.DataFrame(rows)\n\n\nFAMILY_SUMMARY = build_family_summary(PER_SAMPLE)\nDATASET_SUMMARY = build_dataset_summary(PER_SAMPLE, FAMILY_SUMMARY)\nFAMILY_SUMMARY.to_csv(FAMILY_SUMMARY_CSV, index=False)\nDATASET_SUMMARY.to_csv(DATASET_SUMMARY_CSV, index=False)\n\nprint(DATASET_SUMMARY.to_string(index=False))\n\n\ndef create_heatmap(family_summary):\n    dataset_order = [\n        DATASET_CONFIGS[key][\"dataset_label\"]\n        for key in DATASETS_TO_RUN\n        if DATASET_CONFIGS[key][\"dataset_label\"] in set(family_summary[\"dataset\"])\n    ]\n    figure, axes = plt.subplots(\n        1,\n        len(dataset_order),\n        figsize=(6.2 * len(dataset_order), 9),\n        squeeze=False,\n    )\n    for axis, dataset_label in zip(axes[0], dataset_order):\n        subset = family_summary[family_summary[\"dataset\"] == dataset_label].copy()\n        subset = subset.sort_values(\"delta_G_mean\", ascending=False)\n        matrix = subset.set_index(\"family\")[[\n            \"delta_R_mean\", \"delta_G_mean\", \"delta_B_mean\"\n        ]]\n        matrix.columns = [\"R\\n(Byte)\", \"G\\n(Entropy)\", \"B\\n(LBP)\"]\n        sns.heatmap(\n            matrix,\n            annot=True,\n            fmt=\".3f\",\n            cmap=\"RdYlBu_r\",\n            center=0,\n            linewidths=0.4,\n            cbar_kws={\"label\": \"Mean true-class confidence drop\", \"shrink\": 0.75},\n            ax=axis,\n        )\n        axis.set_title(f\"{dataset_label}\\nChannel Occlusion Attribution\")\n        axis.set_xlabel(\"Channel\")\n        axis.set_ylabel(\"Malware family\")\n        axis.tick_params(axis=\"y\", labelsize=8)\n    figure.tight_layout()\n    figure.savefig(HEATMAP_PATH, dpi=300, bbox_inches=\"tight\")\n    plt.show()\n    plt.close(figure)\n\n\ncreate_heatmap(FAMILY_SUMMARY)\n\n","metadata":{},"outputs":[],"execution_count":null},{"id":"5f4c746f-0a62-4be3-81ff-9a7ee55eea70","cell_type":"markdown","source":"## Consistency report\n\nThe report compares the newly reproduced results with the rounded values in the\ncurrent paper, identifies which aggregation reproduces the prose, verifies the\nselected family-level claims, and records checkpoint/split provenance.\n\n","metadata":{}},{"id":"91cb0949-5559-4c26-994e-69639cab8f63","cell_type":"code","source":"# ============================================================\n# CELL 5: Consistency report and deliverable validation\n# ============================================================\ndef markdown_table(frame, float_digits=6):\n    formatted = frame.copy()\n    for column in formatted.select_dtypes(include=[np.number]).columns:\n        formatted[column] = formatted[column].map(\n            lambda value: f\"{value:.{float_digits}f}\" if pd.notna(value) else \"NA\"\n        )\n    headers = list(formatted.columns)\n    lines = [\n        \"| \" + \" | \".join(headers) + \" |\",\n        \"| \" + \" | \".join([\"---\"] * len(headers)) + \" |\",\n    ]\n    for row in formatted.astype(str).itertuples(index=False, name=None):\n        lines.append(\"| \" + \" | \".join(row) + \" |\")\n    return \"\\n\".join(lines)\n\n\ndef paper_rounding_match(observed, expected):\n    return round(float(observed) + 1e-12, 2) == float(expected)\n\n\ndef build_consistency_report(per_sample, family_summary, dataset_summary):\n    lines = [\n        \"# Channel occlusion consistency report\",\n        \"\",\n        \"## Audit scope\",\n        \"\",\n        \"No model was trained or fine-tuned. All probabilities, summaries, and the\",\n        \"regenerated heatmap derive from `channel_occlusion_per_sample.csv`.\",\n        \"The occlusion score uses the true-class probability exactly as defined in\",\n        \"the paper.\",\n        \"\",\n        \"## Checkpoint and split provenance\",\n        \"\",\n    ]\n\n    provenance_rows = []\n    for dataset_key, context in RUN_CONTEXTS.items():\n        label = context[\"config\"][\"dataset_label\"]\n        subset = per_sample[per_sample[\"dataset\"] == label]\n        provenance_rows.append({\n            \"dataset\": label,\n            \"checkpoint\": context[\"checkpoint_path\"],\n            \"checkpoint_sha256\": context[\"checkpoint_sha256\"],\n            \"checkpoint_num_classes\": context[\"checkpoint_num_classes\"],\n            \"training_run\": TRAINING_RUN,\n            \"training_seed\": TRAINING_SEED,\n            \"selection_rule\": \"highest validation F1-Macro within training run 1\",\n            \"split_source\": context[\"split_source\"],\n            \"split_sha256\": context[\"split_sha256\"],\n            \"num_test_samples\": len(subset),\n            \"test_accuracy\": float(subset[\"correct_full\"].mean()),\n        })\n    provenance = pd.DataFrame(provenance_rows)\n    lines.extend([markdown_table(provenance), \"\"])\n\n    lines.extend([\n        \"## Dataset-level aggregation\",\n        \"\",\n        markdown_table(dataset_summary),\n        \"\",\n        \"The `sample_weighted_mean` row gives every test sample equal weight; the\",\n        \"`macro_family_mean` row gives every malware family equal weight after\",\n        \"within-family averaging.\",\n        \"\",\n        \"## Comparison with current paper text\",\n        \"\",\n    ])\n\n    comparison_rows = []\n    conclusions = []\n    for dataset_label, expected_values in PAPER_DATASET_VALUES.items():\n        dataset_rows = dataset_summary[dataset_summary[\"dataset\"] == dataset_label]\n        if dataset_rows.empty:\n            conclusions.append(f\"- {dataset_label}: not included in this audit run.\")\n            continue\n        aggregation_matches = {}\n        for aggregation in (\"sample_weighted_mean\", \"macro_family_mean\"):\n            row = dataset_rows[dataset_rows[\"aggregation\"] == aggregation].iloc[0]\n            matches = []\n            for channel in (\"R\", \"G\", \"B\"):\n                column = f\"delta_{channel}\"\n                match = paper_rounding_match(row[column], expected_values[column])\n                matches.append(match)\n                comparison_rows.append({\n                    \"dataset\": dataset_label,\n                    \"aggregation\": aggregation,\n                    \"channel\": channel,\n                    \"observed\": float(row[column]),\n                    \"paper_rounded\": expected_values[column],\n                    \"rounds_to_paper_value\": match,\n                })\n            aggregation_matches[aggregation] = all(matches)\n        if aggregation_matches[\"sample_weighted_mean\"]:\n            conclusions.append(\n                f\"- {dataset_label}: paper text is reproduced by `sample_weighted_mean`.\"\n            )\n        elif aggregation_matches[\"macro_family_mean\"]:\n            conclusions.append(\n                f\"- {dataset_label}: paper text is reproduced by `macro_family_mean`.\"\n            )\n        else:\n            conclusions.append(\n                f\"- {dataset_label}: neither aggregation reproduces every rounded paper value; \"\n                \"replace the paper values with this audit output after checking provenance.\"\n            )\n    comparison = pd.DataFrame(comparison_rows)\n    lines.extend([markdown_table(comparison), \"\", *conclusions, \"\"])\n\n    lines.extend([\n        \"## Test-set size and model-accuracy check\",\n        \"\",\n    ])\n    test_rows = []\n    for dataset_label, expected in PAPER_TEST_METADATA.items():\n        subset = per_sample[per_sample[\"dataset\"] == dataset_label]\n        if subset.empty:\n            continue\n        actual_n = int(len(subset))\n        actual_accuracy = float(subset[\"correct_full\"].mean())\n        test_rows.append({\n            \"dataset\": dataset_label,\n            \"actual_num_samples\": actual_n,\n            \"paper_num_samples\": expected[\"num_samples\"],\n            \"sample_count_match\": actual_n == expected[\"num_samples\"],\n            \"actual_accuracy\": actual_accuracy,\n            \"paper_accuracy\": expected[\"accuracy\"],\n            \"accuracy_rounding_match\": round(actual_accuracy, 4) == expected[\"accuracy\"],\n        })\n    lines.extend([markdown_table(pd.DataFrame(test_rows)), \"\"])\n\n    lines.extend([\n        \"## Family-cell and paper-claim checks\",\n        \"\",\n        \"Every heatmap cell is read directly from the corresponding `delta_*_mean`\",\n        \"column in `channel_occlusion_family_summary.csv`. The checks below compare\",\n        \"selected values explicitly stated in the current paper.\",\n        \"\",\n    ])\n    claim_rows = []\n    for dataset_label, family, metric, expected in PAPER_FAMILY_CLAIMS:\n        matched = family_summary[\n            (family_summary[\"dataset\"] == dataset_label)\n            & (family_summary[\"family\"] == family)\n        ]\n        mean_column = f\"{metric}_mean\"\n        if matched.empty:\n            observed = np.nan\n            match = False\n            status = \"family missing\"\n        else:\n            observed = float(matched.iloc[0][mean_column])\n            match = paper_rounding_match(observed, expected)\n            status = \"match\" if match else \"mismatch\"\n        claim_rows.append({\n            \"dataset\": dataset_label,\n            \"family\": family,\n            \"metric\": metric,\n            \"observed\": observed,\n            \"paper_rounded\": expected,\n            \"status\": status,\n        })\n    lines.extend([markdown_table(pd.DataFrame(claim_rows)), \"\"])\n\n    all_dataset_matches = bool(len(comparison)) and bool(\n        comparison.groupby([\"dataset\", \"aggregation\"])[\"rounds_to_paper_value\"]\n        .all()\n        .groupby(level=0)\n        .any()\n        .all()\n    )\n    all_test_matches = all(\n        row[\"sample_count_match\"] and row[\"accuracy_rounding_match\"]\n        for row in test_rows\n    )\n    all_claim_matches = all(row[\"status\"] == \"match\" for row in claim_rows)\n\n    lines.extend([\n        \"## Recommendation for the camera-ready paper\",\n        \"\",\n    ])\n    if all_dataset_matches and all_test_matches and all_claim_matches:\n        lines.append(\n            \"The current rounded dataset-level values, test metadata, and selected \"\n            \"family claims are reproducible from this audit. Keep the paper values and \"\n            \"state explicitly whether dataset-level means are sample-weighted.\"\n        )\n    else:\n        lines.append(\n            \"At least one current value or provenance check is not reproduced. Use the \"\n            \"new per-sample and summary CSV files as the source of truth, update the \"\n            \"affected paper text/figure, and retain this report with the submission audit.\"\n        )\n    lines.extend([\n        \"\",\n        \"Suggested caption clarification: `Each heatmap cell is a within-family mean;\",\n        \"dataset-level values in the text are reported as sample-weighted means unless\",\n        \"stated otherwise.`\",\n        \"\",\n    ])\n    return \"\\n\".join(lines)\n\n\nREPORT_TEXT = build_consistency_report(\n    PER_SAMPLE,\n    FAMILY_SUMMARY,\n    DATASET_SUMMARY,\n)\nCONSISTENCY_REPORT.write_text(REPORT_TEXT, encoding=\"utf-8\")\n\nRUN_METADATA = {\n    \"task\": \"Task 2 - channel occlusion audit\",\n    \"training_performed\": False,\n    \"model_architecture\": \"EfficientNet-B3\",\n    \"input_type\": \"PPS RGB\",\n    \"image_size\": IMG_SIZE,\n    \"batch_size\": BATCH_SIZE,\n    \"amp_enabled\": AMP_ENABLED,\n    \"occlusion_definition\": \"set one raw channel to zero before normalization\",\n    \"probability_definition\": \"softmax probability of the true class\",\n    \"training_run\": TRAINING_RUN,\n    \"training_seed\": TRAINING_SEED,\n    \"selection_rule\": \"highest validation F1-Macro within training run 1\",\n    \"datasets\": {\n        dataset_key: {\n            \"dataset_label\": context[\"config\"][\"dataset_label\"],\n            \"checkpoint_path\": context[\"checkpoint_path\"],\n            \"checkpoint_sha256\": context[\"checkpoint_sha256\"],\n            \"split_source\": context[\"split_source\"],\n            \"split_sha256\": context[\"split_sha256\"],\n            \"split_snapshot\": context[\"split_snapshot\"],\n            \"num_test_samples\": len(context[\"test_records\"]),\n        }\n        for dataset_key, context in RUN_CONTEXTS.items()\n    },\n}\nRUN_METADATA_JSON.write_text(json.dumps(RUN_METADATA, indent=2) + \"\\n\", encoding=\"utf-8\")\n\nATTR_README.write_text(\n    \"\"\"# Task 2 - channel occlusion audit\n\nRun `task2_channel_occlusion_audit.ipynb` from top to bottom on a Kaggle GPU.\nAttach both original EfficientNet-B3 checkpoints, both PPS image datasets, the\nBIG-2015 labels, and saved split JSON files when available. No training occurs.\n\nThe per-dataset raw CSV files support resume. Set `FORCE_RERUN=True` only for an\nintentional clean rerun with the same provenance. All final summaries and the\nheatmap are generated from `channel_occlusion_per_sample.csv`.\n\"\"\",\n    encoding=\"utf-8\",\n)\n\nrequired_outputs = [\n    PER_SAMPLE_CSV,\n    FAMILY_SUMMARY_CSV,\n    DATASET_SUMMARY_CSV,\n    CONSISTENCY_REPORT,\n    HEATMAP_PATH,\n    RUN_METADATA_JSON,\n    ATTR_README,\n    ATTR_LOG,\n]\nmissing_outputs = [str(path) for path in required_outputs if not path.exists()]\nif missing_outputs:\n    raise RuntimeError(\"Missing Task 2 outputs:\\n- \" + \"\\n- \".join(missing_outputs))\n\nLOGGER.info(\"Task 2 audit complete: %s\", ATTR_DIR)\nprint(REPORT_TEXT)\nprint(\"\\nTask 2 outputs:\")\nfor output in required_outputs:\n    print(f\"- {output}\")\n","metadata":{},"outputs":[],"execution_count":null}]}