IMF / sample.py
fushinguyenex's picture
Upload completed Trace-iMF training run
cfaf85b verified
Raw History Blame Contribute Delete
2.27 kB
"""Sample the exported custom TraceDiT inference bundle without loading pickle."""
import argparse
import json
from pathlib import Path
import numpy as np
from PIL import Image
from safetensors.torch import load_file
import torch
from models import build_backbone
from objectives import build_objective
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--model-dir", required=True)
parser.add_argument("--output", default="samples.png")
parser.add_argument("--n-per-class", type=int, default=4)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--weights", default="model.safetensors")
parser.add_argument("--nfe", type=int, default=1)
args = parser.parse_args()
if args.n_per_class <= 0:
parser.error("n-per-class must be positive")
folder = Path(args.model_dir)
cfg = json.loads((folder / "config.json").read_text())
device = "cuda" if torch.cuda.is_available() else "cpu"
torch.manual_seed(args.seed)
model = build_backbone(cfg).to(device)
model.load_state_dict(load_file(str(folder / args.weights), device=device), strict=True)
objective = build_objective(cfg)
if args.nfe != 1 and not hasattr(objective, "max_nfe"):
parser.error("This model supports only one NFE")
if objective.num_classes is None:
labels, columns = None, 10
else:
columns = objective.num_classes
labels = torch.arange(columns, device=device).repeat(args.n_per_class)
images = objective.sample(model, n_samples=columns * args.n_per_class,
labels=labels, device=device,
**({"nfe": args.nfe} if hasattr(objective, "max_nfe") else {}))
pixels = (images.permute(0, 2, 3, 1).cpu().numpy() * 255).round().astype(np.uint8)
rows = args.n_per_class
if pixels.shape[-1] == 1:
pixels = pixels[..., 0]
grid = np.concatenate([np.concatenate(pixels[row*columns:(row+1)*columns], axis=1)
for row in range(rows)], axis=0)
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
Image.fromarray(grid).save(output)
print(f"{args.nfe}-NFE samples saved: {output}")
if __name__ == "__main__":
main()