mini-unet-colorizer / colorize_eval.py
User-2468's picture
Rename colorize_eval 7.py to colorize_eval.py
caef50b verified
Raw History Blame
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()