Download sample2.py from M1n1A1/MiniAI-TITS2: direct link, hf CLI and curl.
- Browser
- Download file 2.43 kB
-
https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/sample2.py
- Command line
-
hf download hf://M1n1A1/MiniAI-TITS2/sample2.py
-
curl -L -o sample2.py https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/sample2.py
2.43 kB
| """ | |
| T.I.T.S.2 inference — rectified flow, ~30 Euler steps, 256x256. | |
| Usage: | |
| python sample2.py --checkpoint checkpoints2/latest.pt --prompt "a red bus on a city street" \ | |
| --num_images 4 --guidance 4.0 --out out.png | |
| """ | |
| import argparse | |
| import torch | |
| from safetensors import safe_open | |
| from safetensors.torch import load_file | |
| from torchvision.utils import save_image | |
| from model2 import TITS2 | |
| from preview import decode, sample_latents | |
| from text_encoder2 import FrozenCLIPTextEncoder | |
| def parse_args(): | |
| p = argparse.ArgumentParser() | |
| p.add_argument("--checkpoint", type=str, required=True) | |
| p.add_argument("--prompt", type=str, required=True) | |
| p.add_argument("--num_images", type=int, default=1) | |
| p.add_argument("--steps", type=int, default=30) | |
| p.add_argument("--guidance", type=float, default=4.0) | |
| p.add_argument("--seed", type=int, default=None) | |
| p.add_argument("--raw", action="store_true", help="use raw weights instead of EMA") | |
| p.add_argument("--out", type=str, default="sample2.png") | |
| p.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu") | |
| return p.parse_args() | |
| def main(): | |
| args = parse_args() | |
| text_encoder = FrozenCLIPTextEncoder(device=args.device) | |
| if args.checkpoint.endswith(".safetensors"): | |
| with safe_open(args.checkpoint, framework="pt") as f: | |
| meta = f.metadata() or {} | |
| cfg = {k: int(meta.get(k, d)) for k, d in (("dim", 576), ("depth", 12), ("heads", 9))} | |
| step = meta.get("step", "?") | |
| state = {k: v.float() for k, v in load_file(args.checkpoint).items()} | |
| else: | |
| ck = torch.load(args.checkpoint, map_location=args.device, weights_only=False) | |
| cfg = {k: ck[k] for k in ("dim", "depth", "heads")} | |
| step = ck["step"] | |
| state = ck["model_state_dict"] if args.raw else ck["ema_state_dict"] | |
| model = TITS2(text_dim=text_encoder.embed_dim, **cfg).to(args.device) | |
| model.load_state_dict(state) | |
| model.eval() | |
| lat = sample_latents(model, text_encoder, [args.prompt] * args.num_images, | |
| steps=args.steps, guidance=args.guidance, | |
| device=args.device, seed=args.seed) | |
| imgs = decode(lat, args.device) | |
| save_image(imgs, args.out, nrow=int(args.num_images**0.5) or 1) | |
| print(f"saved {args.num_images} image(s) to {args.out} (step {step})") | |
| if __name__ == "__main__": | |
| main() | |