"""Head and/or NAR-LoRA training on REAL audio with the flow loss as teacher. usage: joint.py [rank] Real window: MERT -> head -> straight-through tokens -> (LoRA'd) NAR flow loss on true VAE latents. Head also gets minted soft-CE each step; in LoRA mode 25% of flow windows are minted (true tokens). Eval = held-out real flow loss (fixed windows/t/noise), minted top-1, repeat rate. Ends by rendering the held-out track with the best head+NAR.""" import os, sys, glob, json, math, time, random, hashlib, numpy as np, torch, torch.nn as nn, torch.nn.functional as F, soundfile as sf from torch.utils.checkpoint import checkpoint os.environ.setdefault("HF_HOME","/workspace/hf"); torch.backends.cuda.matmul.allow_tf32=True from yue2.modeling_yue2 import YuE2ForCausalLM from yue2.modeling_vae import YuE2VAE from yue2.protocol import CODEC_OFFSET, MUSIC_END, SongRequest, token_prefixes from yue2.tokenization_yue2 import YuE2TextTokenizer from yue2.nar import attention as nar_attention, synthesize import torchaudio, subprocess, tempfile AUX_W=float(os.environ.get("AUX_W","0")); AUX_TMAX=float(os.environ.get("AUX_TMAX","0.4")); AUX_FR=int(os.environ.get("AUX_FR","150")); AUX_MARGIN=int(os.environ.get("AUX_MARGIN","25")); HOP=1920 NAME=sys.argv[1]; STEPS=int(sys.argv[2]); TRAIN_HEAD=int(sys.argv[3]); TRAIN_LORA=int(sys.argv[4]); INIT_HEAD=sys.argv[5]; INIT_LORA=sys.argv[6]; RANK=int(sys.argv[7]) if len(sys.argv)>7 else 32; HOLD="05_crossing_the_frame" W="/workspace/tok/full"; ROOT="/workspace/yue2-corpus/tracks"; RP=os.environ.get("RP","/workspace/real/prep"); HOLDS=[h for h in os.environ.get("HOLD","").split(",") if h]; OUT=f"{W}/{NAME}"; os.makedirs(OUT,exist_ok=True); dev="cuda" VOCAB=32768; WIN=512; D=512; L=8; H=8; LR_HEAD=1e-4; LR_LORA=5e-5; LR_IO=2e-5; ALPHA=0.25; TAU=0.05; MB=16; MINTED_FLOW_P=0.25 snap=glob.glob("/workspace/hf/hub/models--m-a-p--YuE2-3B/snapshots/*")[0]; vsnap=glob.glob("/workspace/hf/hub/models--m-a-p--YuE2-Vae/snapshots/*")[0] model=YuE2ForCausalLM.from_pretrained(snap, local_files_only=True, torch_dtype=torch.bfloat16, low_cpu_mem_usage=True).eval().to(dev); model.requires_grad_(False); bb=model.model tok=YuE2TextTokenizer(snap+"/qwen.tiktoken"); Ecodec=bb.embed_tokens.weight[CODEC_OFFSET:CODEC_OFFSET+VOCAB] class LoRALinear(nn.Module): def __init__(s, base, r): super().__init__(); s.base=base; s.A=nn.Parameter(torch.randn(r, base.in_features, device=base.weight.device)*(1/math.sqrt(base.in_features))); s.B=nn.Parameter(torch.zeros(base.out_features, r, device=base.weight.device)) def forward(s,x): return s.base(x)+((x.float()@s.A.T)@s.B.T).to(x.dtype) lora_params=[] for layer in bb.layers: for mod,names in ((layer.nar_self_attn,("q_proj","k_proj","v_proj","o_proj")),(layer.nar_mlp,("gate_proj","up_proj","down_proj"))): for n in names: l=LoRALinear(getattr(mod,n),RANK); setattr(mod,n,l); lora_params+=[l.A,l.B] model.vae2llm.float(); model.llm2vae.float(); io_params=list(model.vae2llm.parameters())+list(model.llm2vae.parameters()) def load_lora(path): ck=torch.load(path,map_location=dev) with torch.no_grad(): for p,v in zip(lora_params,ck["lora"]): p.copy_(v.to(dev)) model.vae2llm.load_state_dict({k:v.float() for k,v in ck["io"]["vae2llm"].items()}); model.llm2vae.load_state_dict({k:v.float() for k,v in ck["io"]["llm2vae"].items()}) def save_lora(path): torch.save({"lora":[p.detach().cpu() for p in lora_params],"io":{"vae2llm":model.vae2llm.state_dict(),"llm2vae":model.llm2vae.state_dict()},"rank":RANK}, path) if INIT_LORA!="none": load_lora(INIT_LORA); print("loaded NAR LoRA", INIT_LORA, flush=True) for p in lora_params+io_params: p.requires_grad_(bool(TRAIN_LORA)) class Tok(nn.Module): def __init__(s, din): super().__init__(); s.inp=nn.Linear(din,D); s.pos=nn.Parameter(torch.zeros(1,WIN,D)) layer=nn.TransformerEncoderLayer(D,H,4*D,dropout=0.1,batch_first=True,norm_first=True,activation="gelu"); s.enc=nn.TransformerEncoder(layer,L); s.norm=nn.LayerNorm(D); s.head=nn.Linear(D,VOCAB) def forward(s,x): return s.head(s.norm(s.enc(s.inp(x)+s.pos[:,:x.shape[1]]))) head=Tok(1024).to(dev); head.load_state_dict(torch.load(INIT_HEAD,map_location=dev)["model"]); head.requires_grad_(bool(TRAIN_HEAD)) groups=[] if TRAIN_HEAD: groups.append({"params":list(head.parameters()),"lr":LR_HEAD,"weight_decay":0.05}) if TRAIN_LORA: groups+=[{"params":lora_params,"lr":LR_LORA,"weight_decay":0.0},{"params":io_params,"lr":LR_IO,"weight_decay":0.0}] opt=torch.optim.AdamW(groups,betas=(0.9,0.95)); base_lrs=[g["lr"] for g in opt.param_groups] def instnorm(x): x=x.astype(np.float32); return (x-x.mean(0))/(x.std(0)+1e-5) real=[]; holds=[] for d in sorted(glob.glob(f"{RP}/*")): item=dict(name=os.path.basename(d), mert=instnorm(np.load(f"{d}/mert.npy")), lat=np.load(f"{d}/lat.npy"), prefix=[int(v) for v in np.load(f"{d}/prefix.npy")]); n=min(len(item["mert"]),len(item["lat"])); item["mert"]=item["mert"][:n]; item["lat"]=item["lat"][:n] if item["name"] in (HOLDS or [HOLD]): holds.append(item) elif len(item["lat"])>=WIN: real.append(item) # tracks shorter than one window cannot be sampled held=lambda p: int(hashlib.md5(p.encode()).hexdigest(),16)%20==0; pids=[os.path.basename(f)[:-4] for f in sorted(glob.glob(f"{W}/feats/*.npy"))] def load_m(p): y=np.load(f"{ROOT}/{p}/semantic.npy").astype(np.int64); a=np.load(f"{W}/feats/{p}.npy",mmap_mode="r"); x=np.asarray(a[3] if a.ndim==3 else a); n=min(len(x),len(y)); return instnorm(x[:n]).astype(np.float16), y[:n] # legacy [4,T,1024] or L20-only [T,1024] MCAP=int(os.environ.get("MINTED_CAP","4000")); _r=random.Random(7); mtrain_pids=[p for p in pids if not held(p)]; _r.shuffle(mtrain_pids); mtrain_pids=sorted(mtrain_pids[:MCAP]); _vp=[p for p in pids if held(p)]; _r.shuffle(_vp) mtrain=[load_m(p) for p in mtrain_pids]; mval=[load_m(p) for p in sorted(_vp[:200])] # capped minted anchor (load time); eval on 200 held-out minted tracks print(f"{NAME}: head {TRAIN_HEAD} lora {TRAIN_LORA} | real {len(real)} tracks, held-out {[h['name'] for h in holds]} | minted {len(mtrain)}/{len(mval)}", flush=True); hold=holds[0] NB=torch.tensor(np.load(f"{W}/sem_nbr_idx.npy").astype(np.int64),device=dev); NW=torch.softmax(torch.tensor(np.load(f"{W}/sem_nbr_cos.npy"),device=dev)/TAU,dim=1) def mbatch(data,bs): xs,ys=[],[] for _ in range(bs): x,y=random.choice(data); s=random.randint(0,max(0,len(x)-WIN)); xw=x[s:s+WIN].astype(np.float32); yw=y[s:s+WIN] if len(xw)0 and t<=AUX_TMAX) else None; la=torch.zeros((),device=dev) if la is None else la if TRAIN_HEAD: x,y=mbatch(mtrain,MB) with torch.autocast("cuda",dtype=torch.bfloat16): lgm=head(x) lc=soft_ce(lgm,y) else: lc=torch.zeros((),device=dev) loss=ln+lc+AUX_W*la; opt.zero_grad(set_to_none=True); loss.backward(); torch.nn.utils.clip_grad_norm_([p for g_ in opt.param_groups for p in g_["params"]],1.0); opt.step() if st<=3 or st%25==0: print(f"step {st} nar {ln.item():.4f} ce {lc.item():.3f} aux {float(la):.3f} {time.time()-t0:.0f}s mem {torch.cuda.max_memory_allocated()/2**30:.1f}G", flush=True) if st%100==0 or st==STEPS: e,a,rp=evaluate(); msg=f"EVAL step {st} real_nar {e:.4f} minted_top1 {a:.4f} real_repeat {rp:.3f} real_mel {eval_mel():.2f}dB {time.time()-t0:.0f}s"; print(msg, flush=True); log.write(msg+"\n"); log.flush() if e=T else WIN//4); out[lo:hi]=pred[lo-s0:hi-s0] return out vae=vae_aux for hd in holds: toks=predict(hd["mert"]); print(f"held-out tokens {hd['name']}: unique {len(set(toks.tolist()))/len(toks):.2f} repeat {float((toks[1:]==toks[:-1]).mean()):.4f}", flush=True) with torch.inference_mode(): z=synthesize(model, hd["prefix"], [int(v) for v in toks], 4242, steps=32).float().cpu() with torch.inference_mode(): audio=vae.decode_tiled(z.T[None].contiguous(), core_frames=750, halo_frames=16, output_device="cpu") sf.write(f"{W}/listen_real/real_pred_{NAME}_{hd['name']}.flac", audio[0].float().clamp(-1,1).T.numpy(), 48000, subtype="PCM_24"); print(f"RENDER DONE {hd['name']}", flush=True)