Publish recovered joint_v6 audio-loss trainer and required neighbor tables
Browse filesPreserve the original training script byte-for-byte, document historical dependencies and limitations, and fix the missing README link. No model weights changed.
- README.md +2 -2
- assets/sem_nbr_cos.npy +3 -0
- assets/sem_nbr_idx.npy +3 -0
- scripts/joint_v6.py +210 -0
- scripts/joint_v6_README.md +97 -0
- scripts/joint_v6_release.json +39 -0
README.md
CHANGED
|
@@ -24,7 +24,7 @@ an artist, and generate new songs or covers.
|
|
| 24 |
| `scripts/` | the training loop and inference scripts (below). `scripts/ckpt_io.py` loads either format. |
|
| 25 |
|
| 26 |
### Safetensors layout
|
| 27 |
-
Same weights as the `.pt` files, bit-exact in fp32 (the `.bf16` variants are half the size; head top-1 agreement with fp32 is 98.6%).
|
| 28 |
extension via `scripts/ckpt_io.load_ckpt(path)`, which returns the same dict the `.pt` files hold.
|
| 29 |
|
| 30 |
- **Head**: the plain `state_dict` of the 8-layer encoder (`inp.*`, `pos`, `enc.layers.{0..7}.*`, `norm.*`, `head.*`), 103 tensors.
|
|
@@ -118,4 +118,4 @@ The v5 head co-trained for 3,000 steps with a rank-32 decoder LoRA on **128 real
|
|
| 118 |
|
| 119 |
Listening (Kytra): v7 clearly better than v6, v8 clearly better than v7, v9 preferred overall. Minted top-1 stays flat across the sweep, so the head remains universal; the latent loss moves by 0.003, so the audio term is not fighting the latent objective. Mel L1 did **not** follow the listening results, LTAS did. Token-choice errors (occasional out-of-tune notes) are unchanged by this loss; that is a head-accuracy problem.
|
| 120 |
|
| 121 |
-
**Use:** tokenize real audio with the v9 (or v8) head and load the matching `nar_lora_joint_v9_comfyui` / `_v8_comfyui` on the decoder at model strength 1.0. Trainer: `scripts/
|
|
|
|
| 24 |
| `scripts/` | the training loop and inference scripts (below). `scripts/ckpt_io.py` loads either format. |
|
| 25 |
|
| 26 |
### Safetensors layout
|
| 27 |
+
Same weights as the `.pt` files, bit-exact in fp32 (the `.bf16` variants are half the size; head top-1 agreement with fp32 is 98.6%). The previously published scripts accept either
|
| 28 |
extension via `scripts/ckpt_io.load_ckpt(path)`, which returns the same dict the `.pt` files hold.
|
| 29 |
|
| 30 |
- **Head**: the plain `state_dict` of the 8-layer encoder (`inp.*`, `pos`, `enc.layers.{0..7}.*`, `norm.*`, `head.*`), 103 tensors.
|
|
|
|
| 118 |
|
| 119 |
Listening (Kytra): v7 clearly better than v6, v8 clearly better than v7, v9 preferred overall. Minted top-1 stays flat across the sweep, so the head remains universal; the latent loss moves by 0.003, so the audio term is not fighting the latent objective. Mel L1 did **not** follow the listening results, LTAS did. Token-choice errors (occasional out-of-tune notes) are unchanged by this loss; that is a head-accuracy problem.
|
| 120 |
|
| 121 |
+
**Use:** tokenize real audio with the v9 (or v8) head and load the matching `nar_lora_joint_v9_comfyui` / `_v8_comfyui` on the decoder at model strength 1.0. Trainer: [`scripts/joint_v6.py`](scripts/joint_v6.py), the recovered original audio-loss variant of `scripts/joint.py` (env `AUX_W`, `AUX_TMAX`, `AUX_FR`). Read the [setup and historical-reproduction notes](scripts/joint_v6_README.md) first: this original script uses `.pt` checkpoints and the old pod paths; unlike the other published scripts, it has not been adapted to `ckpt_io`. The required neighbor tables are included under `assets/`.
|
assets/sem_nbr_cos.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:cf544760a9ef2f66b0bba4fa6ac3d055c628fc109cea47ba9bc6de59c767ad12
|
| 3 |
+
size 2097280
|
assets/sem_nbr_idx.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:081d5f3d41ac2b604ee95db53f5c2069b33d871c317b611d7df8e02b42eda141
|
| 3 |
+
size 2097280
|
scripts/joint_v6.py
ADDED
|
@@ -0,0 +1,210 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Head and/or NAR-LoRA training on REAL audio with the flow loss as teacher.
|
| 2 |
+
usage: joint.py <name> <steps> <train_head 0|1> <train_lora 0|1> <init_head.pt> <init_lora.pt|none> [rank]
|
| 3 |
+
Real window: MERT -> head -> straight-through tokens -> (LoRA'd) NAR flow loss on true VAE latents. Head also gets minted soft-CE each step;
|
| 4 |
+
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.
|
| 5 |
+
Ends by rendering the held-out track with the best head+NAR."""
|
| 6 |
+
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
|
| 7 |
+
from torch.utils.checkpoint import checkpoint
|
| 8 |
+
os.environ.setdefault("HF_HOME","/workspace/hf"); torch.backends.cuda.matmul.allow_tf32=True
|
| 9 |
+
from yue2.modeling_yue2 import YuE2ForCausalLM
|
| 10 |
+
from yue2.modeling_vae import YuE2VAE
|
| 11 |
+
from yue2.protocol import CODEC_OFFSET, MUSIC_END, SongRequest, token_prefixes
|
| 12 |
+
from yue2.tokenization_yue2 import YuE2TextTokenizer
|
| 13 |
+
from yue2.nar import attention as nar_attention, synthesize
|
| 14 |
+
import torchaudio, subprocess, tempfile
|
| 15 |
+
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
|
| 16 |
+
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"
|
| 17 |
+
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"
|
| 18 |
+
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
|
| 19 |
+
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]
|
| 20 |
+
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
|
| 21 |
+
tok=YuE2TextTokenizer(snap+"/qwen.tiktoken"); Ecodec=bb.embed_tokens.weight[CODEC_OFFSET:CODEC_OFFSET+VOCAB]
|
| 22 |
+
class LoRALinear(nn.Module):
|
| 23 |
+
def __init__(s, base, r):
|
| 24 |
+
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))
|
| 25 |
+
def forward(s,x): return s.base(x)+((x.float()@s.A.T)@s.B.T).to(x.dtype)
|
| 26 |
+
lora_params=[]
|
| 27 |
+
for layer in bb.layers:
|
| 28 |
+
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"))):
|
| 29 |
+
for n in names: l=LoRALinear(getattr(mod,n),RANK); setattr(mod,n,l); lora_params+=[l.A,l.B]
|
| 30 |
+
model.vae2llm.float(); model.llm2vae.float(); io_params=list(model.vae2llm.parameters())+list(model.llm2vae.parameters())
|
| 31 |
+
def load_lora(path):
|
| 32 |
+
ck=torch.load(path,map_location=dev)
|
| 33 |
+
with torch.no_grad():
|
| 34 |
+
for p,v in zip(lora_params,ck["lora"]): p.copy_(v.to(dev))
|
| 35 |
+
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()})
|
| 36 |
+
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)
|
| 37 |
+
if INIT_LORA!="none": load_lora(INIT_LORA); print("loaded NAR LoRA", INIT_LORA, flush=True)
|
| 38 |
+
for p in lora_params+io_params: p.requires_grad_(bool(TRAIN_LORA))
|
| 39 |
+
class Tok(nn.Module):
|
| 40 |
+
def __init__(s, din):
|
| 41 |
+
super().__init__(); s.inp=nn.Linear(din,D); s.pos=nn.Parameter(torch.zeros(1,WIN,D))
|
| 42 |
+
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)
|
| 43 |
+
def forward(s,x): return s.head(s.norm(s.enc(s.inp(x)+s.pos[:,:x.shape[1]])))
|
| 44 |
+
head=Tok(1024).to(dev); head.load_state_dict(torch.load(INIT_HEAD,map_location=dev)["model"]); head.requires_grad_(bool(TRAIN_HEAD))
|
| 45 |
+
groups=[]
|
| 46 |
+
if TRAIN_HEAD: groups.append({"params":list(head.parameters()),"lr":LR_HEAD,"weight_decay":0.05})
|
| 47 |
+
if TRAIN_LORA: groups+=[{"params":lora_params,"lr":LR_LORA,"weight_decay":0.0},{"params":io_params,"lr":LR_IO,"weight_decay":0.0}]
|
| 48 |
+
opt=torch.optim.AdamW(groups,betas=(0.9,0.95)); base_lrs=[g["lr"] for g in opt.param_groups]
|
| 49 |
+
def instnorm(x): x=x.astype(np.float32); return (x-x.mean(0))/(x.std(0)+1e-5)
|
| 50 |
+
real=[]; holds=[]
|
| 51 |
+
for d in sorted(glob.glob(f"{RP}/*")):
|
| 52 |
+
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]
|
| 53 |
+
if item["name"] in (HOLDS or [HOLD]): holds.append(item)
|
| 54 |
+
elif len(item["lat"])>=WIN: real.append(item) # tracks shorter than one window cannot be sampled
|
| 55 |
+
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"))]
|
| 56 |
+
def load_m(p):
|
| 57 |
+
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]
|
| 58 |
+
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)
|
| 59 |
+
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
|
| 60 |
+
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]
|
| 61 |
+
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)
|
| 62 |
+
def mbatch(data,bs):
|
| 63 |
+
xs,ys=[],[]
|
| 64 |
+
for _ in range(bs):
|
| 65 |
+
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]
|
| 66 |
+
if len(xw)<WIN: pad=WIN-len(xw); xw=np.pad(xw,((0,pad),(0,0))); yw=np.pad(yw,(0,pad),constant_values=-100)
|
| 67 |
+
xs.append(xw); ys.append(yw)
|
| 68 |
+
return torch.tensor(np.stack(xs),device=dev), torch.tensor(np.stack(ys),device=dev)
|
| 69 |
+
def soft_ce(lg,y):
|
| 70 |
+
m=y!=-100; lg=lg[m].float(); y=y[m]; logp=F.log_softmax(lg,-1); return ((1-ALPHA)*(-logp.gather(1,y[:,None])[:,0])+ALPHA*(-(logp.gather(1,NB[y])*NW[y]).sum(1))).mean()
|
| 71 |
+
def ar_layer(layer,x,cos_,sin_):
|
| 72 |
+
q,k,v=layer.self_attn.project_qkv(layer.input_layernorm(x),cos_,sin_); h=nar_attention(q[0],k[0],v[0],causal=True)
|
| 73 |
+
x=x+layer.self_attn.o_proj(h.flatten(1)[None]); return x+layer.mlp(layer.post_attention_layernorm(x)), k[0], v[0]
|
| 74 |
+
def nar_layer(layer,h,ak,av,ncos,nsin):
|
| 75 |
+
q,k,v=layer.nar_self_attn.project_qkv(layer.nar_input_layernorm(h),ncos,nsin); a=nar_attention(q[0],torch.cat((ak,k[0])),torch.cat((av,v[0])))
|
| 76 |
+
h=h+layer.nar_self_attn.o_proj(a.flatten(1)[None]); return h+layer.nar_mlp(layer.nar_pre_mlp_layernorm(h))
|
| 77 |
+
def flow_loss(prefix, codec_emb, x1, t, noise, grad_ar, grad_nar):
|
| 78 |
+
pre=bb.embed_tokens(torch.tensor([prefix],device=dev))[0]; end=bb.embed_tokens(torch.tensor([MUSIC_END],device=dev)); x=torch.cat((pre,codec_emb.to(pre.dtype),end),0)[None]
|
| 79 |
+
Lq=x.shape[1]; cos_,sin_=bb.rotary_emb(torch.arange(Lq,device=dev)[None]); cache=[]
|
| 80 |
+
if grad_ar:
|
| 81 |
+
for layer in bb.layers: x,k,v=checkpoint(ar_layer,layer,x,cos_,sin_,use_reentrant=False); cache.append((k,v))
|
| 82 |
+
else:
|
| 83 |
+
with torch.no_grad():
|
| 84 |
+
for layer in bb.layers: x,k,v=ar_layer(layer,x,cos_,sin_); cache.append((k,v))
|
| 85 |
+
T=codec_emb.shape[0]; xt=t*noise+(1-t)*x1; target=noise-x1; N=T+2; ncos,nsin=bb.rotary_emb(torch.arange(Lq,Lq+N,device=dev)[None])
|
| 86 |
+
pe=model.latent_pos_embed(torch.arange(N,device=dev).clamp(max=model.config.max_latent_frames-1))[None]; sh=model._shift_t_value(float(np.clip(np.log(t/(1-t)),-20,20)),dev,torch.bfloat16)
|
| 87 |
+
h=model.vae2llm(F.pad(xt,(0,0,1,1))[None].float()).to(torch.bfloat16)+model.time_embedder(sh.expand(N))[None]+pe
|
| 88 |
+
for layer,(ak,av) in zip(bb.layers,cache): h=checkpoint(nar_layer,layer,h,ak,av,ncos,nsin,use_reentrant=False) if (grad_ar or grad_nar) else nar_layer(layer,h,ak,av,ncos,nsin)
|
| 89 |
+
v_hat=model.llm2vae(bb.norm(h)[0,1:-1].float()); return F.mse_loss(v_hat,target), v_hat, xt
|
| 90 |
+
def st_embed(logits):
|
| 91 |
+
p=torch.softmax(logits.float(),-1); idx=p.argmax(-1); hard=F.one_hot(idx,VOCAB).float(); return ((hard+(p-p.detach())).to(Ecodec.dtype))@Ecodec, idx
|
| 92 |
+
def real_window(item,s=None):
|
| 93 |
+
n=len(item["lat"]); s=random.randint(0,n-WIN) if s is None else s; item["_s"]=s; return item["mert"][s:s+WIN], torch.tensor(item["lat"][s:s+WIN],device=dev)
|
| 94 |
+
# ---- audio-domain auxiliary loss (frozen VAE decoder, differentiable w.r.t. its input)
|
| 95 |
+
vae_aux=YuE2VAE.from_pretrained(vsnap, decoder_only=True, device=dev, local_files_only=True); vae_aux.decoder.requires_grad_(False)
|
| 96 |
+
_mel=torchaudio.transforms.MelSpectrogram(48000,n_fft=2048,hop_length=480,n_mels=128,power=1.0).to(dev); _wins={n:torch.hann_window(n,device=dev) for n in (512,1024,2048)}
|
| 97 |
+
def audio_path(name):
|
| 98 |
+
tag,_,base=name.partition("__")
|
| 99 |
+
if tag=="coheed": c=[f"/workspace/real/coheed/{base}.flac"]
|
| 100 |
+
else: c=glob.glob(f"/workspace/ComfyUI/input/train/{tag}/{base}.*")
|
| 101 |
+
c=[p for p in c if p.lower().endswith((".flac",".mp3",".wav"))]; return c[0] if c else None
|
| 102 |
+
_acache={}
|
| 103 |
+
def audio48(item):
|
| 104 |
+
if item["name"] in _acache: return _acache[item["name"]]
|
| 105 |
+
p=audio_path(item["name"]); a=None
|
| 106 |
+
if p:
|
| 107 |
+
try: x,sr=sf.read(p,dtype="float32")
|
| 108 |
+
except Exception:
|
| 109 |
+
with tempfile.NamedTemporaryFile(suffix=".wav") as tf: subprocess.run(["ffmpeg","-v","error","-y","-i",p,"-ac","2","-ar","48000",tf.name],check=True); x,sr=sf.read(tf.name,dtype="float32")
|
| 110 |
+
x=np.stack([x,x],1) if x.ndim==1 else x; a=torch.from_numpy(x.T.copy())
|
| 111 |
+
if sr!=48000: a=torchaudio.functional.resample(a,sr,48000)
|
| 112 |
+
_acache[item["name"]]=a; return a
|
| 113 |
+
def spec_loss(pred, ref):
|
| 114 |
+
"""pred, ref: [2,S] float32 @48k. log-mel L1 + multi-res STFT (spectral convergence + log-mag L1) on mono, + stereo width term."""
|
| 115 |
+
pm,rm=pred.mean(0),ref.mean(0); mel=(torch.log(_mel(pm)+1e-5)-torch.log(_mel(rm)+1e-5)).abs().mean(); mr=0
|
| 116 |
+
for n,w in _wins.items():
|
| 117 |
+
A=torch.stft(pm,n,hop_length=n//4,window=w,return_complex=True).abs(); B=torch.stft(rm,n,hop_length=n//4,window=w,return_complex=True).abs()
|
| 118 |
+
mr=mr+((A-B).norm()/(B.norm()+1e-6)+(torch.log(A+1e-5)-torch.log(B+1e-5)).abs().mean())/3
|
| 119 |
+
side=lambda x: torch.log(((x[0]-x[1])/2).pow(2).mean()+1e-7)-torch.log(((x[0]+x[1])/2).pow(2).mean()+1e-7)
|
| 120 |
+
return mel+0.5*mr+0.5*(side(pred)-side(ref)).abs()
|
| 121 |
+
def aux_loss(item, xt, v_hat, t):
|
| 122 |
+
"""clean-latent estimate x1 = xt - t*v, centre AUX_FR frames (+margin) decoded through the VAE vs the original recording's samples."""
|
| 123 |
+
a=audio48(item)
|
| 124 |
+
if a is None: return None
|
| 125 |
+
s0=item["_s"]+(WIN-AUX_FR)//2-AUX_MARGIN; f0=s0+AUX_MARGIN; ref=a[:, f0*HOP:(f0+AUX_FR)*HOP]
|
| 126 |
+
if ref.shape[1]<AUX_FR*HOP: return None
|
| 127 |
+
x1=(xt-t*v_hat)[s0-item["_s"]:s0-item["_s"]+AUX_FR+2*AUX_MARGIN]
|
| 128 |
+
wav=vae_aux.decoder(x1.T[None].float())[0]; wav=wav[:, AUX_MARGIN*HOP:AUX_MARGIN*HOP+AUX_FR*HOP]
|
| 129 |
+
return spec_loss(wav, ref.to(dev)) if wav.shape[1]==ref.shape[1] else None
|
| 130 |
+
@torch.no_grad()
|
| 131 |
+
def eval_mel():
|
| 132 |
+
"""held-out audio-domain check: log-mel L1 (dB) of the t=0.2 clean estimate vs the original, centre window of each held-out track."""
|
| 133 |
+
tot=0; cnt=0
|
| 134 |
+
for hd in holds:
|
| 135 |
+
a=audio48(hd)
|
| 136 |
+
if a is None: continue
|
| 137 |
+
n=len(hd["lat"]); s=min(n//2, n-WIN); m,z=real_window(hd,s)
|
| 138 |
+
with torch.autocast("cuda",dtype=torch.bfloat16): idx=head(torch.tensor(m[None],device=dev))[0].float().argmax(-1)
|
| 139 |
+
g=torch.Generator(device="cpu").manual_seed(7); noise=torch.randn(WIN,64,generator=g).to(dev); _,vh,xt=flow_loss(hd["prefix"],Ecodec[idx],z,0.2,noise,False,False)
|
| 140 |
+
s0=(WIN-AUX_FR)//2-AUX_MARGIN; x1=(xt-0.2*vh)[s0:s0+AUX_FR+2*AUX_MARGIN]; wav=vae_aux.decoder(x1.T[None].float())[0][:, AUX_MARGIN*HOP:AUX_MARGIN*HOP+AUX_FR*HOP]
|
| 141 |
+
f0=s+s0+AUX_MARGIN; ref=a[:, f0*HOP:(f0+AUX_FR)*HOP].to(dev)
|
| 142 |
+
if ref.shape[1]!=wav.shape[1]: continue
|
| 143 |
+
tot+=(20/math.log(10))*(torch.log(_mel(wav.mean(0))+1e-5)-torch.log(_mel(ref.mean(0))+1e-5)).abs().mean().item(); cnt+=1
|
| 144 |
+
return tot/max(1,cnt)
|
| 145 |
+
@torch.no_grad()
|
| 146 |
+
def evaluate():
|
| 147 |
+
head.eval(); tot=0; g=torch.Generator(device="cpu").manual_seed(123); reps=[]
|
| 148 |
+
for hd in holds: # every held-out artist track, 3 fixed windows x 3 noise levels
|
| 149 |
+
n=len(hd["lat"])
|
| 150 |
+
for s in (n//4,n//2,3*n//4):
|
| 151 |
+
s=min(s,n-WIN) # short tracks: keep the eval window inside the track
|
| 152 |
+
m,z=real_window(hd,s)
|
| 153 |
+
with torch.autocast("cuda",dtype=torch.bfloat16): idx=head(torch.tensor(m[None],device=dev))[0].float().argmax(-1)
|
| 154 |
+
reps.append(float((idx[1:]==idx[:-1]).float().mean())); noise=torch.randn(WIN,64,generator=g).to(dev)
|
| 155 |
+
for t in (0.2,0.5,0.8): tot+=flow_loss(hd["prefix"],Ecodec[idx],z,t,noise,False,False)[0].item()/len(holds)
|
| 156 |
+
t1=tot_=0
|
| 157 |
+
for _ in range(24):
|
| 158 |
+
x,y=mbatch(mval,16)
|
| 159 |
+
with torch.autocast("cuda",dtype=torch.bfloat16): lg=head(x)
|
| 160 |
+
mk=y!=-100; t1+=(lg.float().argmax(-1)==y)[mk].sum().item(); tot_+=mk.sum().item()
|
| 161 |
+
head.train(); return tot/9, t1/tot_, float(np.mean(reps))
|
| 162 |
+
def save_all(tag):
|
| 163 |
+
torch.save({"model":head.state_dict(),"cfg":dict(NAME=NAME,instnorm=True)}, f"{OUT}/head_{tag}.pt")
|
| 164 |
+
if TRAIN_LORA: save_lora(f"{OUT}/lora_{tag}.pt")
|
| 165 |
+
e0,a0,r0=evaluate(); msg=f"EVAL step 0 real_nar {e0:.4f} minted_top1 {a0:.4f} real_repeat {r0:.3f} real_mel {eval_mel():.2f}dB"; print(msg, flush=True); log=open(f"{OUT}/train.log","a"); log.write(msg+"\n"); best=e0; save_all("best"); t0=time.time()
|
| 166 |
+
for st in range(1,STEPS+1):
|
| 167 |
+
mult=min(1,st/50)*(0.2+0.8*0.5*(1+math.cos(math.pi*st/STEPS)))
|
| 168 |
+
for g_,b in zip(opt.param_groups,base_lrs): g_["lr"]=b*mult
|
| 169 |
+
if TRAIN_LORA and random.random()<MINTED_FLOW_P: # minted regularizer for the NAR: true tokens, minted latents
|
| 170 |
+
p=random.choice(mtrain_pids); d=f"{ROOT}/{p}"; r=json.load(open(f"{d}/request.json")); y=np.load(f"{d}/semantic.npy").astype(np.int64); z=np.load(f"{d}/latent.npy"); n=min(len(y),len(z)); s=random.randint(0,max(0,n-WIN))
|
| 171 |
+
pre=token_prefixes(SongRequest(style=r["style"],lyrics=r["lyrics"],cot="off",seed=r["seed"],id=p),tok); t=float(np.clip(np.random.beta(2,2),0.02,0.98))
|
| 172 |
+
ln,_,_=flow_loss(pre,Ecodec[torch.tensor(y[s:s+WIN],device=dev)],torch.tensor(z[s:s+WIN],device=dev),t,torch.randn(WIN,64,device=dev),False,True); lc=torch.zeros((),device=dev); la=torch.zeros((),device=dev)
|
| 173 |
+
else:
|
| 174 |
+
item=random.choice(real); m,z=real_window(item); t=float(np.clip(np.random.beta(2,2),0.05,0.95)); noise=torch.randn(WIN,64,device=dev)
|
| 175 |
+
with torch.autocast("cuda",dtype=torch.bfloat16): lg=head(torch.tensor(m[None],device=dev))[0]
|
| 176 |
+
emb,idx=st_embed(lg) if TRAIN_HEAD else (Ecodec[lg.float().argmax(-1)],None)
|
| 177 |
+
ln,vh,xt=flow_loss(item["prefix"],emb,z,t,noise,bool(TRAIN_HEAD),bool(TRAIN_LORA))
|
| 178 |
+
la=aux_loss(item,xt,vh,t) if (AUX_W>0 and t<=AUX_TMAX) else None; la=torch.zeros((),device=dev) if la is None else la
|
| 179 |
+
if TRAIN_HEAD:
|
| 180 |
+
x,y=mbatch(mtrain,MB)
|
| 181 |
+
with torch.autocast("cuda",dtype=torch.bfloat16): lgm=head(x)
|
| 182 |
+
lc=soft_ce(lgm,y)
|
| 183 |
+
else: lc=torch.zeros((),device=dev)
|
| 184 |
+
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()
|
| 185 |
+
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)
|
| 186 |
+
if st%100==0 or st==STEPS:
|
| 187 |
+
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()
|
| 188 |
+
if e<best: best=e; save_all("best")
|
| 189 |
+
save_all("last")
|
| 190 |
+
print(f"RESULT {NAME}: best real_nar {best:.4f} (start {e0:.4f})", flush=True)
|
| 191 |
+
# ---- render held-out with best
|
| 192 |
+
head.load_state_dict(torch.load(f"{OUT}/head_best.pt",map_location=dev)["model"]); head.eval()
|
| 193 |
+
if TRAIN_LORA: load_lora(f"{OUT}/lora_best.pt")
|
| 194 |
+
model.vae2llm.to(torch.bfloat16); model.llm2vae.to(torch.bfloat16)
|
| 195 |
+
@torch.no_grad()
|
| 196 |
+
def predict(x):
|
| 197 |
+
T=len(x); out=np.zeros(T,dtype=np.int64); starts=list(range(0,max(1,T-WIN+1),WIN//2))
|
| 198 |
+
if starts[-1]+WIN<T: starts.append(max(0,T-WIN))
|
| 199 |
+
for s0 in starts:
|
| 200 |
+
xw=x[s0:s0+WIN]; n=len(xw)
|
| 201 |
+
if n<WIN: xw=np.pad(xw,((0,WIN-n),(0,0)))
|
| 202 |
+
with torch.autocast("cuda",dtype=torch.bfloat16): pred=head(torch.tensor(xw[None],device=dev))[0,:n].float().argmax(-1).cpu().numpy()
|
| 203 |
+
lo=s0+(0 if s0==0 else WIN//4); hi=s0+n-(0 if s0+n>=T else WIN//4); out[lo:hi]=pred[lo-s0:hi-s0]
|
| 204 |
+
return out
|
| 205 |
+
vae=vae_aux
|
| 206 |
+
for hd in holds:
|
| 207 |
+
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)
|
| 208 |
+
with torch.inference_mode(): z=synthesize(model, hd["prefix"], [int(v) for v in toks], 4242, steps=32).float().cpu()
|
| 209 |
+
with torch.inference_mode(): audio=vae.decode_tiled(z.T[None].contiguous(), core_frames=750, halo_frames=16, output_device="cpu")
|
| 210 |
+
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)
|
scripts/joint_v6_README.md
ADDED
|
@@ -0,0 +1,97 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Recovered audio-loss trainer: joint_v6.py
|
| 2 |
+
|
| 3 |
+
This is the original research script recovered on 2026-09-19 from
|
| 4 |
+
`/workspace/tok/full/joint_v6.py`. It is the trainer used for the v6–v9 audio-loss
|
| 5 |
+
weight sweep, not the later reference-to-song adapter trainer. The file is
|
| 6 |
+
byte-for-byte preserved; no training logic or paths were changed for this release.
|
| 7 |
+
Original SHA-256: `07a53c28262cfcda90ee1a0d37facc82ad6a4de4d5c448305f4ae69202836ca9`.
|
| 8 |
+
|
| 9 |
+
Validation for this recovery was source-hash equality, Python syntax checking,
|
| 10 |
+
and lookup-table shape/range checks. The original GPU experiment was not rerun
|
| 11 |
+
as part of publication; this is not a newly validated portable training package.
|
| 12 |
+
|
| 13 |
+
## Objective and sweep
|
| 14 |
+
|
| 15 |
+
The script jointly trains the MERT-feature tokenizer head and the acoustic
|
| 16 |
+
decoder LoRA. Real-audio flow loss backpropagates through straight-through token
|
| 17 |
+
embeddings. A frozen VAE decoder adds log-mel, multi-resolution STFT, and stereo
|
| 18 |
+
width losses on an estimated clean latent crop. Minted examples supply head
|
| 19 |
+
soft-label CE and 25% of decoder flow windows.
|
| 20 |
+
|
| 21 |
+
The same script runs every audio-loss variant:
|
| 22 |
+
|
| 23 |
+
| Run name | `AUX_W` |
|
| 24 |
+
| --- | ---: |
|
| 25 |
+
| joint_v6 | 0.3 |
|
| 26 |
+
| joint_v7 | 1.0 |
|
| 27 |
+
| joint_v8 | 2.0 |
|
| 28 |
+
| joint_v9 | 4.0 |
|
| 29 |
+
|
| 30 |
+
**`AUX_W` defaults to zero**, so set it explicitly to enable the waveform loss.
|
| 31 |
+
Other defaults: `AUX_TMAX=0.4`, `AUX_FR=150`, `AUX_MARGIN=25`, 512-frame training
|
| 32 |
+
windows, and 25 frames/second. The waveform term applies to real windows only
|
| 33 |
+
when sampled flow time is at most `AUX_TMAX`. It can also be skipped when the
|
| 34 |
+
original audio cannot be resolved or the requested crop is too short.
|
| 35 |
+
|
| 36 |
+
## Original environment and required inputs
|
| 37 |
+
|
| 38 |
+
Use the older `yue2-infer` runtime described in the repository README (commit
|
| 39 |
+
`92a73cc7`, original Python 3.12 / torch 2.10 / CUDA 12.8 environment), plus numpy,
|
| 40 |
+
soundfile, torchaudio, and ffmpeg. Current YuE2 APIs may differ. Do not replace a
|
| 41 |
+
working modern training environment with these older dependencies.
|
| 42 |
+
|
| 43 |
+
The original script has hard-coded paths. Recreate them or edit a working copy:
|
| 44 |
+
|
| 45 |
+
- `/workspace/hf/hub/models--m-a-p--YuE2-3B/snapshots/...` and the corresponding
|
| 46 |
+
`YuE2-Vae` snapshot. Both must already exist. The script takes the first glob
|
| 47 |
+
match; isolate one intended snapshot to avoid accidental version selection.
|
| 48 |
+
- `/workspace/tok/full/feats/<minted_id>.npy`: minted MERT L20 features, either
|
| 49 |
+
`[T,1024]` or the legacy `[4,T,1024]` layout (index 3 is selected).
|
| 50 |
+
- `/workspace/yue2-corpus/tracks/<minted_id>/`: `semantic.npy`, `latent.npy`, and
|
| 51 |
+
`request.json` with `style`, `lyrics`, and `seed`. The small AR regularizer pack
|
| 52 |
+
alone is insufficient for this trainer: full features and latents are needed.
|
| 53 |
+
- Copy `assets/sem_nbr_idx.npy` and `assets/sem_nbr_cos.npy` from this repository
|
| 54 |
+
to `/workspace/tok/full/`. These are semantic-token neighbor lookup tables,
|
| 55 |
+
not song data. Their exact hashes are in `joint_v6_release.json`.
|
| 56 |
+
- `RP` points to prepared real recordings. Each child directory needs `mert.npy`
|
| 57 |
+
`[T,1024]`, `lat.npy` `[T,64]`, and integer `prefix.npy`. Training songs require
|
| 58 |
+
at least 512 frames. Set `HOLD` to existing prepared directory names, separated
|
| 59 |
+
by commas; supply at least one sufficiently long held-out recording.
|
| 60 |
+
- Edit `audio_path()` to map each prepared name to its exact original audio.
|
| 61 |
+
The historical mapping recognizes `coheed__<name>` and other `<tag>__<name>`
|
| 62 |
+
layouts. Without a working mapping the auxiliary audio loss may silently
|
| 63 |
+
disappear. Check nonzero `aux` values on eligible real updates; legitimate
|
| 64 |
+
zero values also occur for minted windows and larger flow times.
|
| 65 |
+
- Create `/workspace/tok/full/listen_real/` for final renders. Run names should
|
| 66 |
+
be unique: this historical script does not implement safe optimizer resume.
|
| 67 |
+
|
| 68 |
+
## Checkpoint format and example invocation
|
| 69 |
+
|
| 70 |
+
Unlike the other adapted scripts in this repository, the recovered original
|
| 71 |
+
calls `torch.load` and expects `.pt` dictionaries. For a safetensors release,
|
| 72 |
+
convert a trusted checkpoint with the existing `ckpt_io.py` helper first:
|
| 73 |
+
|
| 74 |
+
```python
|
| 75 |
+
import torch
|
| 76 |
+
from ckpt_io import load_ckpt
|
| 77 |
+
torch.save(load_ckpt("tokenizer_head_v5_30k.safetensors"), "head_init.pt")
|
| 78 |
+
```
|
| 79 |
+
|
| 80 |
+
The historical sweep started from the pretrained head corresponding to
|
| 81 |
+
`tokenizer_head_v5_30k.safetensors` and `nar_lora_joint_v4.pt`. All variants used
|
| 82 |
+
the same initial weights, rather than chaining v6 into v7 into v8 into v9.
|
| 83 |
+
For an already prepared custom dataset, after resolving the paths above:
|
| 84 |
+
|
| 85 |
+
```sh
|
| 86 |
+
mkdir -p /workspace/tok/full/listen_real
|
| 87 |
+
RP=/path/to/prepared_real HOLD=heldout_song_1,heldout_song_2 \
|
| 88 |
+
MINTED_CAP=4000 AUX_W=4.0 AUX_TMAX=0.4 AUX_FR=150 \
|
| 89 |
+
python scripts/joint_v6.py my_audio_loss_run 3000 1 1 \
|
| 90 |
+
/path/to/head_init.pt /path/to/nar_lora_joint_v4.pt 32
|
| 91 |
+
```
|
| 92 |
+
|
| 93 |
+
This example is a historical recipe, not a claim that 3,000 steps or those
|
| 94 |
+
hyperparameters are optimal for a new dataset. `best` is selected by held-out
|
| 95 |
+
latent flow loss, not the auxiliary spectral metric or a listening score.
|
| 96 |
+
The outputs include head/LoRA best and last weights and held-out reconstructions;
|
| 97 |
+
they are not new-song artist-conditioning demonstrations.
|
scripts/joint_v6_release.json
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"recovered_date": "2026-09-19",
|
| 3 |
+
"source_path": "/workspace/tok/full/joint_v6.py",
|
| 4 |
+
"training_logic_changed": false,
|
| 5 |
+
"validation": [
|
| 6 |
+
"source SHA256 matches recovered original",
|
| 7 |
+
"Python syntax",
|
| 8 |
+
"lookup arrays shape/range/finite",
|
| 9 |
+
"credential pattern scan"
|
| 10 |
+
],
|
| 11 |
+
"gpu_training_rerun": false,
|
| 12 |
+
"files": [
|
| 13 |
+
{
|
| 14 |
+
"path": "README.md",
|
| 15 |
+
"bytes": 12028,
|
| 16 |
+
"sha256": "05d63d2479f5a65956505b0dedfa3cc39f473f73bcd7d1dc55dd4a12a74c8ee2"
|
| 17 |
+
},
|
| 18 |
+
{
|
| 19 |
+
"path": "assets/sem_nbr_cos.npy",
|
| 20 |
+
"bytes": 2097280,
|
| 21 |
+
"sha256": "cf544760a9ef2f66b0bba4fa6ac3d055c628fc109cea47ba9bc6de59c767ad12"
|
| 22 |
+
},
|
| 23 |
+
{
|
| 24 |
+
"path": "assets/sem_nbr_idx.npy",
|
| 25 |
+
"bytes": 2097280,
|
| 26 |
+
"sha256": "081d5f3d41ac2b604ee95db53f5c2069b33d871c317b611d7df8e02b42eda141"
|
| 27 |
+
},
|
| 28 |
+
{
|
| 29 |
+
"path": "scripts/joint_v6.py",
|
| 30 |
+
"bytes": 18364,
|
| 31 |
+
"sha256": "07a53c28262cfcda90ee1a0d37facc82ad6a4de4d5c448305f4ae69202836ca9"
|
| 32 |
+
},
|
| 33 |
+
{
|
| 34 |
+
"path": "scripts/joint_v6_README.md",
|
| 35 |
+
"bytes": 5125,
|
| 36 |
+
"sha256": "0672cc32e1ddf89a231c820b11b0bb5a90b61572512fd95398bb7663b31ea345"
|
| 37 |
+
}
|
| 38 |
+
]
|
| 39 |
+
}
|