Download colorize_eval.py from User-2468/mini-unet-colorizer: direct link, hf CLI and curl.
- Browser
- Download file 16.2 kB
-
https://huggingface.co/User-2468/mini-unet-colorizer/resolve/078af3e22d82482c922c8679cca466523f87acf8/colorize_eval.py
- Command line
-
hf download hf://User-2468/mini-unet-colorizer@078af3e22d82482c922c8679cca466523f87acf8/colorize_eval.py
-
curl -L -o colorize_eval.py https://huggingface.co/User-2468/mini-unet-colorizer/resolve/078af3e22d82482c922c8679cca466523f87acf8/colorize_eval.py
16.2 kB
| # /// 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() | |