""" 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()