{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# v27 — JEPA-AASIST: đổi backend (notebook độc lập, chỉ train)\n",
    "\n",
    "Chạy trên phiên trống. Không cần tải tập test nào. Dựng mel từ ASVspoof2019 train, huấn luyện ba hệ, lưu checkpoint sang Drive. Chấm điểm để `v24` / `v26` làm.\n",
    "\n",
    "## Vì sao\n",
    "\n",
    "Trong bảng arena có đúng **một cặp cùng front-end, khác backend**:\n",
    "\n",
    "| | ADD22-T1 | ASVspoof19 |\n",
    "|---|---|---|\n",
    "| wav2vec2 + ECAPA | 46,43 | 29,69 |\n",
    "| wav2vec2 + AASIST | 31,04 | 0,22 |\n",
    "| **chênh do backend** | **15,39** | **29,47** |\n",
    "\n",
    "Hệ hiện tại của ta dùng gộp thống kê chú ý — cùng loại với ECAPA, tức đang ở vế thua. Sáu hệ xếp trên ta đều dùng backend chuyên dụng.\n",
    "\n",
    "## Thay đổi\n",
    "\n",
    "| | hệ cũ | JEPA-AASIST |\n",
    "|---|---|---|\n",
    "| patch | (8, 32) → lưới 32×4 | **(16, 16) → lưới 16×8** |\n",
    "| trọng số patch | nội suy song khối từ 16×16 | **dùng nguyên bản gốc** |\n",
    "| backend | gộp thống kê chú ý + MLP | **đồ thị chú ý AASIST** |\n",
    "\n",
    "Phải đổi patch vì AASIST dựng đồ thị trên trục tần số rồi gộp tỉ lệ 0,5. Với 4 nút tần số thì còn 2 — suy biến. Lưới 16×8 cho 8 nút tần số và 16 nút thời gian.\n",
    "\n",
    "Hai biến cùng đổi, nên **đây là hệ mới**, không phải ablation của hệ cũ. So sánh có kiểm soát JEPA / MAE / ngẫu nhiên làm **bên trong** hệ mới, tất cả dùng chung cấu hình.\n",
    "\n",
    "Chi phí: mel ~20 phút (có cache), 3 lượt train × ~15 phút. Khoảng 1,2 giờ."
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 0 · Checkpoint Audio-JEPA có chứa predictor không? (rẻ, chạy trước)\n",
    "import subprocess, sys, collections\n",
    "subprocess.run([sys.executable,\"-m\",\"pip\",\"install\",\"-q\",\"huggingface_hub\"],check=False)\n",
    "import torch\n",
    "from huggingface_hub import hf_hub_download\n",
    "_raw=torch.load(hf_hub_download(\"ltuncay/Audio-JEPA\",\"JEPA.ckpt\"),\n",
    "                map_location=\"cpu\",weights_only=False)\n",
    "_sd=_raw.get(\"state_dict\",_raw)\n",
    "pref=collections.Counter(k.split(\".\")[0] for k in _sd)\n",
    "print(\"các nhánh trong checkpoint:\",dict(pref))\n",
    "for p in sorted(pref):\n",
    "    ks=[k for k in _sd if k.startswith(p+\".\")]\n",
    "    print(f\"\\n--- {p}  ({len(ks)} tensor) ---\")\n",
    "    for k in ks[:5]: print(\"   \",k,tuple(_sd[k].shape))\n",
    "has_pred=any(p.startswith(\"predictor\") for p in pref)\n",
    "print(\"\\n\"+\"=\"*62)\n",
    "print(\"CÓ predictor -> hướng 'sai số dự đoán làm điểm phát hiện' KHẢ THI\"\n",
    "      if has_pred else \"KHÔNG có predictor -> hướng đó cần tiền huấn luyện lại\")\n",
    "print(\"=\"*62)"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 1 · Môi trường\n",
    "import os, glob, json, time, math, random, shutil, io, types\n",
    "subprocess.run([sys.executable,\"-m\",\"pip\",\"install\",\"-q\",\"timm\",\"soundfile\",\n",
    "  \"datasets>=2.19,<4.0\",\"scikit-learn\",\"pytorch_lightning\"],check=False)\n",
    "import numpy as np, 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\"; ARENA=\"/content/arena/speech_df_arena\"\n",
    "for url,d in [(\"https://github.com/LudovicTuncay/Audio-JEPA.git\",REPO),\n",
    "              (\"https://github.com/Speech-Arena/speech_df_arena.git\",ARENA)]:\n",
    "    if not os.path.isdir(d):\n",
    "        os.makedirs(os.path.dirname(d),exist_ok=True)\n",
    "        subprocess.run([\"git\",\"clone\",\"--depth\",\"1\",\"-q\",url,d],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",
    "\n",
    "SR_MODEL,N_MELS,T_BINS=32_000,128,256\n",
    "EMBED_DIM,DEPTH,CROP_S=768,12,2.56\n",
    "PATCH=(16,16)                                   # gốc Audio-JEPA -> lưới 16x8\n",
    "EPOCHS,BATCH,LR_HEAD,WD=6,32,1e-3,0.05\n",
    "LR={\"jepa\":3e-5,\"mae\":3e-4,\"random\":1e-4}\n",
    "MAE_HUB=\"hf_hub:gaunernst/vit_base_patch16_1024_128.audiomae_as2m\"\n",
    "_ENC_SD={k[8:]:v for k,v in _sd.items()\n",
    "         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",
    "    assert tuple(w.shape[-2:])==tuple(PATCH), \"patch (16,16) là gốc, không cần nội suy\"\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",
    "_e=build_encoder(\"jepa\")\n",
    "NT,NF=_e.patch_embed.num_patches_h,_e.patch_embed.num_patches_w\n",
    "print(f\"✓ lưới {NT} (thời gian) x {NF} (tần số) = {_e.patch_embed.num_patches} token\")\n",
    "assert NF>=8, f\"chỉ {NF} nút tần số — đồ thị phổ của AASIST sẽ suy biến\""
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 2 · Backend AASIST đặt lên token của encoder\n",
    "sys.path.insert(0,ARENA)\n",
    "import importlib.util\n",
    "_sp=importlib.util.spec_from_file_location(\"aa_src\",f\"{ARENA}/Models/aasist.py\")\n",
    "_aa=importlib.util.module_from_spec(_sp); _sp.loader.exec_module(_aa)\n",
    "GAT,HtrgGAT,GPool = _aa.GraphAttentionLayer,_aa.HtrgGraphAttentionLayer,_aa.GraphPool\n",
    "GD=[64,32]; PR=[0.5,0.7,0.5,0.5]; TP=[2.0,2.0,100.0,100.0]   # cấu hình AASIST gốc\n",
    "\n",
    "class JepaAASIST(nn.Module):\n",
    "    \"\"\"Encoder ViT -> (B,C,F,T) -> phần đồ thị của AASIST, nguyên bản.\"\"\"\n",
    "    def __init__(s,enc,n_t,n_f,c=64):\n",
    "        super().__init__(); s.enc=enc; s.n_t,s.n_f=n_t,n_f\n",
    "        s.layer_w=nn.Parameter(torch.zeros(DEPTH+1))\n",
    "        s.proj=nn.Conv2d(EMBED_DIM,c,kernel_size=1)      # 768 -> 64 kênh\n",
    "        s.bn=nn.BatchNorm2d(c); s.selu=nn.SELU(inplace=True)\n",
    "        s.pos_S=nn.Parameter(torch.randn(1,n_f,c))\n",
    "        s.master1=nn.Parameter(torch.randn(1,1,GD[0]))\n",
    "        s.master2=nn.Parameter(torch.randn(1,1,GD[0]))\n",
    "        s.GAT_S=GAT(c,GD[0],temperature=TP[0]); s.GAT_T=GAT(c,GD[0],temperature=TP[1])\n",
    "        s.HST11=HtrgGAT(GD[0],GD[1],temperature=TP[2]); s.HST12=HtrgGAT(GD[1],GD[1],temperature=TP[2])\n",
    "        s.HST21=HtrgGAT(GD[0],GD[1],temperature=TP[2]); s.HST22=HtrgGAT(GD[1],GD[1],temperature=TP[2])\n",
    "        s.pool_S=GPool(PR[0],GD[0],0.3); s.pool_T=GPool(PR[1],GD[0],0.3)\n",
    "        s.pool_hS1=GPool(PR[2],GD[1],0.3); s.pool_hT1=GPool(PR[2],GD[1],0.3)\n",
    "        s.pool_hS2=GPool(PR[2],GD[1],0.3); s.pool_hT2=GPool(PR[2],GD[1],0.3)\n",
    "        s.drop=nn.Dropout(0.5); s.drop_way=nn.Dropout(0.2)\n",
    "        s.out_layer=nn.Linear(5*GD[1],2)\n",
    "\n",
    "    def _tokens(s,x):\n",
    "        h=s.enc.patch_embed(x)\n",
    "        h=h+s.enc.interpolate_pos_encoding(h,s.enc.pos_embed,cls_token=False)\n",
    "        o=[]\n",
    "        for blk in s.enc.blocks: h=blk(h); o.append(h)\n",
    "        o.append(s.enc.norm(h))\n",
    "        w=torch.softmax(s.layer_w,0)\n",
    "        return sum(w[i]*o[i] for i in range(len(o)))       # (B, n_t*n_f, 768)\n",
    "\n",
    "    def forward(s,x):\n",
    "        B=x.shape[0]\n",
    "        t=s._tokens(x)\n",
    "        assert t.shape[1]==s.n_t*s.n_f, (t.shape, s.n_t, s.n_f)\n",
    "        e=t.view(B,s.n_t,s.n_f,EMBED_DIM).permute(0,3,2,1)  # (B,C,F,T)\n",
    "        e=s.selu(s.bn(s.proj(e)))\n",
    "        e_S,_=torch.max(torch.abs(e),dim=3); e_S=e_S.transpose(1,2)+s.pos_S\n",
    "        out_S=s.pool_S(s.GAT_S(e_S))\n",
    "        e_T,_=torch.max(torch.abs(e),dim=2); e_T=e_T.transpose(1,2)\n",
    "        out_T=s.pool_T(s.GAT_T(e_T))\n",
    "        m1=s.master1.expand(B,-1,-1); m2=s.master2.expand(B,-1,-1)\n",
    "        oT1,oS1,m1=s.HST11(out_T,out_S,master=m1)\n",
    "        oS1=s.pool_hS1(oS1); oT1=s.pool_hT1(oT1)\n",
    "        a,b,c_=s.HST12(oT1,oS1,master=m1); oT1,oS1,m1=oT1+a,oS1+b,m1+c_\n",
    "        oT2,oS2,m2=s.HST21(out_T,out_S,master=m2)\n",
    "        oS2=s.pool_hS2(oS2); oT2=s.pool_hT2(oT2)\n",
    "        a,b,c_=s.HST22(oT2,oS2,master=m2); oT2,oS2,m2=oT2+a,oS2+b,m2+c_\n",
    "        oT1,oT2,oS1,oS2=map(s.drop_way,(oT1,oT2,oS1,oS2))\n",
    "        m1,m2=s.drop_way(m1),s.drop_way(m2)\n",
    "        oT=torch.max(oT1,oT2); oS=torch.max(oS1,oS2); mm=torch.max(m1,m2)\n",
    "        Tm,_=torch.max(torch.abs(oT),dim=1); Ta=torch.mean(oT,dim=1)\n",
    "        Sm,_=torch.max(torch.abs(oS),dim=1); Sa=torch.mean(oS,dim=1)\n",
    "        h=torch.cat([Tm,Ta,Sm,Sa,mm.squeeze(1)],dim=1)\n",
    "        return s.out_layer(s.drop(h))\n",
    "\n",
    "_m=JepaAASIST(build_encoder(\"random\"),NT,NF).to(DEVICE)\n",
    "with torch.no_grad(): _o=_m(torch.randn(2,1,T_BINS,N_MELS,device=DEVICE))\n",
    "n_enc=sum(p.numel() for p in _m.enc.parameters())\n",
    "n_bk=sum(p.numel() for p in _m.parameters())-n_enc\n",
    "print(f\"✓ chạy thử {tuple(_o.shape)} | encoder {n_enc/1e6:.1f}M | backend {n_bk/1e6:.2f}M\")\n",
    "del _m; torch.cuda.empty_cache()"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 3 · Mel ASVspoof train (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 mel\")\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); np.savez(mp,y=Y_TR,sys=S_TR,n=i)\n",
    "    print(f\"xong {i} | {time.perf_counter()-t0:.0f}s\")\n",
    "assert len(Y_TR)>=25_000 and int(Y_TR.sum())==2580\n",
    "print(\"bonafide\",int(Y_TR.sum()),\"| spoof\",int((1-Y_TR).sum()))"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 4 · Huấn luyện ba hệ (holdout theo attack A05)\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",
    "class DS(Dataset):\n",
    "    def __init__(s,idx,y,train=False): s.idx,s.y,s.train=idx,y,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),int(s.y[i])\n",
    "\n",
    "rng=np.random.default_rng(0)\n",
    "i_ho=np.where(S_TR==\"A05\")[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",
    "dl_tr=DataLoader(DS(IDX_TR,Y_TR,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,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",
    "print(f\"train {len(IDX_TR)} | holdout A05 {len(IDX_VA)}\")\n",
    "\n",
    "def train_one(source,seed,tag):\n",
    "    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)\n",
    "    mdl=JepaAASIST(build_encoder(source),NT,NF).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,y in dl_tr:\n",
    "            x,y=x.to(DEVICE,non_blocking=True),y.to(DEVICE)\n",
    "            op.zero_grad(set_to_none=True)\n",
    "            with torch.amp.autocast(\"cuda\"): loss=crit(mdl(x),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,y in dl_va:\n",
    "                with torch.amp.autocast(\"cuda\"): lo=mdl(x.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\"); print(f\"  [lưu] {tag}.pth -> Drive\")\n",
    "    del mdl; torch.cuda.empty_cache(); return e\n",
    "\n",
    "R={}\n",
    "for src,tag in [(\"jepa\",\"jaasist_jepa_s1\"),(\"mae\",\"jaasist_mae_s1\"),\n",
    "                (\"random\",\"jaasist_rand_s1\")]:\n",
    "    if os.path.exists(f\"{DRIVE}/{tag}.pth\"): print(f\"[bỏ qua] {tag}\"); continue\n",
    "    R[tag]=train_one(src,1234,tag)\n",
    "print(\"\\n=== A05 epoch cuối ===\")\n",
    "for k,v in R.items(): print(f\"   {k:20s} {v:6.2f}\")\n",
    "print(\"\\nA05 chỉ theo dõi, KHÔNG dùng chọn siêu tham số.\")\n",
    "print(\"Chấm điểm: chạy v24 (ASVspoof19) và v26 (ADD22-T1) với các tag jaasist_*.\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Sau khi xong\n",
    "\n",
    "Trên Drive có `jaasist_jepa_s1.pth`, `jaasist_mae_s1.pth`, `jaasist_rand_s1.pth`.\n",
    "\n",
    "**Lưu ý khi chấm:** `v24`/`v26` viết file model từ `MODEL_SOURCE` của hệ **cũ** (gộp thống kê). Phải thay bằng lớp `JepaAASIST` với `PATCH=(16,16)`, nếu không `load_state_dict` sẽ báo thừa/thiếu hàng loạt — và đó là dấu hiệu tốt, nó chặn ngay chứ không cho ra số sai.\n",
    "\n",
    "Ba khả năng khi có kết quả:\n",
    "\n",
    "| | nghĩa là |\n",
    "|---|---|\n",
    "| khá lên nhiều (≥10 điểm) | backend đúng là nút thắt; JEPA-AASIST thành hệ chính của bài, lập luận 86 M so với 318 M bắt đầu có sức nặng |\n",
    "| khá lên ít | lợi thế AASIST đặc thù cho đặc trưng wav2vec2, không tổng quát — vẫn là kết quả đáng báo cáo |\n",
    "| tệ đi | đồ thị AASIST không hợp với token ViT; ghi nhận, quay lại gộp thống kê |\n",
    "\n",
    "Ô 0 trả lời câu còn lại: nếu checkpoint có nhánh `predictor` thì hướng *dùng sai số dự đoán của JEPA làm điểm phát hiện* chạy được mà không cần tiền huấn luyện lại — và đó là thứ duy nhất trong bài mà không SSL nào khác làm được."
   ]
  }
 ],
 "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
}