import torch, torch.nn as nn, numpy as np, json, math, os, argparse torch.set_default_device('cuda') LAT=64 class Dec(nn.Module): def __init__(s): super().__init__(); s.net=nn.Sequential(nn.Conv2d(LAT,192,3,1,1),nn.GELU(),nn.ConvTranspose2d(192,96,4,2,1),nn.GELU(),nn.ConvTranspose2d(96,48,4,2,1),nn.GELU(),nn.ConvTranspose2d(48,3,4,2,1)) def forward(s,z): return s.net(z) class TagEnc(nn.Module): def __init__(s, vocab, emb=48, d=256): super().__init__(); s.emb=nn.Embedding(vocab,emb); s.gru=nn.GRU(emb,d,batch_first=True) def forward(s,seq,lens): packed=torch.nn.utils.rnn.pack_padded_sequence(s.emb(seq).to('cuda'),lens.cpu(),batch_first=True,enforce_sorted=False) return s.gru(packed)[1].squeeze(0) def t_emb(t,dim): half=dim//2; f=torch.exp(-math.log(10000)*torch.arange(half,device='cuda')/half); a=t[:,None].float()*f[None,:] return torch.cat([torch.cos(a),torch.sin(a)],1) class CDiff(nn.Module): def __init__(s,cin=LAT,tagdim=256,base=256): super().__init__(); s.temb=nn.Sequential(nn.Linear(64,base),nn.GELU(),nn.Linear(base,cin)); s.tagproj=nn.Linear(tagdim,cin) s.d1=nn.Sequential(nn.Conv2d(cin,base,3,1,1),nn.GroupNorm(8,base),nn.GELU(),nn.Conv2d(base,base,3,1,1),nn.GroupNorm(8,base),nn.GELU()) s.d2=nn.Sequential(nn.Conv2d(base,base,3,2,1),nn.GroupNorm(8,base),nn.GELU()); s.mid=nn.Sequential(nn.Conv2d(base,base,3,1,1),nn.GroupNorm(8,base),nn.GELU()) s.u2=nn.Sequential(nn.ConvTranspose2d(base,base,3,2,1,1),nn.GroupNorm(8,base),nn.GELU()); s.outc=nn.Conv2d(base,cin,3,1,1) def forward(s,x,t,tag): te=s.temb(t_emb(t,64)).unsqueeze(-1).unsqueeze(-1); tg=s.tagproj(tag).unsqueeze(-1).unsqueeze(-1) x=x+te+tg; h1=s.d1(x); h2=s.mid(s.d2(h1)); h=s.u2(h2)+h1; return s.outc(h) def load_model(ckpt): ck=torch.load(ckpt,weights_only=False,map_location='cuda'); tok2id=ck['tok2id'] dec=Dec(); te=TagEnc(vocab=len(tok2id)+1); cd=CDiff(LAT) dec.load_state_dict(ck['dec']); te.load_state_dict(ck['te']); cd.load_state_dict(ck['cd']) dec.eval(); te.eval(); cd.eval() pc_mean=ck.get('pc_mean',torch.zeros(1,64,1,1)).to('cuda'); pc_std=ck.get('pc_std',torch.ones(1,64,1,1)).to('cuda') return dec,te,cd,tok2id,pc_mean,pc_std def generate(prompt, dec,te,cd,tok2id,pc_mean,pc_std, n=3, steps=30, cfg=4.0, seed=0): MAXL=40 def encp(t): toks=[tok2id[x.strip()] for x in t.split(',') if x.strip() in tok2id][:MAXL]; return toks+[0]*(MAXL-len(toks)) torch.manual_seed(seed); dev='cuda' TMAX=200; beta=torch.linspace(1e-4,0.02,TMAX,device=dev); alpha=1-beta; abar=torch.cumprod(alpha,0) toks=torch.tensor([encp(prompt)]*n,device=dev); lens=torch.tensor([max(1,min(len([x for x in encp(prompt) if x!=0]),MAXL))]*n) utok=torch.zeros_like(toks); ulen=torch.ones_like(lens) st=torch.linspace(TMAX-1,0,steps).long(); zn=torch.randn(n,LAT,12,12,device=dev) with torch.no_grad(): c=te(toks,lens); ucond=te(utok,ulen) for idx in range(len(st)): t=st[idx]; tn=st[idx+1] if idx+1=0: ab_tn=abar[tn].view(-1,1,1,1); z0=(zn-torch.sqrt(1-ab_t)*ep)/torch.sqrt(ab_t) zn=torch.sqrt(ab_tn)*z0+torch.sqrt(torch.clamp(1-ab_tn,min=0))*ep else: zn=(zn-torch.sqrt(1-ab_t)*ep)/torch.sqrt(ab_t) z=zn*pc_std+pc_mean; imgs=dec(z) return imgs.clamp(0,1) if __name__=='__main__': ap=argparse.ArgumentParser(); ap.add_argument('--ckpt',required=True); ap.add_argument('--prompt',required=True); ap.add_argument('--out',default='out'); ap.add_argument('-n',type=int,default=3); ap.add_argument('--steps',type=int,default=30); ap.add_argument('--cfg',type=float,default=4.0) a=ap.parse_args() dec,te,cd,tok,pm,ps=load_model(a.ckpt) imgs=generate(a.prompt,dec,te,cd,tok,pm,ps,n=a.n,steps=a.steps,cfg=a.cfg) from PIL import Image os.makedirs(a.out,exist_ok=True) for i in range(a.n): im=imgs[i].permute(1,2,0).cpu().numpy() Image.fromarray((np.clip(im,0,1)*255).astype(np.uint8)).save(f"{a.out}/gen_{i}.png") print(f"saved {a.n} images to {a.out}")