File size: 8,336 Bytes
b0e01a5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 | #!/usr/bin/env python3
"""Extract a small batch of Guillaumin 2014 images to disk for sweep tests.
Reads `data/gtsegs_ijcv.mat` and writes the first N images (or a sampled set)
as JPEGs named after their canonical ImageNet ids (`nXXXXXXXX_NNNN.JPEG`).
This makes the sweep able to look up GT masks by filename stem and compute
FER / mIoU / mAP end-to-end.
Also generates `metadata.json` mapping each filename to its ImageNet-1k class
index (when the synset is in ILSVRC-2012; null otherwise). The mapping format
is consumed by `experiments.run_attack_sweep.load_sample_ground_truth_map`.
Background: Guillaumin 2014 samples 445 synsets from the FULL ImageNet pool
(Deng et al. 2009, "IN-21k"), of which only 95 fall in ImageNet-1k. So roughly
22% of the 4,276 images get a valid `ground_truth` integer; the rest have
`imagenet_id: null` and the sweep falls back to the model's clean prediction
for label-dependent metrics (ASR, confidence drop, top-k drop). See
`wiki/pesquisa-vit/metodologia-fer.md` for the full protocol rationale.
Usage:
python scripts/export_guillaumin_samples.py --n 5 \
--out data/guillaumin_samples
python scripts/export_guillaumin_samples.py --all \
--out data/guillaumin_samples_full
"""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
from typing import Dict, Optional
def _find_project_root() -> Path:
cur = Path(__file__).resolve().parent
for parent in [cur, *cur.parents]:
if (parent / "requirements.txt").exists():
return parent
raise RuntimeError(f"project root not found from {__file__}")
PROJECT_ROOT = _find_project_root()
sys.path.insert(0, str(PROJECT_ROOT))
import numpy as np # noqa: E402
from PIL import Image # noqa: E402
from utils.foreground import GTMaskLoader, _decode_id # noqa: E402
def _build_in1k_synset_index() -> Dict[str, int]:
"""Return mapping `synset (e.g. 'n01440764') → IN-1k class index 0..999`.
Uses timm's canonical IN-1k synset list. Order matches the standard
`synset_to_idx` mapping used by ImageNet-1k pretrained classifiers.
"""
try:
from timm.data.imagenet_info import ImageNetInfo
except ImportError as exc:
raise ImportError(
"timm is required for ImageNet-1k synset mapping. "
"Install: pip install timm"
) from exc
info = ImageNetInfo()
synsets = list(info.label_descriptions(as_dict=True).keys())
return {syn: idx for idx, syn in enumerate(synsets)}
def _synset_for_image_id(image_id: str) -> str:
"""`'n01322343_1025'` → `'n01322343'`."""
return image_id.split("_", 1)[0]
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--n",
type=int,
default=5,
help="Number of images to export (ignored if --all is set).",
)
parser.add_argument(
"--all",
action="store_true",
help="Export every image in the .mat file (overrides --n / --start).",
)
parser.add_argument(
"--in1k-only",
action="store_true",
help="Only export images whose synset is in ImageNet-1k. Useful for "
"smoke tests where ground_truth coverage matters.",
)
parser.add_argument(
"--diverse-synsets",
action="store_true",
help="Pick at most one image per synset (round-robin). Useful with "
"--in1k-only for a varied smoke set instead of multiple images "
"of the same class.",
)
parser.add_argument(
"--out",
type=Path,
default=PROJECT_ROOT / "data" / "guillaumin_samples",
help="Output directory (created if missing).",
)
parser.add_argument(
"--start",
type=int,
default=0,
help="Start index in the .mat file (default 0; ignored with --all).",
)
args = parser.parse_args()
loader = GTMaskLoader()
if not loader.available:
print(f"ERROR: GT mask file not found at {loader.mat_path}")
return 1
args.out.mkdir(parents=True, exist_ok=True)
loader._ensure_open()
val = loader._h5["value"]
total = int(np.asarray(val["n"]).squeeze())
in1k_map = _build_in1k_synset_index()
print(f"IN-1k synset table loaded ({len(in1k_map)} synsets)")
# Build the iteration order: indices into the .mat file.
if args.all:
candidate_indices = range(total)
else:
if args.start >= total:
print(f"ERROR: --start {args.start} >= total images {total}")
return 1
candidate_indices = range(args.start, total)
target_count = total if args.all else args.n
if args.in1k_only:
target_count = min(target_count, total)
print(f"Output dir: {args.out}")
print(f"Filter: in1k_only={args.in1k_only}, target_count={target_count}")
downloaded_files: list[str] = []
suggested_classes: list[Dict[str, Optional[object]]] = []
in1k_count = 0
skipped_shape = 0
skipped_ood = 0
skipped_duplicate_synset = 0
seen_synsets: set[str] = set()
for i in candidate_indices:
if len(downloaded_files) >= target_count:
break
ref = val["id"][i, 0]
image_id = _decode_id(loader._h5[ref][()])
synset = _synset_for_image_id(image_id)
in1k_idx = in1k_map.get(synset)
if args.in1k_only and in1k_idx is None:
skipped_ood += 1
continue
if args.diverse_synsets and synset in seen_synsets:
skipped_duplicate_synset += 1
continue
img = loader._image_at(i)
if img.ndim != 3 or img.shape[2] != 3:
print(f" [{i}] {image_id} — unexpected image shape {img.shape}, skipping")
skipped_shape += 1
continue
out_path = args.out / f"{image_id}.JPEG"
Image.fromarray(img).save(out_path, format="JPEG", quality=95)
downloaded_files.append(out_path.name)
seen_synsets.add(synset)
if in1k_idx is not None:
in1k_count += 1
suggested_classes.append({
"synset": synset,
"imagenet_id": in1k_idx if in1k_idx is not None else None,
"in_imagenet_1k": in1k_idx is not None,
})
h, w = img.shape[:2]
in1k_marker = "[IN-1k]" if in1k_idx is not None else "[OOD]"
print(f" [{i}] {image_id} ({h}x{w}) -> {out_path.name} {in1k_marker}")
metadata = {
"description": (
"Guillaumin 2014 ImageNet-Segmentation samples extracted from "
"data/gtsegs_ijcv.mat. Each image is named after its canonical "
"ImageNet id (synset + sample number). suggested_classes has the "
"same length and ordering as downloaded_files; "
"imagenet_id is the IN-1k class index (0..999) when the synset "
"is in ILSVRC-2012, and null otherwise. The sweep loader skips "
"null entries and falls back to clean prediction as reference."
),
"source": "data/gtsegs_ijcv.mat (Guillaumin et al. 2014, IJCV)",
"imagenet_namespace": "ILSVRC-2012 (1000 classes); synsets outside "
"this set come from the broader ImageNet-21k pool.",
"exported": len(downloaded_files),
"in_imagenet_1k": in1k_count,
"out_of_distribution": len(downloaded_files) - in1k_count,
"downloaded_files": downloaded_files,
"suggested_classes": suggested_classes,
}
metadata_path = args.out / "metadata.json"
with open(metadata_path, "w", encoding="utf-8") as f:
json.dump(metadata, f, indent=2, ensure_ascii=False)
print(f"\nDone:")
print(f" exported: {len(downloaded_files)} files in {args.out}")
print(f" IN-1k coverage: {in1k_count}/{len(downloaded_files)} "
f"({100 * in1k_count / max(len(downloaded_files), 1):.1f}%)")
if skipped_shape:
print(f" skipped (bad shape): {skipped_shape}")
if skipped_ood:
print(f" skipped (OOD synset, --in1k-only): {skipped_ood}")
if skipped_duplicate_synset:
print(f" skipped (duplicate synset, --diverse-synsets): {skipped_duplicate_synset}")
print(f" metadata: {metadata_path}")
return 0
if __name__ == "__main__":
sys.exit(main())
|