"""AR dataset in score-free (cot=off) format: prefix ids (instructions+tags+lyrics) + codec tokens.
Artist tracks: tokens from the joint_v1 head. Minted: true semantic tokens. -> /workspace/real/ar/dataset.pt"""
import sys, os as _os; sys.path.insert(0,_os.path.dirname(_os.path.abspath(__file__))); from ckpt_io import load_ckpt
import os, glob, json, hashlib, numpy as np, torch, torch.nn as nn, sys
os.environ.setdefault("HF_HOME","/workspace/hf")
from yue2.protocol import SongRequest, token_prefixes
from yue2.tokenization_yue2 import YuE2TextTokenizer
HEAD_CK=sys.argv[1] if len(sys.argv)>1 else "/workspace/tok/full/joint_v1/head_best.pt"
W="/workspace/tok/full"; ROOT="/workspace/yue2-corpus/tracks"; RP="/workspace/real/prep"; dev="cuda"; VOCAB=32768; WIN=512; D=512; L=8; H=8
snap=glob.glob("/workspace/hf/hub/models--m-a-p--YuE2-3B/snapshots/*")[0]; tok=YuE2TextTokenizer(snap+"/qwen.tiktoken")
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).eval(); head.load_state_dict(load_ckpt(HEAD_CK,dev)["model"])
def instnorm(x): x=x.astype(np.float32); return (x-x.mean(0))/(x.std(0)+1e-5)
@torch.no_grad()
def predict(x):
    T=len(x); out=np.zeros(T,dtype=np.int64); starts=list(range(0,max(1,T-WIN+1),WIN//2))
    if starts[-1]+WIN<T: starts.append(max(0,T-WIN))
    for s0 in starts:
        xw=x[s0:s0+WIN]; n=len(xw)
        if n<WIN: xw=np.pad(xw,((0,WIN-n),(0,0)))
        with torch.autocast("cuda",dtype=torch.bfloat16): pred=head(torch.tensor(xw[None],device=dev))[0,:n].float().argmax(-1).cpu().numpy()
        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]
    return out
data=[]
for d in sorted(glob.glob(f"{RP}/*")):
    name=os.path.basename(d); toks=predict(instnorm(np.load(f"{d}/mert.npy")))
    src="/workspace/real/artist"; cap=open(f"{src}/{name}.txt").read().split("===LYRICS===")[0].replace("Global Metadata:","").strip(); style=" ".join(cap.split())[:1500]
    fl=f"/workspace/real/artist_lyrics/{name}.lyrics.txt"; lyr=open(fl).read().strip() if os.path.exists(fl) else (open(f"{src}/{name}.lyrics.txt").read().strip() if os.path.exists(f"{src}/{name}.lyrics.txt") else "[instrumental]")
    data.append(dict(name=name, src="artist", style=style, lyrics=lyr, prefix=token_prefixes(SongRequest(style=style,lyrics=lyr,cot="off",seed=1,id="c"),tok), codec=toks.astype(np.int32)))
print("artist", len(data), "mean frames", int(np.mean([len(x["codec"]) for x in data])), flush=True)
held=lambda p: int(hashlib.md5(p.encode()).hexdigest(),16)%20==0; nm=0
for d in sorted(glob.glob(f"{ROOT}/*")):
    if not os.path.exists(f"{d}/item.json"): continue
    p=os.path.basename(d); r=json.load(open(f"{d}/request.json")); y=np.load(f"{d}/semantic.npy").astype(np.int32)
    data.append(dict(name=p, src="minted_val" if held(p) else "minted", style=r["style"], lyrics=r["lyrics"], prefix=token_prefixes(SongRequest(style=r["style"],lyrics=r["lyrics"],cot="off",seed=r["seed"],id=p),tok), codec=y)); nm+=1
print("minted", nm, flush=True); torch.save(data, "/workspace/real/ar/dataset.pt"); print("AR PREP DONE", len(data), flush=True)
