Download sample.py from fushinguyenex/IMF: direct link, hf CLI and curl.
- Browser
- Download file 2.27 kB
-
https://huggingface.co/fushinguyenex/IMF/resolve/main/sample.py
- Command line
-
hf download hf://fushinguyenex/IMF/sample.py
-
curl -L -o sample.py https://huggingface.co/fushinguyenex/IMF/resolve/main/sample.py
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() | |