{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# v25 — Huấn luyện nhiều seed (notebook độc lập, chỉ train)\n",
    "\n",
    "Chạy trên phiên trống. **Không cần tải ADD22.** Chỉ dựng mel từ ASVspoof2019 train (HuggingFace) rồi huấn luyện, lưu checkpoint sang Drive. Việc chấm điểm để `v26` làm.\n",
    "\n",
    "**Tại sao cần.** Trên ADD22-T1: JEPA 41,75 · MAE 42,42 · ngẫu nhiên 45,53. Khoảng cách JEPA–MAE là **0,67 điểm với một seed duy nhất** — không kết luận được gì. Notebook này chạy thêm hai seed cho mỗi hệ.\n",
    "\n",
    "Ba khả năng sau khi có ba seed:\n",
    "\n",
    "| kết quả | nghĩa là | dùng được? |\n",
    "|---|---|---|\n",
    "| độ lệch giữa seed < 0,67 | JEPA hơn MAE, nhẹ nhưng thật | có |\n",
    "| độ lệch > 0,67 | mục tiêu SSL không đáng kể, **miền dữ liệu mới quyết định** | có, và ngược kỳ vọng của ngành nên còn đáng chú ý hơn |\n",
    "| MAE thắng ở seed khác | như trên, và ta biết trước reviewer | có |\n",
    "\n",
    "Không có khả năng nào cần giấu MAE.\n",
    "\n",
    "**Chi phí:** dựng mel ~20 phút (một lần, có cache) + 4 lần train × ~18 phút ≈ 1,5 giờ.\n",
    "\n",
    "**Lưu ý về tập validation.** Chia ngẫu nhiên 5 % ASVspoof train là vô nghĩa — chỉ có 20 người nói và 6 attack, đều xuất hiện hai bên, nên val EER = 0,00 % từ epoch 2. Notebook này dùng **holdout theo attack A05** thay thế, để ít nhất có tín hiệu thật. A05 dùng để theo dõi được, **không** dùng để chọn learning rate (Spearman với chuyển miền = −1,00)."
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 1 · Môi trường\n",
    "import os, sys, glob, json, time, math, random, shutil, subprocess, io, types\n",
    "subprocess.run([sys.executable,\"-m\",\"pip\",\"install\",\"-q\",\"timm\",\"soundfile\",\n",
    "  \"datasets>=2.19,<4.0\",\"huggingface_hub\",\"scikit-learn\"],check=False)\n",
    "import numpy as np, torch, torch.nn as nn, torch.nn.functional as F, torchaudio, soundfile as sf\n",
    "DEVICE=\"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
    "assert DEVICE==\"cuda\", \"cần GPU\"\n",
    "print(\"GPU:\",torch.cuda.get_device_name(0))\n",
    "from google.colab import drive; drive.mount(\"/content/drive\")\n",
    "DRIVE=\"/content/drive/MyDrive/jepa_spoof\"; os.makedirs(DRIVE,exist_ok=True)\n",
    "REPO=\"/content/audio-jepa\"\n",
    "if not os.path.isdir(REPO):\n",
    "    subprocess.run([\"git\",\"clone\",\"--depth\",\"1\",\"-q\",\n",
    "                    \"https://github.com/LudovicTuncay/Audio-JEPA.git\",REPO],check=False)\n",
    "\n",
    "from importlib.machinery import ModuleSpec\n",
    "from pathlib import Path\n",
    "def _shim():\n",
    "    class CpuMHA(nn.Module):\n",
    "        def __init__(s,embed_dim,num_heads,dropout=0.0,qkv_proj_bias=True,**_):\n",
    "            super().__init__(); s.num_heads=num_heads; s.head_dim=embed_dim//num_heads\n",
    "            s.dropout=dropout\n",
    "            s.qkv=nn.Linear(embed_dim,3*embed_dim,bias=qkv_proj_bias)\n",
    "            s.proj=nn.Linear(embed_dim,embed_dim,bias=True)\n",
    "        def forward(s,x):\n",
    "            b,n,d=x.shape\n",
    "            qkv=s.qkv(x).reshape(b,n,3,s.num_heads,s.head_dim)\n",
    "            q,k,v=qkv.permute(2,0,3,1,4).unbind(0)\n",
    "            o=F.scaled_dot_product_attention(q,k,v,dropout_p=s.dropout if s.training else 0.0)\n",
    "            return s.proj(o.transpose(1,2).reshape(b,n,d))\n",
    "    def mk(n):\n",
    "        m=types.ModuleType(n); m.__spec__=ModuleSpec(n,loader=None); return m\n",
    "    fa,mo,mh=mk(\"flash_attn\"),mk(\"flash_attn.modules\"),mk(\"flash_attn.modules.mha\")\n",
    "    fa.__version__=\"0.0.0-shim\"; mh.MHA=CpuMHA\n",
    "    sys.modules.update({\"flash_attn\":fa,\"flash_attn.modules\":mo,\"flash_attn.modules.mha\":mh})\n",
    "    for c in list(sys.modules):\n",
    "        if c==\"src\" or c.startswith(\"src.\"): del sys.modules[c]\n",
    "    for n in (\"src\",\"src.utils\",\"src.models\",\"src.models.components\",\"src.masks\",\n",
    "              \"src.masks.components\",\"src.data\",\"src.data.components\"):\n",
    "        m=types.ModuleType(n); m.__path__=[str(Path(REPO)/n.replace(\".\",\"/\"))]; sys.modules[n]=m\n",
    "_shim()\n",
    "if REPO not in sys.path: sys.path.insert(0,REPO)\n",
    "from src.models.components.vision_transformer import VisionTransformer\n",
    "from huggingface_hub import hf_hub_download\n",
    "\n",
    "SR_MODEL,N_MELS,T_BINS=32_000,128,256\n",
    "EMBED_DIM,DEPTH,PATCH,CROP_S=768,12,(8,32),2.56\n",
    "EPOCHS,BATCH,LR_HEAD,WD=6,32,1e-3,0.05\n",
    "LR={\"jepa\":3e-5,\"mae\":3e-4,\"random\":1e-4}      # chọn trên ASVspoof, KHÔNG đụng ADD22\n",
    "MAE_HUB=\"hf_hub:gaunernst/vit_base_patch16_1024_128.audiomae_as2m\"\n",
    "_CK=torch.load(hf_hub_download(\"ltuncay/Audio-JEPA\",\"JEPA.ckpt\"),map_location=\"cpu\",weights_only=False)\n",
    "_CK=_CK.get(\"state_dict\",_CK)\n",
    "_ENC_SD={k[8:]:v for k,v in _CK.items() if k.startswith(\"encoder.\") and not k.startswith(\"encoder_\")}\n",
    "_MAE=None\n",
    "def _mae_sd():\n",
    "    global _MAE\n",
    "    if _MAE is None:\n",
    "        import timm\n",
    "        m=timm.create_model(MAE_HUB,pretrained=True,num_classes=0)\n",
    "        _MAE={k:v for k,v in m.state_dict().items()\n",
    "              if k not in (\"cls_token\",\"pos_embed\",\"mask_token\",\"reg_token\")\n",
    "              and not k.startswith((\"head\",\"fc_norm\",\"pre_logits\"))}\n",
    "    return _MAE\n",
    "def build_encoder(source):\n",
    "    enc=VisionTransformer(input_size=(T_BINS,N_MELS),patch_size=PATCH,in_chans=1,\n",
    "                          embed_dim=EMBED_DIM,depth=DEPTH,num_heads=12,mlp_ratio=4.0,\n",
    "                          use_flash_attn=False)\n",
    "    if source==\"random\": return enc\n",
    "    sd=dict(_ENC_SD) if source==\"jepa\" else dict(_mae_sd())\n",
    "    w=sd[\"patch_embed.proj.weight\"]\n",
    "    if tuple(w.shape[-2:])!=tuple(PATCH):\n",
    "        sd[\"patch_embed.proj.weight\"]=F.interpolate(w.float(),size=PATCH,mode=\"bicubic\",align_corners=False)\n",
    "    sd.pop(\"pos_embed\",None)\n",
    "    miss,unexp=enc.load_state_dict(sd,strict=False)\n",
    "    assert unexp==[] and miss==[\"pos_embed\"], (miss,unexp[:5])\n",
    "    assert len(sd)-len(unexp)>=140\n",
    "    return enc\n",
    "print(\"✓ encoder sẵn sàng | token:\",build_encoder(\"jepa\").patch_embed.num_patches)"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 2 · Mel ASVspoof train theo quy ước arena (có cache)\n",
    "from datasets import load_dataset, Audio\n",
    "ROOT=\"/content\"; TAG=\"tr_arena\"; N_TR=25_380\n",
    "_RS=torchaudio.transforms.Resample(16_000,SR_MODEL).to(DEVICE)\n",
    "def pad_tile(x,L=64600):\n",
    "    x=np.asarray(x,dtype=np.float32)\n",
    "    if len(x)==0: return np.zeros(L,np.float32)\n",
    "    return x[:L] if len(x)>=L else np.tile(x,int(L/len(x))+1)[:L]\n",
    "def mel_from_wav16(w16):\n",
    "    w32=_RS(w16)[:, :int(CROP_S*SR_MODEL)]; out=[]\n",
    "    for i in range(w32.shape[0]):\n",
    "        v=w32[i:i+1]\n",
    "        s=torchaudio.compliance.kaldi.fbank(v-v.mean(),sample_frequency=SR_MODEL,\n",
    "            frame_length=25.0,frame_shift=10.0,num_mel_bins=N_MELS,low_freq=20,\n",
    "            high_freq=8000,use_log_fbank=True,window_type=\"hanning\")\n",
    "        n=s.shape[0]\n",
    "        if n<T_BINS: s=torch.cat([s,torch.zeros(T_BINS-n,N_MELS,device=s.device)],0)\n",
    "        out.append(s[:T_BINS])\n",
    "    return torch.stack(out).unsqueeze(1)\n",
    "def decode(a):\n",
    "    if isinstance(a,dict) and a.get(\"bytes\"):\n",
    "        x,sr=sf.read(io.BytesIO(a[\"bytes\"]),dtype=\"float32\",always_2d=False)\n",
    "    elif isinstance(a,dict) and a.get(\"array\") is not None:\n",
    "        return np.asarray(a[\"array\"],np.float32),int(a[\"sampling_rate\"])\n",
    "    elif isinstance(a,dict) and a.get(\"path\"):\n",
    "        x,sr=sf.read(a[\"path\"],dtype=\"float32\",always_2d=False)\n",
    "    else:\n",
    "        s=a.get_all_samples(); return s.data.mean(0).numpy().astype(np.float32),int(s.sample_rate)\n",
    "    return (x if x.ndim==1 else x.mean(1)).astype(np.float32),sr\n",
    "\n",
    "fp,mp=f\"{ROOT}/{TAG}.npy\",f\"{ROOT}/{TAG}_meta.npz\"\n",
    "if os.path.exists(fp) and os.path.exists(mp):\n",
    "    z=np.load(mp,allow_pickle=True); Y_TR,S_TR=z[\"y\"],z[\"sys\"]; print(\"dùng lại cache\")\n",
    "else:\n",
    "    mm=np.lib.format.open_memmap(fp,mode=\"w+\",dtype=np.float16,shape=(N_TR,T_BINS,N_MELS))\n",
    "    Y=np.zeros(N_TR,np.int64); S=[]; i=0; buf=[]; meta=[]; t0=time.perf_counter()\n",
    "    ds=load_dataset(\"Bisher/ASVspoof_2019_LA\",split=\"train\",streaming=True)\n",
    "    try: ds=ds.cast_column(\"audio\",Audio(decode=False))\n",
    "    except Exception: pass\n",
    "    def flush():\n",
    "        global i\n",
    "        if not buf: return\n",
    "        wb=torch.from_numpy(np.stack([pad_tile(v) for v in buf])).to(DEVICE)\n",
    "        with torch.no_grad(): ml=mel_from_wav16(wb).squeeze(1).cpu().numpy()\n",
    "        for k in range(len(buf)):\n",
    "            mm[i]=ml[k].astype(np.float16); Y[i]=meta[k][0]; S.append(meta[k][1]); i+=1\n",
    "        buf.clear(); meta.clear()\n",
    "    for ex in ds:\n",
    "        if i+len(buf)>=N_TR: break\n",
    "        sid=str(ex.get(\"system_id\",\"-\")).strip(); x,sr=decode(ex[\"audio\"])\n",
    "        if sr!=16_000:\n",
    "            x=torchaudio.functional.resample(torch.from_numpy(x)[None],sr,16_000).numpy().ravel()\n",
    "        buf.append(x); meta.append((1 if sid==\"-\" else 0,sid))\n",
    "        if len(buf)==64:\n",
    "            flush()\n",
    "            if i%5000<64: print(f\"   {i}/{N_TR}  {time.perf_counter()-t0:.0f}s\")\n",
    "    flush(); mm.flush(); Y_TR,S_TR=Y[:i],np.array(S)\n",
    "    np.savez(mp,y=Y_TR,sys=S_TR,n=i); print(f\"xong {i} | {time.perf_counter()-t0:.0f}s\")\n",
    "MEL=np.load(fp,mmap_mode=\"r\")\n",
    "assert len(Y_TR)>=25_000 and int(Y_TR.sum())==2580, (len(Y_TR),int(Y_TR.sum()))\n",
    "print(\"bonafide\",int(Y_TR.sum()),\"| spoof\",int((1-Y_TR).sum()),\"| attack\",sorted(set(S_TR[Y_TR==0])))"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 3 · Mô hình + vòng train (holdout theo ATTACK A05, không chia ngẫu nhiên)\n",
    "from torch.utils.data import Dataset,DataLoader\n",
    "from sklearn.metrics import roc_curve\n",
    "def eer(l,s):\n",
    "    l=np.asarray(l).astype(int); s=np.asarray(s,float)\n",
    "    if l.sum()==0 or (1-l).sum()==0: return float(\"nan\")\n",
    "    fpr,tpr,_=roc_curve(l,s); fnr=1-tpr\n",
    "    i=int(np.nanargmin(np.abs(fnr-fpr))); return float((fpr[i]+fnr[i])/2)\n",
    "assert abs(eer(np.r_[np.zeros(50),np.ones(50)],np.zeros(100))-0.5)<1e-9\n",
    "\n",
    "def forward_layers(enc,x):\n",
    "    h=enc.patch_embed(x); h=h+enc.interpolate_pos_encoding(h,enc.pos_embed,cls_token=False)\n",
    "    o=[]\n",
    "    for blk in enc.blocks: h=blk(h); o.append(h)\n",
    "    o.append(enc.norm(h)); return o\n",
    "class ASP(nn.Module):\n",
    "    def __init__(s,d,hid=128):\n",
    "        super().__init__(); s.att=nn.Sequential(nn.Linear(d,hid),nn.Tanh(),nn.Linear(hid,1))\n",
    "    def forward(s,h,mask):\n",
    "        e=s.att(h).squeeze(-1).masked_fill(mask==0,-1e4)\n",
    "        a=torch.softmax(e,1).unsqueeze(-1)\n",
    "        mu=(a*h).sum(1); sg=((a*(h-mu.unsqueeze(1))**2).sum(1)).clamp(min=1e-8).sqrt()\n",
    "        return torch.cat([mu,sg],-1)\n",
    "class JepaCM(nn.Module):\n",
    "    def __init__(s,enc):\n",
    "        super().__init__(); s.enc=enc\n",
    "        s.layer_w=nn.Parameter(torch.zeros(DEPTH+1)); s.pool=ASP(EMBED_DIM)\n",
    "        s.head=nn.Sequential(nn.LayerNorm(2*EMBED_DIM),nn.Linear(2*EMBED_DIM,256),\n",
    "                             nn.GELU(),nn.Dropout(0.2),nn.Linear(256,2))\n",
    "    def forward(s,x,m):\n",
    "        o=forward_layers(s.enc,x); w=torch.softmax(s.layer_w,0)\n",
    "        return s.head(s.pool(sum(w[i]*o[i] for i in range(len(o))),m))\n",
    "class DS(Dataset):\n",
    "    def __init__(s,idx,y,n_tok,train=False):\n",
    "        s.idx,s.y,s.n_tok,s.train=idx,y,n_tok,train; s.mm=None\n",
    "    def __len__(s): return len(s.idx)\n",
    "    def __getitem__(s,j):\n",
    "        if s.mm is None: s.mm=np.load(fp,mmap_mode=\"r\")\n",
    "        i=int(s.idx[j]); a=torch.from_numpy(np.asarray(s.mm[i],dtype=np.float32))\n",
    "        if s.train:\n",
    "            t0=np.random.randint(0,T_BINS-16); a[t0:t0+np.random.randint(0,17)]=0\n",
    "            f0=np.random.randint(0,N_MELS-16); a[:,f0:f0+np.random.randint(0,17)]=0\n",
    "        return a.unsqueeze(0),torch.ones(s.n_tok),int(s.y[i])\n",
    "\n",
    "HOLD=\"A05\"                                   # holdout theo ATTACK, không ngẫu nhiên\n",
    "rng=np.random.default_rng(0)\n",
    "i_ho=np.where(S_TR==HOLD)[0]; bo=np.where(Y_TR==1)[0].copy(); rng.shuffle(bo)\n",
    "IDX_VA=np.concatenate([i_ho,bo[:600]]); IDX_TR=np.setdiff1d(np.arange(len(Y_TR)),IDX_VA)\n",
    "print(f\"train {len(IDX_TR)} (attack {sorted(set(S_TR[IDX_TR][Y_TR[IDX_TR]==0]))}) | \"\n",
    "      f\"holdout {len(IDX_VA)} (A05 + 600 bonafide)\")\n",
    "\n",
    "N_TOK=build_encoder(\"random\").patch_embed.num_patches\n",
    "dl_tr=DataLoader(DS(IDX_TR,Y_TR,N_TOK,True),batch_size=BATCH,shuffle=True,\n",
    "                 num_workers=2,pin_memory=True,drop_last=True)\n",
    "dl_va=DataLoader(DS(IDX_VA,Y_TR,N_TOK,False),batch_size=64,num_workers=2)\n",
    "nb_,ns_=int(Y_TR[IDX_TR].sum()),int((1-Y_TR[IDX_TR]).sum())\n",
    "crit=nn.CrossEntropyLoss(weight=torch.tensor([1.0,ns_/max(nb_,1)],device=DEVICE))\n",
    "\n",
    "def train_one(source,seed,tag):\n",
    "    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)\n",
    "    mdl=JepaCM(build_encoder(source)).to(DEVICE)\n",
    "    hp=[p for n_,p in mdl.named_parameters() if not n_.startswith(\"enc.\")]\n",
    "    op=torch.optim.AdamW([{\"params\":mdl.enc.parameters(),\"lr\":LR[source]},\n",
    "                          {\"params\":hp,\"lr\":LR_HEAD}],weight_decay=WD)\n",
    "    sc=torch.optim.lr_scheduler.OneCycleLR(op,max_lr=[LR[source],LR_HEAD],\n",
    "                                           total_steps=EPOCHS*len(dl_tr),pct_start=0.1)\n",
    "    sca=torch.amp.GradScaler(\"cuda\")\n",
    "    print(f\"\\n=== {tag}  source={source} seed={seed} lr={LR[source]:g} ===\")\n",
    "    for ep in range(EPOCHS):\n",
    "        mdl.train(); t0=time.perf_counter(); tot=0\n",
    "        for x,m,y in dl_tr:\n",
    "            x,m,y=x.to(DEVICE,non_blocking=True),m.to(DEVICE),y.to(DEVICE)\n",
    "            op.zero_grad(set_to_none=True)\n",
    "            with torch.amp.autocast(\"cuda\"): loss=crit(mdl(x,m),y)\n",
    "            sca.scale(loss).backward(); sca.step(op); sca.update(); sc.step(); tot+=loss.item()\n",
    "        mdl.eval(); S=[];L=[]\n",
    "        with torch.no_grad():\n",
    "            for x,m,y in dl_va:\n",
    "                with torch.amp.autocast(\"cuda\"): lo=mdl(x.to(DEVICE),m.to(DEVICE))\n",
    "                lo=lo.float(); S.append((lo[:,1]-lo[:,0]).cpu().numpy()); L.append(y.numpy())\n",
    "        e=100*eer(np.concatenate(L),np.concatenate(S))\n",
    "        print(f\"  epoch {ep+1}/{EPOCHS} | loss {tot/len(dl_tr):.4f} | \"\n",
    "              f\"A05 EER {e:.2f}% | {time.perf_counter()-t0:.0f}s\")\n",
    "    torch.save(mdl.state_dict(),f\"{DRIVE}/{tag}.pth\")\n",
    "    print(f\"  [lưu] {tag}.pth -> Drive\")\n",
    "    del mdl; torch.cuda.empty_cache()\n",
    "    return e"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 4 · Chạy 4 lượt (seed 2 và 3 cho JEPA và MAE)\n",
    "PLAN=[(\"jepa\",2,\"audio_jepa_s2\"),(\"mae\",2,\"audio_mae_s2\"),\n",
    "      (\"jepa\",3,\"audio_jepa_s3\"),(\"mae\",3,\"audio_mae_s3\")]\n",
    "final={}\n",
    "for src,sd_,tag in PLAN:\n",
    "    if os.path.exists(f\"{DRIVE}/{tag}.pth\"):\n",
    "        print(f\"[bỏ qua] {tag} đã có trên Drive\"); continue\n",
    "    final[tag]=train_one(src,sd_,tag)\n",
    "print(\"\\n=== A05 epoch cuối ===\")\n",
    "for k,v in final.items(): print(f\"   {k:18s} {v:6.2f}\")\n",
    "print(\"\\nA05 chỉ để theo dõi. KHÔNG dùng chọn learning rate\")\n",
    "print(\"(Spearman với chuyển miền sang codec = -1,00).\")\n",
    "print(\"\\nChạy v26 để chấm các checkpoint này trên ADD22-T1.\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Sau khi xong\n",
    "\n",
    "Trên Drive sẽ có: `audio_jepa.pth`, `audio_mae.pth`, `audio_rand.pth` (seed 1234 từ v23) cộng bốn file `_s2` / `_s3`.\n",
    "\n",
    "Chạy `v26` để chấm tất cả trên ADD22-T1, rồi tính độ lệch chuẩn giữa các seed. So sánh nó với khoảng cách 0,67 giữa JEPA và MAE — đó là toàn bộ mục đích của notebook này."
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python",
   "version": "3.11"
  },
  "colab": {
   "provenance": [],
   "toc_visible": true
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}