{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"7335cd87-d8bd-4fc3-8069-ea1751f8955f","cell_type":"markdown","source":"# ResNet-50 Microsoft BIG-2015 Grayscale Retraining\n\nThis notebook trains three **new reproduction runs** with seeds 43, 44, and 45. It does not recover the lost historical checkpoint bytes.\n\nRequired Kaggle inputs: `vnhtbo/microsoft` containing `train_gray/train_gray/*.png`, the `malware-classification` competition containing `trainLabels.csv`, and `vnhtbo/dataconfig` containing `split_microsoft_rgb.json`. The grayscale and RGB flat image sets share the competition CSV row order, so the fixed RGB split indices are reused. Enable a GPU accelerator. Enable Internet for the first torchvision ImageNet-weight download, or set `IMAGENET_WEIGHTS_PATH_OVERRIDE` to an attached official state dict.\n\nOutputs are written to `/kaggle/working/microsoft_gray_retrain/`, including three checkpoints, the fixed or reconstructed split, per-run metrics, aggregate metrics, SHA-256 hashes, environment metadata, and a ZIP. Keep this reproduction set separate from the historical `dataconfig` Dataset unless its provenance is clearly relabelled.\n","metadata":{}},{"id":"4ec77b51-761d-459e-90fb-6f0d8418fd6d","cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"Train three new ResNet-50 Microsoft BIG-2015 grayscale reproduction runs.\n\nThis code follows the documented PIXEL grayscale protocol. It creates new\nreproduction checkpoints; it does not recover the lost historical weights.\n\"\"\"\n\nimport csv\nimport gc\nimport hashlib\nimport json\nimport os\nfrom pathlib import Path\nimport platform\nimport random\nimport zipfile\n\nimport numpy as np\nimport pandas as pd\nimport psutil\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom PIL import Image\nfrom sklearn.metrics import classification_report, f1_score\nfrom torch.amp import GradScaler, autocast\nfrom torch.utils.data import DataLoader, Dataset, Subset, WeightedRandomSampler, random_split\nfrom torchvision import models, transforms\n\n\n# -----------------------------------------------------------------------------\n# User configuration\n# -----------------------------------------------------------------------------\n\nDATASET_ROOT_OVERRIDE = \"/kaggle/input/datasets/vnhtbo/microsoft/train_gray/train_gray\"\nDATASET_ROOT_CANDIDATES = [\n    \"/kaggle/input/datasets/vnhtbo/microsoft/train_gray/train_gray\",\n]\nMICROSOFT_LABELS_CSV = (\n    \"/kaggle/input/competitions/malware-classification/trainLabels.csv\"\n)\n\n# Leave this as None when Kaggle Internet is enabled. For an offline notebook,\n# attach the official torchvision ResNet-50 IMAGENET1K_V2 state dict and set its\n# mounted path here.\nIMAGENET_WEIGHTS_PATH_OVERRIDE = None\n\n# Set this only to force a particular split file. By default, the notebook\n# reuses the Microsoft RGB split from dataconfig because both flat image sets\n# follow trainLabels.csv row order. If it is unavailable, a seed-42 split is\n# reconstructed and that provenance is recorded explicitly.\nSPLIT_JSON_OVERRIDE = None\nSPLIT_JSON_CANDIDATES = [\n    \"/kaggle/input/datasets/vnhtbo/dataconfig/split_microsoft_rgb.json\",\n]\nALLOW_RECONSTRUCTED_SPLIT = True\n\nOUTPUT_ROOT = Path(\"/kaggle/working/microsoft_gray_retrain\")\nRUNS_TO_TRAIN = [1, 2, 3]\nREUSE_COMPLETED_RUNS = True\n\nEXPECTED_CLASSES = 9\nEXPECTED_SAMPLES = 10868\n\nFREEZE_EPOCHS = 7\nUNFREEZE_EPOCHS = 18\nTOTAL_EPOCHS = FREEZE_EPOCHS + UNFREEZE_EPOCHS\nEARLY_STOP_PATIENCE = 7\nIMG_SIZE = 224\nBATCH_SIZE = 32\nACCUMULATION_STEPS = 2\nLR_HEAD = 1e-3\nLR_FINETUNE = 1e-4\nVAL_RATIO = 0.15\nTEST_RATIO = 0.15\nSPLIT_SEED = 42\nNUM_WORKERS = 2\n\nPAPER_REFERENCE = {\n    \"accuracy_mean\": 0.9796,\n    \"accuracy_std\": 0.0085,\n    \"f1_macro_mean\": 0.9558,\n    \"f1_macro_std\": 0.0175,\n}\n\n\n# -----------------------------------------------------------------------------\n# Environment and deterministic identifiers\n# -----------------------------------------------------------------------------\n\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\nos.environ[\"OMP_NUM_THREADS\"] = \"2\"\nos.environ[\"MKL_NUM_THREADS\"] = \"2\"\n\nif not torch.cuda.is_available():\n    raise RuntimeError(\"A CUDA GPU is required. Enable a Kaggle GPU accelerator.\")\n\nDEVICE = torch.device(\"cuda\")\nOUTPUT_ROOT.mkdir(parents=True, exist_ok=True)\n\n\ndef sha256_file(path):\n    digest = hashlib.sha256()\n    with open(path, \"rb\") as handle:\n        for chunk in iter(lambda: handle.read(1024 * 1024), b\"\"):\n            digest.update(chunk)\n    return digest.hexdigest()\n\n\ndef seed_run(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n\ndef clear_memory():\n    gc.collect()\n    torch.cuda.empty_cache()\n    torch.cuda.synchronize()\n\n\ndef resolve_dataset_root():\n    candidates = []\n    if DATASET_ROOT_OVERRIDE:\n        candidates.append(DATASET_ROOT_OVERRIDE)\n    candidates.extend(DATASET_ROOT_CANDIDATES)\n    diagnostics = []\n    for value in candidates:\n        root = Path(value)\n        if not root.is_dir():\n            diagnostics.append(f\"missing: {root}\")\n            continue\n        image_count = sum(1 for path in root.iterdir() if path.suffix.lower() == \".png\")\n        if image_count != EXPECTED_SAMPLES:\n            diagnostics.append(f\"wrong PNG count ({image_count}): {root}\")\n            continue\n        return root\n    raise FileNotFoundError(\n        \"No valid flat Microsoft grayscale PNG directory was found.\\n\"\n        + \"\\n\".join(diagnostics)\n    )\n\n\ndef resolve_labels_csv():\n    path = Path(MICROSOFT_LABELS_CSV)\n    if not path.is_file():\n        raise FileNotFoundError(\n            f\"Microsoft trainLabels.csv not found: {path}. \"\n            \"Add the malware-classification competition as a Kaggle input.\"\n        )\n    return path\n\n\nclass MicrosoftGrayDataset(Dataset):\n    \"\"\"Flat Id.png directory labelled and ordered by trainLabels.csv.\"\"\"\n\n    def __init__(self, root, labels_csv, transform=None):\n        self.root = Path(root)\n        frame = pd.read_csv(labels_csv, dtype={\"Id\": str})\n        required = {\"Id\", \"Class\"}\n        if not required.issubset(frame.columns):\n            raise ValueError(f\"Labels CSV must contain columns {sorted(required)}\")\n        raw_classes = sorted(frame[\"Class\"].unique().tolist())\n        self.class_to_idx = {\n            raw_class: index for index, raw_class in enumerate(raw_classes)\n        }\n        self.classes = [str(raw_class) for raw_class in raw_classes]\n        existing = {path.name for path in self.root.glob(\"*.png\")}\n        self.samples = []\n        missing = []\n        for row in frame.itertuples(index=False):\n            filename = f\"{row.Id}.png\"\n            if filename not in existing:\n                missing.append(filename)\n                continue\n            self.samples.append(\n                (str(self.root / filename), self.class_to_idx[row.Class])\n            )\n        if missing:\n            raise ValueError(\n                f\"Missing {len(missing)} labelled grayscale PNGs; first: {missing[:5]}\"\n            )\n        self.targets = [label for _, label in self.samples]\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, index):\n        path, label = self.samples[index]\n        with Image.open(path) as source:\n            image = source.convert(\"RGB\")\n        if self.transform:\n            image = self.transform(image)\n        return image, label\n\n\ndef sample_order_sha256(dataset):\n    digest = hashlib.sha256()\n    for path, label in dataset.samples:\n        relative = Path(path).relative_to(dataset.root).as_posix()\n        digest.update(f\"{relative}\\t{label}\\n\".encode(\"utf-8\"))\n    return digest.hexdigest()\n\n\ndef validate_dataset(dataset):\n    if len(dataset.classes) != EXPECTED_CLASSES:\n        raise ValueError(\n            f\"Expected {EXPECTED_CLASSES} classes, found {len(dataset.classes)}\"\n        )\n    if len(dataset) != EXPECTED_SAMPLES:\n        raise ValueError(\n            f\"Expected {EXPECTED_SAMPLES} samples, found {len(dataset)}. \"\n            \"Use the same Microsoft BIG-2015 Dataset version as the paper.\"\n        )\n    counts = np.bincount(dataset.targets, minlength=EXPECTED_CLASSES)\n    print(f\"Dataset root: {dataset.root}\")\n    print(f\"Classes ({len(dataset.classes)}): {dataset.classes}\")\n    print(f\"Samples: {len(dataset)} | min/max family: {counts.min()}/{counts.max()}\")\n\n\ndef validate_split(split, sample_count, order_hash):\n    required = (\"train\", \"val\", \"test\")\n    for key in required:\n        if key not in split or not isinstance(split[key], list):\n            raise ValueError(f\"Split is missing list field: {key}\")\n    if split.get(\"total\") != sample_count:\n        raise ValueError(\n            f\"Split total {split.get('total')} does not match {sample_count} samples\"\n        )\n    if split.get(\"seed\") != SPLIT_SEED:\n        raise ValueError(\n            f\"Split seed {split.get('seed')} does not match required seed {SPLIT_SEED}\"\n        )\n    merged = split[\"train\"] + split[\"val\"] + split[\"test\"]\n    if len(merged) != sample_count or len(set(merged)) != sample_count:\n        raise ValueError(\"Split indices do not form a disjoint full partition\")\n    if min(merged) != 0 or max(merged) != sample_count - 1:\n        raise ValueError(\"Split indices are outside the dataset range\")\n    recorded_hash = split.get(\"sample_order_sha256\")\n    if recorded_hash and recorded_hash != order_hash:\n        raise ValueError(\"Split sample-order fingerprint does not match this Dataset\")\n\n\ndef load_or_create_split(dataset, order_hash):\n    output_path = OUTPUT_ROOT / \"split_microsoft_gray.json\"\n    if SPLIT_JSON_OVERRIDE:\n        source_path = Path(SPLIT_JSON_OVERRIDE)\n        if not source_path.is_file():\n            raise FileNotFoundError(f\"SPLIT_JSON_OVERRIDE not found: {source_path}\")\n        split = json.loads(source_path.read_text(encoding=\"utf-8\"))\n        split_origin = f\"historical override: {source_path}\"\n    elif any(Path(value).is_file() for value in SPLIT_JSON_CANDIDATES):\n        source_path = next(\n            Path(value) for value in SPLIT_JSON_CANDIDATES if Path(value).is_file()\n        )\n        split = json.loads(source_path.read_text(encoding=\"utf-8\"))\n        split_origin = (\n            f\"shared Microsoft RGB split: {source_path}; valid because both flat \"\n            \"datasets are ordered by the same trainLabels.csv rows\"\n        )\n    elif output_path.is_file():\n        split = json.loads(output_path.read_text(encoding=\"utf-8\"))\n        split_origin = f\"existing reproduction output: {output_path}\"\n    else:\n        if not ALLOW_RECONSTRUCTED_SPLIT:\n            raise FileNotFoundError(\n                \"Historical split is unavailable and reconstructed split is disabled\"\n            )\n        n_test = int(len(dataset) * TEST_RATIO)\n        n_val = int(len(dataset) * VAL_RATIO)\n        n_train = len(dataset) - n_val - n_test\n        generator = torch.Generator().manual_seed(SPLIT_SEED)\n        train_set, val_set, test_set = random_split(\n            dataset, [n_train, n_val, n_test], generator=generator\n        )\n        split = {\n            \"train\": train_set.indices,\n            \"val\": val_set.indices,\n            \"test\": test_set.indices,\n            \"seed\": SPLIT_SEED,\n            \"total\": len(dataset),\n            \"sample_order_sha256\": order_hash,\n            \"provenance\": (\n                \"reconstructed from the documented seed; not claimed to be the \"\n                \"lost historical split\"\n            ),\n        }\n        split_origin = \"new reconstruction\"\n    validate_split(split, len(dataset), order_hash)\n    if \"sample_order_sha256\" not in split:\n        split[\"sample_order_sha256\"] = order_hash\n        split[\"sample_order_note\"] = (\n            \"fingerprint added during reproduction; historical file did not bind ordering\"\n        )\n    output_path.write_text(json.dumps(split, indent=2), encoding=\"utf-8\")\n    print(\n        f\"Split: {split_origin} | train={len(split['train'])}, \"\n        f\"val={len(split['val'])}, test={len(split['test'])}\"\n    )\n    print(f\"Split SHA-256: {sha256_file(output_path)}\")\n    return split, output_path, split_origin\n\n\n# -----------------------------------------------------------------------------\n# Exact model and evaluation protocol from the original notebook\n# -----------------------------------------------------------------------------\n\n\ndef get_transforms():\n    base = [\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.Grayscale(num_output_channels=3),\n    ]\n    train_transform = transforms.Compose(\n        base\n        + [\n            transforms.ColorJitter(brightness=0.2, contrast=0.2),\n            transforms.ToTensor(),\n            transforms.Normalize([0.5] * 3, [0.5] * 3),\n        ]\n    )\n    eval_transform = transforms.Compose(\n        base\n        + [\n            transforms.ToTensor(),\n            transforms.Normalize([0.5] * 3, [0.5] * 3),\n        ]\n    )\n    return train_transform, eval_transform\n\n\ndef build_resnet50(num_classes):\n    if IMAGENET_WEIGHTS_PATH_OVERRIDE:\n        weights_path = Path(IMAGENET_WEIGHTS_PATH_OVERRIDE)\n        if not weights_path.is_file():\n            raise FileNotFoundError(\n                f\"IMAGENET_WEIGHTS_PATH_OVERRIDE not found: {weights_path}\"\n            )\n        model = models.resnet50(weights=None)\n        state = torch.load(weights_path, map_location=\"cpu\", weights_only=True)\n        model.load_state_dict(state, strict=True)\n        print(f\"Loaded offline ImageNet weights: {weights_path}\")\n    else:\n        try:\n            model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)\n        except Exception as exc:\n            raise RuntimeError(\n                \"Could not load torchvision ResNet-50 IMAGENET1K_V2 weights. \"\n                \"Enable Kaggle Internet for the first run or set \"\n                \"IMAGENET_WEIGHTS_PATH_OVERRIDE to an attached official state dict.\"\n            ) from exc\n    for parameter in model.parameters():\n        parameter.requires_grad = False\n    model.fc = nn.Sequential(nn.Dropout(0.3), nn.Linear(2048, num_classes))\n    return model.to(DEVICE)\n\n\ndef unfreeze_last_blocks(model):\n    for parameter in model.layer3.parameters():\n        parameter.requires_grad = True\n    for parameter in model.layer4.parameters():\n        parameter.requires_grad = True\n\n\ndef train_one_epoch(model, loader, optimizer, criterion, scaler):\n    model.train()\n    total_loss = 0.0\n    correct = 0\n    total = 0\n    optimizer.zero_grad(set_to_none=True)\n    for step, (images, labels) in enumerate(loader):\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n        with autocast(\"cuda\"):\n            outputs = model(images)\n            loss = criterion(outputs, labels) / ACCUMULATION_STEPS\n        scaler.scale(loss).backward()\n        if (step + 1) % ACCUMULATION_STEPS == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n        total_loss += loss.item() * ACCUMULATION_STEPS\n        correct += (outputs.detach().argmax(1) == labels).sum().item()\n        total += labels.size(0)\n    return total_loss / len(loader), correct / total\n\n\ndef evaluate(model, loader, criterion):\n    model.eval()\n    total_loss = 0.0\n    predictions = []\n    targets = []\n    with torch.inference_mode():\n        for images, labels in loader:\n            images = images.to(DEVICE, non_blocking=True)\n            labels = labels.to(DEVICE, non_blocking=True)\n            with autocast(\"cuda\"):\n                outputs = model(images)\n                total_loss += criterion(outputs, labels).item()\n            predictions.extend(outputs.argmax(1).cpu().tolist())\n            targets.extend(labels.cpu().tolist())\n    accuracy = np.mean(np.asarray(predictions) == np.asarray(targets))\n    macro_f1 = f1_score(targets, predictions, average=\"macro\", zero_division=0)\n    return total_loss / len(loader), float(accuracy), float(macro_f1), predictions, targets\n\n\ndef run_result_path(run):\n    return OUTPUT_ROOT / f\"retrained_result_microsoft_gray_run{run}.json\"\n\n\ndef checkpoint_path(run):\n    return OUTPUT_ROOT / f\"best_microsoft_gray_run{run}.pt\"\n\n\ndef train_run(\n    run,\n    root,\n    labels_csv,\n    classes,\n    all_targets,\n    class_counts,\n    split,\n    split_path,\n    order_hash,\n    train_tf,\n    eval_tf,\n):\n    result_path = run_result_path(run)\n    best_path = checkpoint_path(run)\n    if REUSE_COMPLETED_RUNS and result_path.is_file() and best_path.is_file():\n        existing = json.loads(result_path.read_text(encoding=\"utf-8\"))\n        if existing.get(\"split_sha256\") != sha256_file(split_path):\n            raise ValueError(f\"Completed run {run} belongs to a different split\")\n        if existing.get(\"dataset_sample_order_sha256\") != order_hash:\n            raise ValueError(f\"Completed run {run} belongs to different sample ordering\")\n        print(f\"Reuse completed run {run}: {best_path.name}\")\n        return existing\n\n    run_seed = SPLIT_SEED + run\n    seed_run(run_seed)\n    clear_memory()\n\n    train_full = MicrosoftGrayDataset(root, labels_csv, transform=train_tf)\n    eval_full = MicrosoftGrayDataset(root, labels_csv, transform=eval_tf)\n    train_set = Subset(train_full, split[\"train\"])\n    val_set = Subset(eval_full, split[\"val\"])\n    test_set = Subset(eval_full, split[\"test\"])\n\n    sample_weights = np.asarray(\n        [1.0 / class_counts[all_targets[index]] for index in split[\"train\"]]\n    )\n    sampler = WeightedRandomSampler(\n        torch.as_tensor(sample_weights, dtype=torch.float32),\n        num_samples=len(sample_weights),\n        replacement=True,\n    )\n    loader_args = {\n        \"batch_size\": BATCH_SIZE,\n        \"num_workers\": NUM_WORKERS,\n        \"pin_memory\": True,\n    }\n    train_loader = DataLoader(train_set, sampler=sampler, **loader_args)\n    val_loader = DataLoader(val_set, shuffle=False, **loader_args)\n    test_loader = DataLoader(test_set, shuffle=False, **loader_args)\n\n    weights = 1.0 / class_counts\n    weights = torch.as_tensor(\n        weights / weights.sum() * len(weights), dtype=torch.float32, device=DEVICE\n    )\n    criterion = nn.CrossEntropyLoss(weight=weights)\n\n    model = build_resnet50(len(classes))\n    scaler = GradScaler(\"cuda\")\n    optimizer = optim.AdamW(model.fc.parameters(), lr=LR_HEAD, weight_decay=1e-4)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=FREEZE_EPOCHS)\n\n    best_val_f1 = 0.0\n    best_val_loss = float(\"inf\")\n    best_epoch = 0\n    no_improve = 0\n    phase = 1\n    history = []\n\n    print(f\"\\n{'=' * 72}\\nRUN {run}/3 | seed={run_seed}\\n{'=' * 72}\")\n    for epoch in range(1, TOTAL_EPOCHS + 1):\n        if epoch == FREEZE_EPOCHS + 1:\n            phase = 2\n            unfreeze_last_blocks(model)\n            optimizer = optim.AdamW(\n                [\n                    {\"params\": model.layer3.parameters(), \"lr\": LR_FINETUNE / 2},\n                    {\"params\": model.layer4.parameters(), \"lr\": LR_FINETUNE},\n                    {\"params\": model.fc.parameters(), \"lr\": LR_FINETUNE * 2},\n                ],\n                weight_decay=1e-4,\n            )\n            scheduler = optim.lr_scheduler.CosineAnnealingLR(\n                optimizer, T_max=UNFREEZE_EPOCHS\n            )\n            scaler = GradScaler(\"cuda\")\n            no_improve = 0\n\n        train_loss, train_accuracy = train_one_epoch(\n            model, train_loader, optimizer, criterion, scaler\n        )\n        val_loss, val_accuracy, val_f1, _, _ = evaluate(\n            model, val_loader, criterion\n        )\n        scheduler.step()\n        history.append(\n            {\n                \"epoch\": epoch,\n                \"phase\": phase,\n                \"train_loss\": round(float(train_loss), 6),\n                \"train_accuracy\": round(float(train_accuracy), 6),\n                \"val_loss\": round(float(val_loss), 6),\n                \"val_accuracy\": round(float(val_accuracy), 6),\n                \"val_f1_macro\": round(float(val_f1), 6),\n            }\n        )\n\n        if val_f1 > best_val_f1:\n            best_val_f1 = val_f1\n            best_val_loss = val_loss\n            best_epoch = epoch\n            no_improve = 0\n            torch.save(model.state_dict(), best_path)\n            marker = \"saved\"\n        else:\n            no_improve += 1\n            marker = \"\"\n\n        ram = psutil.virtual_memory()\n        print(\n            f\"epoch={epoch:02d} phase={phase} train_loss={train_loss:.4f} \"\n            f\"train_acc={train_accuracy:.4f} val_loss={val_loss:.4f} \"\n            f\"val_acc={val_accuracy:.4f} val_f1={val_f1:.4f} \"\n            f\"patience={no_improve}/{EARLY_STOP_PATIENCE} \"\n            f\"RAM={ram.used / 1e9:.1f}GB {marker}\"\n        )\n        if no_improve >= EARLY_STOP_PATIENCE:\n            print(\"Early stopping on validation F1-Macro\")\n            break\n\n    if not best_path.is_file():\n        raise RuntimeError(f\"No best checkpoint was saved for run {run}\")\n    state = torch.load(best_path, map_location=DEVICE, weights_only=True)\n    model.load_state_dict(state, strict=True)\n    _, test_accuracy, test_f1, predictions, targets = evaluate(\n        model, test_loader, criterion\n    )\n    print(f\"Run {run} test accuracy={test_accuracy:.6f}, F1-Macro={test_f1:.6f}\")\n    print(\n        classification_report(\n            targets,\n            predictions,\n            labels=list(range(len(classes))),\n            target_names=classes,\n            zero_division=0,\n        )\n    )\n\n    result = {\n        \"artifact_type\": \"new reproduction run; not a recovered historical checkpoint\",\n        \"dataset\": \"Microsoft BIG-2015 grayscale\",\n        \"labels_csv_sha256\": sha256_file(labels_csv),\n        \"run\": run,\n        \"training_seed\": run_seed,\n        \"split_seed\": SPLIT_SEED,\n        \"split_sha256\": sha256_file(split_path),\n        \"dataset_sample_order_sha256\": order_hash,\n        \"best_epoch\": best_epoch,\n        \"best_val_f1_macro\": float(best_val_f1),\n        \"best_val_loss\": float(best_val_loss),\n        \"test_accuracy\": test_accuracy,\n        \"test_f1_macro\": test_f1,\n        \"checkpoint\": best_path.name,\n        \"checkpoint_sha256\": sha256_file(best_path),\n        \"history\": history,\n    }\n    result_path.write_text(json.dumps(result, indent=2), encoding=\"utf-8\")\n\n    del model, train_loader, val_loader, test_loader\n    clear_memory()\n    return result\n\n\ndef write_environment(\n    dataset_root, labels_csv, sample_hash, split_path, split_origin\n):\n    environment = {\n        \"python\": platform.python_version(),\n        \"torch\": torch.__version__,\n        \"torchvision\": __import__(\"torchvision\").__version__,\n        \"cuda_runtime\": torch.version.cuda,\n        \"gpu\": torch.cuda.get_device_name(0),\n        \"dataset_root\": str(dataset_root),\n        \"labels_csv\": str(labels_csv),\n        \"labels_csv_sha256\": sha256_file(labels_csv),\n        \"dataset_sample_order_sha256\": sample_hash,\n        \"split_path\": str(split_path),\n        \"split_sha256\": sha256_file(split_path),\n        \"split_origin\": split_origin,\n        \"protocol\": {\n            \"image_size\": IMG_SIZE,\n            \"batch_size\": BATCH_SIZE,\n            \"accumulation_steps\": ACCUMULATION_STEPS,\n            \"freeze_epochs\": FREEZE_EPOCHS,\n            \"unfreeze_epochs\": UNFREEZE_EPOCHS,\n            \"early_stop_patience\": EARLY_STOP_PATIENCE,\n            \"lr_head\": LR_HEAD,\n            \"lr_finetune\": LR_FINETUNE,\n            \"normalization_mean\": [0.5, 0.5, 0.5],\n            \"normalization_std\": [0.5, 0.5, 0.5],\n        },\n    }\n    path = OUTPUT_ROOT / \"retrained_microsoft_gray_environment.json\"\n    path.write_text(json.dumps(environment, indent=2), encoding=\"utf-8\")\n\n\ndef write_test_manifest(dataset, split):\n    path = OUTPUT_ROOT / \"retrained_microsoft_gray_test_samples.csv\"\n    with path.open(\"w\", newline=\"\", encoding=\"utf-8\") as handle:\n        writer = csv.DictWriter(\n            handle,\n            fieldnames=[\"split_index\", \"dataset_index\", \"relative_path\", \"class_index\"],\n        )\n        writer.writeheader()\n        for split_index, dataset_index in enumerate(split[\"test\"]):\n            sample_path, class_index = dataset.samples[dataset_index]\n            writer.writerow(\n                {\n                    \"split_index\": split_index,\n                    \"dataset_index\": dataset_index,\n                    \"relative_path\": Path(sample_path).relative_to(dataset.root).as_posix(),\n                    \"class_index\": class_index,\n                }\n            )\n    print(f\"Test manifest: {path} | SHA-256: {sha256_file(path)}\")\n\n\ndef aggregate_results(results, split_path, sample_hash):\n    ordered = sorted(results, key=lambda item: item[\"run\"])\n    accuracies = [item[\"test_accuracy\"] for item in ordered]\n    f1_values = [item[\"test_f1_macro\"] for item in ordered]\n    summary = {\n        \"artifact_type\": \"new three-run reproduction; not recovered historical runs\",\n        \"description\": \"Microsoft BIG-2015 Grayscale ResNet-50 reproduction\",\n        \"num_runs\": len(ordered),\n        \"runs\": {\n            f\"run_{item['run']}\": {\n                \"training_seed\": item[\"training_seed\"],\n                \"test_accuracy\": item[\"test_accuracy\"],\n                \"test_f1_macro\": item[\"test_f1_macro\"],\n                \"checkpoint\": item[\"checkpoint\"],\n                \"checkpoint_sha256\": item[\"checkpoint_sha256\"],\n            }\n            for item in ordered\n        },\n        \"mean_accuracy\": float(np.mean(accuracies)),\n        \"std_accuracy\": float(np.std(accuracies, ddof=0)),\n        \"mean_f1_macro\": float(np.mean(f1_values)),\n        \"std_f1_macro\": float(np.std(f1_values, ddof=0)),\n        \"split\": split_path.name,\n        \"split_sha256\": sha256_file(split_path),\n        \"dataset_sample_order_sha256\": sample_hash,\n        \"accepted_paper_reference_only\": PAPER_REFERENCE,\n    }\n    filename = (\n        \"retrained_results_microsoft_gray.json\"\n        if len(ordered) == 3\n        else \"retrained_results_microsoft_gray_partial.json\"\n    )\n    output_path = OUTPUT_ROOT / filename\n    output_path.write_text(json.dumps(summary, indent=2), encoding=\"utf-8\")\n    print(\"\\nReproduction aggregate\")\n    print(json.dumps({key: summary[key] for key in (\n        \"num_runs\", \"mean_accuracy\", \"std_accuracy\", \"mean_f1_macro\", \"std_f1_macro\"\n    )}, indent=2))\n    print(\"Accepted-paper values above are reference-only, not overwritten.\")\n    return output_path\n\n\ndef write_checksums_and_zip():\n    checksum_path = OUTPUT_ROOT / \"SHA256SUMS\"\n    files = sorted(\n        path for path in OUTPUT_ROOT.iterdir()\n        if path.is_file() and path.name not in {\"SHA256SUMS\", \"microsoft_gray_retrain.zip\"}\n    )\n    checksum_path.write_text(\n        \"\".join(f\"{sha256_file(path)}  {path.name}\\n\" for path in files),\n        encoding=\"utf-8\",\n    )\n    zip_path = OUTPUT_ROOT / \"microsoft_gray_retrain.zip\"\n    with zipfile.ZipFile(zip_path, \"w\", compression=zipfile.ZIP_DEFLATED) as archive:\n        for path in sorted(OUTPUT_ROOT.iterdir()):\n            if path.is_file() and path != zip_path:\n                archive.write(path, arcname=path.name)\n    print(f\"Artifact ZIP: {zip_path} ({zip_path.stat().st_size / 1024**2:.1f} MiB)\")\n    print(f\"ZIP SHA-256: {sha256_file(zip_path)}\")\n\n\ndef main():\n    invalid_runs = [run for run in RUNS_TO_TRAIN if run not in (1, 2, 3)]\n    if invalid_runs or len(set(RUNS_TO_TRAIN)) != len(RUNS_TO_TRAIN):\n        raise ValueError(f\"RUNS_TO_TRAIN must be a unique subset of [1,2,3]: {RUNS_TO_TRAIN}\")\n\n    print(f\"Device: {DEVICE} | GPU: {torch.cuda.get_device_name(0)}\")\n    print(\"This creates new reproduction checkpoints, not recovered historical files.\")\n    root = resolve_dataset_root()\n    labels_csv = resolve_labels_csv()\n    train_tf, eval_tf = get_transforms()\n    metadata_dataset = MicrosoftGrayDataset(root, labels_csv, transform=None)\n    validate_dataset(metadata_dataset)\n    order_hash = sample_order_sha256(metadata_dataset)\n    print(f\"Sample-order SHA-256: {order_hash}\")\n    split, split_path, split_origin = load_or_create_split(metadata_dataset, order_hash)\n\n    all_targets = list(metadata_dataset.targets)\n    classes = list(metadata_dataset.classes)\n    class_counts = np.bincount(all_targets, minlength=len(classes))\n    write_environment(root, labels_csv, order_hash, split_path, split_origin)\n    write_test_manifest(metadata_dataset, split)\n\n    results = []\n    for run in RUNS_TO_TRAIN:\n        results.append(\n            train_run(\n                run,\n                root,\n                labels_csv,\n                classes,\n                all_targets,\n                class_counts,\n                split,\n                split_path,\n                order_hash,\n                train_tf,\n                eval_tf,\n            )\n        )\n    aggregate_results(results, split_path, order_hash)\n    write_checksums_and_zip()\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T17:04:46.860195Z","iopub.execute_input":"2026-07-20T17:04:46.860596Z"}},"outputs":[],"execution_count":null}]}