"""Load a 1-bit T.I.T.S.2 export and generate from it. See quantize.py for the format.""" import argparse import torch from safetensors.torch import load_file from torchvision.utils import save_image from model2 import TITS2 from preview import decode, sample_latents from quantize import unpack from text_encoder2 import FrozenCLIPTextEncoder def load_1bit(path): raw = load_file(path) sd = {} for k, v in raw.items(): if k.endswith(".packed"): base = k[: -len(".packed")] shape = raw[base + ".shape"].tolist() sd[base] = unpack(v, raw[base + ".scale"], shape[1]) elif not k.endswith((".scale", ".shape")): sd[k] = v.float() return sd def main(): p = argparse.ArgumentParser() p.add_argument("--weights", default="exports/tits2_1bit.safetensors") p.add_argument("--prompts", nargs="+", default=["a red double decker bus on a city street"]) p.add_argument("--steps", type=int, default=30) p.add_argument("--guidance", type=float, default=4.5) p.add_argument("--seed", type=int, default=7) p.add_argument("--out", default="out_1bit.png") p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") args = p.parse_args() te = FrozenCLIPTextEncoder(device=args.device) sd = load_1bit(args.weights) if "1bit" in args.weights else {k: v.float() for k, v in load_file(args.weights).items()} model = TITS2(dim=576, depth=12, heads=9, text_dim=te.embed_dim).to(args.device) model.load_state_dict(sd) model.eval() lat = sample_latents(model, te, args.prompts, steps=args.steps, guidance=args.guidance, device=args.device, seed=args.seed) save_image(decode(lat, args.device), args.out, nrow=2, padding=4, pad_value=1) print(f"saved {args.out}") if __name__ == "__main__": main()