MiniAI-TITS2 / sample2.py
Codeminute's picture
sample2.py: load .safetensors directly (as the card instructs)
a30f2d4 verified
Raw History Blame Contribute Delete
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()