Download sample_quant.py from M1n1A1/MiniAI-TITS2: direct link, hf CLI and curl.
- Browser
- Download file 1.88 kB
-
https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/sample_quant.py
- Command line
-
hf download hf://M1n1A1/MiniAI-TITS2/sample_quant.py
-
curl -L -o sample_quant.py https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/sample_quant.py
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() | |