|
|
| 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<len(st) else -1; tt=torch.tensor([t.item()]*n,device=dev) |
| ep_c=cd(zn,tt,c); ep_u=cd(zn,tt,ucond); ep=ep_u+cfg*(ep_c-ep_u) |
| ab_t=abar[t].view(-1,1,1,1) |
| if tn>=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}") |
|
|