MiniAI-TITS2 / sample_quant.py
Codeminute's picture
T.I.T.S.2 — 93M DiT, rectified flow, 256px, 24 epochs on cleaned CC12M
2ca14fe verified
Raw History Blame Contribute Delete
1.88 kB
"""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()