File size: 1,879 Bytes
2ca14fe | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 | """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()
|