Asilarkness's picture
Upload infer.py with huggingface_hub
72d9b5e verified
Raw
History Blame Contribute Delete
4.31 kB
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}")