# /// script # requires-python = ">=3.10" # dependencies = [ # "torch>=2.3", # "torchvision>=0.18", # "datasets>=2.19", # "huggingface_hub>=0.24", # "safetensors>=0.4", # "scikit-image>=0.22", # "pillow>=10.0", # "numpy", # ] # /// """ Evaluate a SmallUNetColorizer checkpoint with a large, class-stratified sample grid -- unlike the small fixed grid saved during training (which, before a fix, was drawn from a single class), this samples N images from *every* class in the dataset, labeled, so weaknesses specific to one subject category (e.g. reflective electronics) are visible instead of hidden. Runs entirely on CPU -- inference-only on a ~3.6M param model, no GPU needed. Designed for `hf jobs uv run --flavor cpu-basic`; since a Jobs container is destroyed when it finishes, the resulting grid is pushed to a Hub repo rather than saved somewhere you could copy it from afterward. hf jobs uv run --flavor cpu-basic --timeout 1h -s HF_TOKEN \\ colorize_eval.py \\ --model User-2468/mini-unet-colorizer \\ --per-class 5 --temperature 0.38 \\ --push-to-hub --hub-model-id User-2468/mini-unet-colorizer """ import argparse from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from huggingface_hub import HfApi, PyTorchModelHubMixin from PIL import Image, ImageDraw, ImageFont from skimage.color import lab2rgb, rgb2lab # -------------------------------------------------------------------------- # Model (identical to colorize_train.py -- kept self-contained since Jobs # scripts don't have access to sibling files). # -------------------------------------------------------------------------- def double_conv(in_ch, out_ch): return nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) class DilatedContextBlock(nn.Module): def __init__(self, channels, mid_ch=96, dilations=(2, 4, 8)): super().__init__() self.proj_in = nn.Sequential( nn.Conv2d(channels, mid_ch, 1, bias=False), nn.BatchNorm2d(mid_ch), nn.ReLU(inplace=True), ) layers = [] for d in dilations: layers += [ nn.Conv2d(mid_ch, mid_ch, 3, padding=d, dilation=d, bias=False), nn.BatchNorm2d(mid_ch), nn.ReLU(inplace=True), ] self.dilated = nn.Sequential(*layers) self.proj_out = nn.Sequential( nn.Conv2d(mid_ch, channels, 1, bias=False), nn.BatchNorm2d(channels), ) self.relu = nn.ReLU(inplace=True) def forward(self, x): y = self.proj_in(x) y = self.dilated(y) y = self.proj_out(y) return self.relu(x + y) class SmallUNetColorizer( nn.Module, PyTorchModelHubMixin, pipeline_tag="image-to-image", license="apache-2.0", tags=["colorization", "unet", "image-to-image", "classification"], ): def __init__(self, bin_centers, in_ch: int = 1, base: int = 44, context_mid_ch: int = 96, context_dilations=(2, 4, 8)): super().__init__() self.in_ch, self.base = in_ch, base num_bins = len(bin_centers) self.num_bins = num_bins self.register_buffer("bin_centers", torch.tensor(bin_centers, dtype=torch.float32)) self.enc1 = double_conv(in_ch, base) self.enc2 = double_conv(base, base * 2) self.enc3 = double_conv(base * 2, base * 4) self.enc4 = double_conv(base * 4, base * 8) self.pool = nn.MaxPool2d(2) self.context = DilatedContextBlock(base * 8, mid_ch=context_mid_ch, dilations=tuple(context_dilations)) self.up3 = nn.ConvTranspose2d(base * 8, base * 4, 2, stride=2) self.dec3 = double_conv(base * 8, base * 4) self.up2 = nn.ConvTranspose2d(base * 4, base * 2, 2, stride=2) self.dec2 = double_conv(base * 4, base * 2) self.up1 = nn.ConvTranspose2d(base * 2, base, 2, stride=2) self.dec1 = double_conv(base * 2, base) self.out_conv = nn.Conv2d(base, num_bins, 1) def forward(self, x): e1 = self.enc1(x) e2 = self.enc2(self.pool(e1)) e3 = self.enc3(self.pool(e2)) e4 = self.context(self.enc4(self.pool(e3))) d3 = self.dec3(torch.cat([self.up3(e4), e3], dim=1)) d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1)) d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1)) return self.out_conv(d1) def decode(self, logits, temperature: float = 0.38): logp = F.log_softmax(logits, dim=1) probs_t = F.softmax(logp / temperature, dim=1) return torch.einsum("bqhw,qc->bchw", probs_t, self.bin_centers) # -------------------------------------------------------------------------- # Dataset loading (identical fallback logic to colorize_train.py, for # datasets like frgfm/imagenette that ship a legacy loading script). # -------------------------------------------------------------------------- # Some Hub repos for imagenette label classes by WordNet synset ID rather # than a readable name (e.g. johnowhitaker/imagenette2-320). Standard, # well-documented imagenette synset -> name mapping so captions stay # readable regardless of which repo's labeling convention is in use. SYNSET_TO_NAME = { "n01440764": "tench", "n02102040": "English springer", "n02979186": "cassette player", "n03000684": "chain saw", "n03028079": "church", "n03394916": "French horn", "n03417042": "garbage truck", "n03425413": "gas pump", "n03445777": "golf ball", "n03888257": "parachute", } def get_dataset_splits(dataset_name, dataset_config): from datasets import Image as HFImage from datasets import load_dataset print(f"Loading {dataset_name}" + (f" ({dataset_config})" if dataset_config else "") + " ...") try: ds = load_dataset(dataset_name, dataset_config) if dataset_config else load_dataset(dataset_name) except RuntimeError as e: if "Dataset scripts are no longer supported" not in str(e): raise print("Dataset uses a legacy loading script; loading its Parquet " f"files directly from refs/convert/parquet/{dataset_config}/ ...") base = f"hf://datasets/{dataset_name}@refs%2Fconvert%2Fparquet/{dataset_config}" try: ds = load_dataset("parquet", data_files={ "train": f"{base}/train/*.parquet", "validation": f"{base}/validation/*.parquet", }) except ValueError as e2: if "all the data_files are invalid" not in str(e2): raise raise RuntimeError( f"No Parquet mirror found for {dataset_name} config '{dataset_config}' at " f"refs/convert/parquet right now -- this can happen even for a config that " f"worked before (transient Hub-side issue with this repo, not a bug here). " f"Consider a dataset stored natively as Parquet instead (no legacy-script " f"fallback needed at all), e.g. --dataset johnowhitaker/imagenette2-320." ) from e2 ds = ds.cast_column("image", HFImage()) train_split = ds["train"] if "validation" in ds: eval_split = ds["validation"] elif "test" in ds: eval_split = ds["test"] else: # Some repos (e.g. johnowhitaker/imagenette2-320) ship a single # unsplit "train" -- fine for a qualitative visual check like this # script does, just note it's not a strictly held-out set. print(f"Note: {dataset_name} has no validation/test split -- sampling " f"from 'train' instead. Fine for a visual check, but these may " f"overlap with what the model was actually trained on.") eval_split = train_split return train_split, eval_split # -------------------------------------------------------------------------- # Evaluation # -------------------------------------------------------------------------- def colorize_one(model, img, size, temperature, grayscale_chroma_threshold=3.0): """Returns (gray_rgb, pred_rgb, true_rgb, is_grayscale_source).""" img = img.convert("RGB").resize((size, size)) arr = np.asarray(img).astype(np.float32) / 255.0 lab = rgb2lab(arr).astype(np.float32) L = torch.from_numpy(lab[:, :, 0:1] / 50.0 - 1.0).permute(2, 0, 1)[None] with torch.no_grad(): logits = model(L) ab_pred = model.decode(logits, temperature=temperature)[0].permute(1, 2, 0).numpy() L_raw = lab[:, :, 0] ab_true = lab[:, :, 1:3] is_grayscale = float(np.sqrt((ab_true ** 2).sum(-1)).mean()) < grayscale_chroma_threshold def to_rgb(ab): lab_full = np.concatenate([L_raw[:, :, None], ab], axis=-1) return (np.clip(lab2rgb(lab_full), 0, 1) * 255).astype(np.uint8) gray_rgb = (np.stack([L_raw] * 3, axis=-1) / 100.0 * 255).astype(np.uint8) return gray_rgb, to_rgb(ab_pred), to_rgb(ab_true), is_grayscale def build_stratified_sample(val_split, per_class, seed): """Returns list of (index, class_name), `per_class` images from every class in the dataset -- not just whichever class happens to sort first.""" labels = val_split["label"] raw_names = val_split.features["label"].names class_names = [SYNSET_TO_NAME.get(n, n) for n in raw_names] # readable, if known rng = np.random.default_rng(seed) by_class = {c: [] for c in range(len(class_names))} for i, lbl in enumerate(labels): by_class[lbl].append(i) sample = [] for c, indices in by_class.items(): chosen = rng.choice(indices, size=min(per_class, len(indices)), replace=False) sample.extend((int(i), class_names[c]) for i in chosen) return sample def build_grid_image(rows_by_class, thumb_size, per_class): """rows_by_class: dict class_name -> list of (gray, pred, true, is_grayscale_source).""" font = ImageFont.load_default() header_h = 24 pad = 2 row_w = thumb_size * 3 + pad * 2 class_block_h = header_h + per_class * (thumb_size + pad) class_names = list(rows_by_class.keys()) total_h = class_block_h * len(class_names) canvas = Image.new("RGB", (row_w, total_h), "white") draw = ImageDraw.Draw(canvas) y = 0 for name in class_names: rows = rows_by_class[name] n_gray_source = sum(1 for r in rows if r[3]) header_text = f"{name} (gray | prediction | ground truth)" if n_gray_source: header_text += f" -- {n_gray_source}/{len(rows)} grayscale-source" draw.rectangle([0, y, row_w, y + header_h], fill=(30, 30, 30)) draw.text((6, y + 5), header_text, fill="white", font=font) y += header_h for gray, pred, true, is_gray_source in rows: x = 0 for arr in (gray, pred, true): im = Image.fromarray(arr).resize((thumb_size, thumb_size)) canvas.paste(im, (x, y)) x += thumb_size + pad if is_gray_source: # Tag the ground-truth thumbnail so it's obvious *which* row # had nothing to colorize, not just that the class has some. tag_x = thumb_size * 2 + pad * 2 draw.rectangle([tag_x, y, tag_x + thumb_size, y + 14], fill=(200, 30, 30)) draw.text((tag_x + 3, y + 2), "B&W source", fill="white", font=font) y += thumb_size + pad return canvas def main(): p = argparse.ArgumentParser(description=__doc__) p.add_argument("--model", required=True, help="Hub model id or local checkpoint path") p.add_argument("--dataset", default="johnowhitaker/imagenette2-320", help="Default is stored natively as Parquet (no legacy-script fallback " "needed), which has been more reliable than frgfm/imagenette's " "auto-converted mirror. It has no separate validation split (just " "'train'), so this samples from the full set rather than a held-out " "portion -- fine for a visual check. frgfm/imagenette is still an " "option (--dataset-config 160px or 320px) if its mirror comes back.") p.add_argument("--dataset-config", default=None, help="Config name, if the dataset needs one (frgfm/imagenette does: " "'160px' or '320px'). Leave unset for single-config datasets like " "the johnowhitaker default.") p.add_argument("--image-size", type=int, default=256, help="Resolution used for inference") p.add_argument("--thumb-size", type=int, default=128, help="Display size in the grid") p.add_argument("--per-class", type=int, default=5, help="Sampled images per class") p.add_argument("--temperature", type=float, default=0.38) p.add_argument("--grayscale-chroma-threshold", type=float, default=3.0, help="Mean Lab chroma below which a ground-truth image is flagged as " "genuinely grayscale-source and tagged in the grid, rather than a " "real colorization miss.") p.add_argument("--seed", type=int, default=0) p.add_argument("--output-dir", default="./eval") p.add_argument("--output-name", default="eval_grid.png") p.add_argument("--push-to-hub", action="store_true") p.add_argument("--hub-model-id", default=None, help="Repo to push the grid to (defaults to --model). Since a Jobs " "container is destroyed on completion, --push-to-hub is the only " "way to get the result back out when running as a Job.") args = p.parse_args() print(f"Loading {args.model} ...") model = SmallUNetColorizer.from_pretrained(args.model).eval() print(f"({model.num_bins} color bins)") _, val_split = get_dataset_splits(args.dataset, args.dataset_config) sample = build_stratified_sample(val_split, args.per_class, args.seed) class_names = [SYNSET_TO_NAME.get(n, n) for n in val_split.features["label"].names] print(f"Evaluating {len(sample)} images across {len(class_names)} classes " f"({args.per_class} per class) ...") rows_by_class = {name: [] for name in class_names} n_grayscale_total = 0 for n, (idx, class_name) in enumerate(sample): img = val_split[idx]["image"] gray, pred, true, is_gray_source = colorize_one( model, img, args.image_size, args.temperature, args.grayscale_chroma_threshold) rows_by_class[class_name].append((gray, pred, true, is_gray_source)) n_grayscale_total += is_gray_source if (n + 1) % 10 == 0 or n + 1 == len(sample): print(f" {n + 1}/{len(sample)}") print(f"{n_grayscale_total}/{len(sample)} sampled images are grayscale-source " f"(tagged in red in the grid) -- these have no color to recover, so weak-looking " f"results there aren't a model failure.") grid = build_grid_image(rows_by_class, args.thumb_size, args.per_class) out_dir = Path(args.output_dir) out_dir.mkdir(parents=True, exist_ok=True) out_path = out_dir / args.output_name grid.save(out_path) print(f"Saved grid ({grid.width}x{grid.height}) to {out_path}") if args.push_to_hub: hub_model_id = args.hub_model_id or args.model try: HfApi().upload_file(path_or_fileobj=str(out_path), path_in_repo=args.output_name, repo_id=hub_model_id) print(f"Pushed to https://huggingface.co/{hub_model_id}/blob/main/{args.output_name}") except Exception as e: print(f"WARNING: push failed ({e}). Grid was still generated at {out_path}, " f"but that path won't survive a Jobs container exiting.") if __name__ == "__main__": main()