{
 "nbformat": 4,
 "nbformat_minor": 0,
 "metadata": {
  "colab": {
   "provenance": [],
   "toc_visible": true
  },
  "kernelspec": {
   "name": "python3",
   "display_name": "Python 3"
  },
  "language_info": {
   "name": "python"
  },
  "accelerator": "GPU"
 },
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# v38 — Ablation mục tiêu SSL ngay trên tập chủ lực, cấu hình đóng băng\n",
    "\n",
    "Tải lên Colab → `Runtime → Change runtime type → T4 GPU` → `Runtime → Run all`. Không cần gõ gì. **~2 giờ.**\n",
    "\n",
    "## Lỗ hổng cần bịt\n",
    "\n",
    "Bài hiện có hai nửa nằm ở hai chỗ khác nhau:\n",
    "\n",
    "| | tập | cấu hình |\n",
    "|---|---|---|\n",
    "| kết quả chủ lực | DFADD | encoder **đóng băng** |\n",
    "| ablation mục tiêu SSL | ASVspoof19 | **fine-tune** |\n",
    "\n",
    "Reviewer sẽ hỏi ngay: mục tiêu SSL đóng vai trò gì ở chính chỗ hệ mạnh nhất? Notebook này trả lời.\n",
    "\n",
    "## Thí nghiệm\n",
    "\n",
    "Ba encoder — **Audio-JEPA**, **AudioMAE**, **khởi tạo ngẫu nhiên** — cùng ViT-Base, cùng mel 2,56 s / hop 10 ms / fmax 8000 / patch (8,32), cùng back-end, cùng `lr_head` 1e-3, 6 epoch, ba seed. **Tất cả đều đóng băng encoder.** Chấm trên DFADD kèm cổng `aasist`.\n",
    "\n",
    "Một điểm mạnh về mặt phương pháp so với ablation fine-tune hiện có: ở đó mỗi nguồn dùng một learning rate encoder khác nhau (JEPA 3e-5, MAE 3e-4, ngẫu nhiên 1e-4), tức có một biến gây nhiễu. Đóng băng thì encoder không có learning rate, nên **so sánh này sạch hơn** — chỉ còn đúng một biến là trọng số tiền huấn luyện.\n",
    "\n",
    "## Chốt trước khi thấy số\n",
    "\n",
    "Nếu AudioMAE đóng băng cũng đạt tầm 8–10 trên DFADD thì **vẫn in ra**, và kết luận đổi thành: trên DFADD mục tiêu SSL không phải yếu tố quyết định. Bài vẫn đứng, chỉ là câu chuyện khác. Ghi ở đây để sau không tự lừa mình.\n",
    "\n",
    "## Ô 3 in thêm một thứ miễn phí\n",
    "\n",
    "Trọng số softmax của 13 tầng, lưu ra `layerw_*.npy`. Ranjan et al. (arXiv 2502.03559) kết luận tầng thấp phân biệt tốt nhất cho deepfake; nếu trọng số của ta cũng dồn về tầng thấp thì đó là một hình và một đoạn phân tích, không tốn thêm GPU.\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 1 · Môi trường và ba encoder\n",
    "import os, sys, glob, json, time, math, random, shutil, io, types, subprocess, gc, ast, re as _re\n",
    "subprocess.run([sys.executable,\"-m\",\"pip\",\"install\",\"-q\",\"pytorch_lightning\",\"rich\",\"librosa\",\n",
    "  \"soundfile\",\"pandas\",\"scikit-learn\",\"timm\",\"datasets>=2.19,<4.0\",\"huggingface_hub\"],check=False)\n",
    "import numpy as np, pandas as pd, torch, torch.nn as nn, torch.nn.functional as F\n",
    "import torchaudio, soundfile as sf\n",
    "DEVICE=\"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
    "assert DEVICE==\"cuda\", \"cần GPU: Runtime -> Change runtime type -> T4 GPU\"\n",
    "print(\"GPU:\",torch.cuda.get_device_name(0),flush=True)\n",
    "from google.colab import drive; drive.mount(\"/content/drive\")\n",
    "DRIVE=\"/content/drive/MyDrive/jepa_spoof\"; os.makedirs(DRIVE,exist_ok=True)\n",
    "BASE=\"/content/arena\"; CKPT=f\"{BASE}/checkpoints\"; PROTO=f\"{BASE}/protocols\"\n",
    "for d in (BASE,CKPT,PROTO): os.makedirs(d,exist_ok=True)\n",
    "os.environ[\"DF_ARENA_CHECKPOINTS_DIR\"]=CKPT; os.environ[\"DF_ARENA_PROTOCOL_FILES_DIR\"]=PROTO\n",
    "os.chdir(BASE)\n",
    "REPO=\"/content/audio-jepa\"\n",
    "for url,d in [(\"https://github.com/Speech-Arena/speech_df_arena.git\",f\"{BASE}/speech_df_arena\"),\n",
    "              (\"https://github.com/clovaai/aasist.git\",f\"{BASE}/aasist_src\"),\n",
    "              (\"https://github.com/LudovicTuncay/Audio-JEPA.git\",REPO)]:\n",
    "    if not os.path.isdir(d): subprocess.run([\"git\",\"clone\",\"--depth\",\"1\",\"-q\",url,d],check=False)\n",
    "_a=torch.load(f\"{BASE}/aasist_src/models/weights/AASIST.pth\",map_location=\"cpu\")\n",
    "_a=_a.get(\"state_dict\",_a)\n",
    "torch.save({k.replace(\"module.\",\"\"):v for k,v in _a.items()},f\"{CKPT}/aasist.pth\")\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",
    "# --- CẤU HÌNH GIỐNG HỆT v36 (bản đã cho 8,47 trên DFADD) ---\n",
    "SR_MODEL,N_MELS,T_BINS=32_000,128,256\n",
    "CLIP_S,HOP_MS,FRAME_MS=2.56,10.0,25.0\n",
    "F_MIN,F_MAX=20,8000\n",
    "EMBED_DIM,DEPTH,PATCH=768,12,(8,32)\n",
    "BATCH,LR_HEAD,WD=32,1e-3,0.05\n",
    "EPOCHS=6\n",
    "MAE_HUB=\"hf_hub:gaunernst/vit_base_patch16_1024_128.audiomae_as2m\"\n",
    "\n",
    "_CK=torch.load(hf_hub_download(\"ltuncay/Audio-JEPA\",\"JEPA.ckpt\"),map_location=\"cpu\",\n",
    "               weights_only=False); _CK=_CK.get(\"state_dict\",_CK)\n",
    "_JEPA_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",
    "\n",
    "def build_encoder(source):\n",
    "    \"\"\"source: 'jepa' | 'mae' | 'random'.  Kiến trúc y hệt nhau, chỉ khác trọng số.\"\"\"\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(_JEPA_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\",\n",
    "                                                    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\"], (source,miss,unexp[:5])\n",
    "    assert len(sd)>=140, (source,len(sd))\n",
    "    return enc\n",
    "\n",
    "for _s in (\"jepa\",\"mae\",\"random\"):\n",
    "    _e=build_encoder(_s)\n",
    "    print(f\"✓ {_s:7s} | lưới {_e.patch_embed.num_patches_h}x{_e.patch_embed.num_patches_w}\"\n",
    "          f\" = {_e.patch_embed.num_patches} token\",flush=True)\n",
    "    del _e\n",
    "gc.collect()\n",
    "print(f\"mel {CLIP_S}s | hop {HOP_MS}ms | frame {FRAME_MS}ms | fmax {F_MAX} | patch {PATCH}\",flush=True)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 2 · Mel ASVspoof train (giống hệt v36, có cache)\n",
    "from datasets import load_dataset, Audio\n",
    "ROOT=\"/content\"; TAG=\"tr_old\"; 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_old(w16):\n",
    "    w=_RS(w16)[:, :int(CLIP_S*SR_MODEL)]; out=[]\n",
    "    for i in range(w.shape[0]):\n",
    "        v=w[i:i+1]\n",
    "        f=torchaudio.compliance.kaldi.fbank(v-v.mean(),sample_frequency=SR_MODEL,\n",
    "            frame_length=FRAME_MS,frame_shift=HOP_MS,num_mel_bins=N_MELS,\n",
    "            low_freq=F_MIN,high_freq=F_MAX,use_log_fbank=True,window_type=\"hanning\")\n",
    "        n=f.shape[0]\n",
    "        if n<T_BINS: f=torch.cat([f,torch.zeros(T_BINS-n,N_MELS,device=f.device)],0)\n",
    "        out.append(f[: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",
    "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_old(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\",flush=True)\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\",flush=True)\n",
    "assert int(Y_TR.sum())==2580, int(Y_TR.sum())\n",
    "print(\"bonafide\",int(Y_TR.sum()),\"| spoof\",int((1-Y_TR).sum()),flush=True)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 3 · Train sáu lượt đóng băng: AudioMAE và ngẫu nhiên, ba seed mỗi bê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",
    "    fpr,tpr,_=roc_curve(l,s); fnr=1-tpr\n",
    "    i=int(np.nanargmin(np.abs(fnr-fpr))); return float(100*(fpr[i]+fnr[i])/2)\n",
    "def _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):\n",
    "        a=torch.softmax(s.att(h).squeeze(-1),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 Net(nn.Module):\n",
    "    def __init__(s,enc,freeze=True):\n",
    "        super().__init__(); s.enc=enc; s.freeze=freeze\n",
    "        if freeze:\n",
    "            for p in s.enc.parameters(): p.requires_grad=False\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):\n",
    "        if s.freeze:\n",
    "            s.enc.eval()\n",
    "            with torch.no_grad(): o=[t.detach() for t in _layers(s.enc,x)]\n",
    "        else: o=_layers(s.enc,x)\n",
    "        w=torch.softmax(s.layer_w,0)\n",
    "        return s.head(s.pool(sum(w[i]*o[i] for i in range(len(o)))))\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",
    "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,num_workers=2,\n",
    "                 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)} | ENCODER ĐÓNG BĂNG cho MỌI nguồn\",flush=True)\n",
    "print(\"A05 chỉ để theo dõi. KHÔNG dùng để chọn siêu tham số (Spearman -1,00).\",flush=True)\n",
    "\n",
    "def train_frozen(source,seed,tag):\n",
    "    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)\n",
    "    mdl=Net(build_encoder(source),freeze=True).to(DEVICE)\n",
    "    hp=[p for n_,p in mdl.named_parameters() if not n_.startswith(\"enc.\")]\n",
    "    n_tr=sum(p.numel() for p in hp)\n",
    "    op=torch.optim.AdamW([{\"params\":hp,\"lr\":LR_HEAD}],weight_decay=WD)\n",
    "    sc=torch.optim.lr_scheduler.OneCycleLR(op,max_lr=[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} | \"\n",
    "          f\"tham số huấn luyện {n_tr/1e6:.2f} M / 86 M ===\",flush=True)\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",
    "        sc_=np.concatenate(S); e=eer(np.concatenate(L),sc_)\n",
    "        print(f\"  epoch {ep+1}/{EPOCHS} | loss {tot/len(dl_tr):.4f} | A05 {e:.2f}% | \"\n",
    "              f\"|điểm| tb {np.abs(sc_).mean():.2f} | {time.perf_counter()-t0:.0f}s\",flush=True)\n",
    "    lw=torch.softmax(mdl.layer_w.detach().cpu(),0).numpy()\n",
    "    print(\"  trọng số 13 tầng: \"+\" \".join(f\"{v:.3f}\" for v in lw),flush=True)\n",
    "    print(f\"  tầng nặng nhất: {int(np.argmax(lw))} (0 = sau block 1, 12 = sau norm cuối)\",flush=True)\n",
    "    torch.save(mdl.state_dict(),f\"{DRIVE}/{tag}.pth\")\n",
    "    np.save(f\"{DRIVE}/layerw_{tag}.npy\",lw)\n",
    "    print(f\"  [lưu] {tag}.pth + layerw_{tag}.npy\")\n",
    "    del mdl; gc.collect(); torch.cuda.empty_cache()\n",
    "\n",
    "PLAN=[(\"mae\",1234,\"mfrz_s1\"),(\"mae\",2,\"mfrz_s2\"),(\"mae\",3,\"mfrz_s3\"),\n",
    "      (\"random\",1234,\"rfrz_s1\"),(\"random\",2,\"rfrz_s2\"),(\"random\",3,\"rfrz_s3\")]\n",
    "for src,sd_,tag in PLAN:\n",
    "    if os.path.exists(f\"{DRIVE}/{tag}.pth\"): print(f\"[bỏ qua] {tag} đã có trên Drive\"); continue\n",
    "    train_frozen(src,sd_,tag)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 4 · File model cho arena\n",
    "MODEL_SOURCE = r\"\"\"\n",
    "import os, sys, types, numpy as np, torch, torch.nn as nn, torch.nn.functional as F\n",
    "import torchaudio, pytorch_lightning as pl\n",
    "from importlib.machinery import ModuleSpec\n",
    "from pathlib import Path\n",
    "SR_MODEL, N_MELS, T_BINS = 32000, 128, 256\n",
    "CLIP_S   = 2.56\n",
    "HOP_MS   = 10.0\n",
    "FRAME_MS = 25.0\n",
    "F_MIN, F_MAX = 20, 8000\n",
    "EMBED_DIM, DEPTH, PATCH = 768, 12, (8,32)\n",
    "REPO = Path(\"/content/audio-jepa\")\n",
    "\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\",\n",
    "              \"src.masks\",\"src.masks.components\",\"src.data\",\"src.data.components\"):\n",
    "        m=types.ModuleType(n); m.__path__=[str(REPO/n.replace(\".\",\"/\"))]; sys.modules[n]=m\n",
    "_shim()\n",
    "if str(REPO) not in sys.path: sys.path.insert(0,str(REPO))\n",
    "from src.models.components.vision_transformer import VisionTransformer\n",
    "\n",
    "def _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",
    "\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):\n",
    "        a=torch.softmax(s.att(h).squeeze(-1),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",
    "\n",
    "class Net(pl.LightningModule):\n",
    "    def __init__(s, out_score_file_name=None):\n",
    "        super().__init__()\n",
    "        s.out_score_file_name=out_score_file_name\n",
    "        s.enc=VisionTransformer(input_size=(T_BINS,N_MELS),patch_size=PATCH,in_chans=1,\n",
    "                                embed_dim=EMBED_DIM,depth=DEPTH,num_heads=12,\n",
    "                                mlp_ratio=4.0,use_flash_attn=False)\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",
    "        s._rs=torchaudio.transforms.Resample(16000,SR_MODEL)\n",
    "\n",
    "    def _mel(s,w16):\n",
    "        w=s._rs.to(w16.device)(w16)\n",
    "        L=int(CLIP_S*SR_MODEL)\n",
    "        if w.shape[1]<L: w=w.repeat(1,int(L/w.shape[1])+1)\n",
    "        w=w[:, :L]\n",
    "        out=[]\n",
    "        for i in range(w.shape[0]):\n",
    "            v=w[i:i+1]\n",
    "            f=torchaudio.compliance.kaldi.fbank(\n",
    "                v-v.mean(),sample_frequency=SR_MODEL,frame_length=FRAME_MS,\n",
    "                frame_shift=HOP_MS,num_mel_bins=N_MELS,low_freq=F_MIN,high_freq=F_MAX,\n",
    "                use_log_fbank=True,window_type=\"hanning\")\n",
    "            n=f.shape[0]\n",
    "            if n<T_BINS: f=torch.cat([f,torch.zeros(T_BINS-n,N_MELS,device=f.device)],0)\n",
    "            out.append(f[:T_BINS])\n",
    "        return torch.stack(out).unsqueeze(1)\n",
    "\n",
    "    def forward(s,x):\n",
    "        if x.dim()==3: x=x.squeeze(1)\n",
    "        mel=s._mel(x.float())\n",
    "        o=_layers(s.enc,mel); w=torch.softmax(s.layer_w,0)\n",
    "        lg=s.head(s.pool(sum(w[i]*o[i] for i in range(len(o)))))\n",
    "        d=(lg[:,1]-lg[:,0]).unsqueeze(1)\n",
    "        return lg, torch.cat([-d,d],dim=1)\n",
    "\n",
    "    def test_step(s,batch,batch_idx):\n",
    "        x,utt=batch\n",
    "        _,out=s(x)\n",
    "        sc=out[:,1].data.cpu().numpy().ravel()\n",
    "        with open(s.out_score_file_name,\"a+\") as fh:\n",
    "            for f,c in zip(utt,sc.tolist()): fh.write(\"{} {}\\n\".format(f,c))\n",
    "\n",
    "def load_model(model_path,out_score_file_name):\n",
    "    m=Net(out_score_file_name)\n",
    "    sd=torch.load(model_path,map_location=\"cpu\"); sd=sd.get(\"state_dict\",sd)\n",
    "    miss,unexp=m.load_state_dict(sd,strict=False)\n",
    "    print(f\"[load] thieu {len(miss)} thua {len(unexp)}\")\n",
    "    assert not unexp, unexp[:6]\n",
    "    assert set(miss)<={\"enc.pos_embed\",\"_rs.kernel\"}, miss[:6]\n",
    "    m.out_score_file_name=out_score_file_name\n",
    "    return m\n",
    "\"\"\"\n",
    "CODE=MODEL_SOURCE\n",
    "ast.parse(CODE)\n",
    "for k in [\"CLIP_S   = 2.56\",\"F_MIN, F_MAX = 20, 8000\",\"(8,32)\"]: assert k in CODE, k\n",
    "\n",
    "WANT=[\"mfrz_s1\",\"mfrz_s2\",\"mfrz_s3\",\"rfrz_s1\",\"rfrz_s2\",\"rfrz_s3\",\n",
    "      \"jfrz_s1\",\"jfrz_s2\",\"jfrz_s3\"]\n",
    "TAGS=[t for t in WANT if os.path.exists(f\"{DRIVE}/{t}.pth\")]\n",
    "for t in TAGS:\n",
    "    shutil.copy(f\"{DRIVE}/{t}.pth\",f\"{CKPT}/{t}.pth\")\n",
    "    open(f\"{BASE}/speech_df_arena/Models/{t}.py\",\"w\").write(CODE)\n",
    "assert not [(a,b) for a in TAGS for b in TAGS if a!=b and b.startswith(a)], \"va chạm tiền tố\"\n",
    "missing=[t for t in WANT if t not in TAGS]\n",
    "if missing: print(\"[thiếu, sẽ bỏ qua]\",missing)\n",
    "print(\"sẽ chấm:\",[\"aasist\"]+TAGS,flush=True)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 5 · Chấm trên DFADD (kèm cổng aasist)\n",
    "from datasets import load_dataset as _ld\n",
    "def build_hf(name):\n",
    "    P=f\"{PROTO}/{name}.csv\"; W=f\"{BASE}/wav_{name}\"; os.makedirs(W,exist_ok=True)\n",
    "    if os.path.exists(P) and len(glob.glob(f\"{W}/*.wav\"))>10: return\n",
    "    ds=_ld(f\"SpeechAntiSpoofingBenchmarks/{name}\",split=\"test\",streaming=True)\n",
    "    rows=[]; seen={}; t0=time.perf_counter(); n=0\n",
    "    for ex in ds:\n",
    "        a=ex[\"audio\"]; x=np.asarray(a[\"array\"],dtype=np.float32); sr=int(a[\"sampling_rate\"])\n",
    "        if x.ndim>1: x=x.mean(1)\n",
    "        lb=ex[\"label\"]\n",
    "        if isinstance(lb,(int,np.integer)): lb=\"bonafide\" if int(lb)==0 else \"spoof\"\n",
    "        lb=str(lb).strip().lower()\n",
    "        lb=\"bonafide\" if lb in (\"bonafide\",\"bona-fide\",\"real\",\"genuine\",\"0\") else \"spoof\"\n",
    "        orig=os.path.basename(str(ex.get(\"path\") or f\"u{n:07d}\"))\n",
    "        stem=_re.sub(r\"[^A-Za-z0-9_.-]\",\"_\",os.path.splitext(orig)[0]) or f\"u{n:07d}\"\n",
    "        if stem in seen: seen[stem]+=1; stem=f\"{stem}__{seen[stem]}\"\n",
    "        else: seen[stem]=0\n",
    "        sf.write(f\"{W}/{stem}.wav\",x,sr,subtype=\"PCM_16\"); rows.append((f\"{W}/{stem}.wav\",lb)); n+=1\n",
    "        del x,a,ex\n",
    "        if n%2000==0: gc.collect(); print(f\"   {name} {n} | {time.perf_counter()-t0:.0f}s\",flush=True)\n",
    "    pd.DataFrame(rows,columns=[\"file_name\",\"label\"]).to_csv(P,index=False)\n",
    "    print(f\"{name}: {n} tệp | {time.perf_counter()-t0:.0f}s\",flush=True)\n",
    "    for c in [\"/root/.cache/huggingface/datasets\",\"/root/.cache/huggingface/hub\"]:\n",
    "        try: shutil.rmtree(c)\n",
    "        except Exception: pass\n",
    "\n",
    "os.chdir(f\"{BASE}/speech_df_arena\")\n",
    "env=dict(os.environ,DF_ARENA_CHECKPOINTS_DIR=CKPT,DF_ARENA_PROTOCOL_FILES_DIR=PROTO)\n",
    "name=\"DFADD\"; build_hf(name)\n",
    "d=pd.read_csv(f\"{PROTO}/{name}.csv\")\n",
    "assert len(d)>=3700, f\"protocol DFADD chỉ có {len(d)} dòng — kiểm tra lại\"\n",
    "print(f\"\\n{name}: {len(d)} dòng {dict(d.label.value_counts())}\",flush=True)\n",
    "for tag in [\"aasist\"]+TAGS:\n",
    "    t0=time.perf_counter()\n",
    "    with open(f\"/content/v38_{name}_{tag}.log\",\"w\") as fh:\n",
    "        r=subprocess.run([sys.executable,\"evaluate.py\",\"--models\",tag,\"--protocol_files\",name,\n",
    "                          \"--batch_size\",\"32\",\"--fix_length\",\"--num_workers\",\"1\",\n",
    "                          \"--device\",\"cuda\"],env=env,stdout=fh,stderr=subprocess.STDOUT)\n",
    "    print(f\"   {name:8s} {tag:10s} mã {r.returncode} | {time.perf_counter()-t0:.0f}s\",flush=True)\n",
    "    if r.returncode!=0:\n",
    "        print(open(f\"/content/v38_{name}_{tag}.log\").read()[-2000:]); break\n",
    "    gc.collect()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 6 · Bảng ablation ba dòng trên DFADD, cùng cấu hình đóng băng\n",
    "res={}\n",
    "for f in sorted(glob.glob(\"logs/summary_*.json\")):\n",
    "    for k,v in json.load(open(f)).get(\"Average\",{}).items():\n",
    "        for ds_,d in v.items(): res.setdefault(ds_,{})[k]=d[\"EER (%)\"]\n",
    "r=res.get(\"DFADD\",{})\n",
    "assert r, \"chưa có kết quả DFADD — chạy lại ô 5\"\n",
    "\n",
    "GROUPS=[(\"Audio-JEPA\",\"jfrz_\"),(\"AudioMAE\",\"mfrz_\"),(\"khởi tạo ngẫu nhiên\",\"rfrz_\")]\n",
    "print(\"=\"*64); print(\"  DFADD — encoder ĐÓNG BĂNG, chỉ 0,5 M tham số huấn luyện\")\n",
    "print(\"  cùng ViT-Base, cùng mel, cùng back-end, cùng lr đầu, 6 epoch, 3 seed\")\n",
    "print(\"=\"*64)\n",
    "aa=r.get(\"aasist\",float(\"nan\"))\n",
    "print(f\"   {'aasist (cổng)':22s} {aa:7.2f}   [công bố 41,86]\")\n",
    "print()\n",
    "summ={}\n",
    "for label,pref in GROUPS:\n",
    "    vals=[r[k] for k in sorted(r) if k.startswith(pref)]\n",
    "    if not vals: print(f\"   {label:22s}   (chưa có)\"); continue\n",
    "    m=float(np.mean(vals)); sd=float(np.std(vals,ddof=1)) if len(vals)>1 else float(\"nan\")\n",
    "    summ[label]=(m,sd,vals)\n",
    "    seeds=\" \".join(f\"{v:.2f}\" for v in vals)\n",
    "    print(f\"   {label:22s} {m:7.2f} ± {sd:4.2f}   n={len(vals)}   [{seeds}]\")\n",
    "print()\n",
    "if \"Audio-JEPA\" in summ and \"AudioMAE\" in summ:\n",
    "    mj,sj,vj=summ[\"Audio-JEPA\"]; mm_,sm,vm=summ[\"AudioMAE\"]\n",
    "    pooled=float(np.sqrt((sj**2+sm**2)/2))\n",
    "    print(f\"   MAE - JEPA = {mm_-mj:+.2f} điểm   |  d gộp = {(mm_-mj)/pooled:+.2f}\")\n",
    "    print(f\"   dải seed chồng nhau? JEPA[{min(vj):.2f},{max(vj):.2f}] \"\n",
    "          f\"MAE[{min(vm):.2f},{max(vm):.2f}] -> \"\n",
    "          f\"{'CÓ' if max(vj)>=min(vm) and max(vm)>=min(vj) else 'KHÔNG'}\")\n",
    "if \"khởi tạo ngẫu nhiên\" in summ:\n",
    "    mr=summ[\"khởi tạo ngẫu nhiên\"][0]\n",
    "    for label in (\"Audio-JEPA\",\"AudioMAE\"):\n",
    "        if label in summ:\n",
    "            print(f\"   {label} lợi so với ngẫu nhiên: {mr-summ[label][0]:+.2f} điểm\")\n",
    "print()\n",
    "lw={}\n",
    "for t in TAGS:\n",
    "    p=f\"{DRIVE}/layerw_{t}.npy\"\n",
    "    if os.path.exists(p): lw[t]=np.load(p)\n",
    "if lw:\n",
    "    print(\"   Trọng số tầng đã học (13 giá trị, softmax):\")\n",
    "    for t,v in sorted(lw.items()):\n",
    "        print(f\"     {t:9s} nặng nhất tầng {int(np.argmax(v)):2d} | \"+\" \".join(f\"{x:.2f}\" for x in v))\n",
    "np.savez(f\"{DRIVE}/res_v38_ablation_frozen_dfadd.npz\",\n",
    "         **{f\"DFADD|{k}\":v for k,v in r.items()},\n",
    "         **{f\"layerw|{k}\":v for k,v in lw.items()})\n",
    "print(\"\\n[lưu] res_v38_ablation_frozen_dfadd.npz -> Drive\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 7 · (tuỳ chọn) Chấm thêm trên ASVspoof2019 — bật cờ nếu còn thời gian\n",
    "CHAM_ASV19 = False   #@param {type:\"boolean\"}\n",
    "\n",
    "if CHAM_ASV19:\n",
    "    def build_asv19():\n",
    "        P=f\"{PROTO}/asvspoof_2019.csv\"; W=f\"{BASE}/asv19\"; os.makedirs(W,exist_ok=True)\n",
    "        if os.path.exists(P) and len(glob.glob(f\"{W}/*.wav\"))>=71_000: return\n",
    "        from datasets import load_dataset, Audio\n",
    "        def dec(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()\n",
    "                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",
    "        ds=load_dataset(\"Bisher/ASVspoof_2019_LA\",split=\"test\",streaming=True)\n",
    "        try: ds=ds.cast_column(\"audio\",Audio(decode=False))\n",
    "        except Exception: pass\n",
    "        rows=[]; t0=time.perf_counter()\n",
    "        for k,ex in enumerate(ds):\n",
    "            x,sr=dec(ex[\"audio\"])\n",
    "            if sr!=16_000:\n",
    "                x=torchaudio.functional.resample(torch.from_numpy(x)[None],sr,16_000).numpy().ravel()\n",
    "            fn=f\"{W}/LA_E_{k:07d}.wav\"\n",
    "            if not os.path.exists(fn): sf.write(fn,x,16_000,subtype=\"PCM_16\")\n",
    "            rows.append((fn,\"bonafide\" if str(ex.get(\"system_id\",\"-\")).strip()==\"-\" else \"spoof\"))\n",
    "            if (k+1)%10_000==0: print(f\"   asv19 {k+1} | {time.perf_counter()-t0:.0f}s\",flush=True)\n",
    "        pd.DataFrame(rows,columns=[\"file_name\",\"label\"]).to_csv(P,index=False)\n",
    "    os.chdir(f\"{BASE}/speech_df_arena\"); build_asv19()\n",
    "    nm=\"asvspoof_2019\"\n",
    "    d=pd.read_csv(f\"{PROTO}/{nm}.csv\"); print(f\"{nm}: {len(d)} dòng\",flush=True)\n",
    "    for tag in [\"aasist\"]+TAGS:\n",
    "        t0=time.perf_counter()\n",
    "        with open(f\"/content/v38_{nm}_{tag}.log\",\"w\") as fh:\n",
    "            r2=subprocess.run([sys.executable,\"evaluate.py\",\"--models\",tag,\"--protocol_files\",nm,\n",
    "                               \"--batch_size\",\"32\",\"--fix_length\",\"--num_workers\",\"1\",\n",
    "                               \"--device\",\"cuda\"],env=env,stdout=fh,stderr=subprocess.STDOUT)\n",
    "        print(f\"   {nm} {tag:10s} mã {r2.returncode} | {time.perf_counter()-t0:.0f}s\",flush=True)\n",
    "        if r2.returncode!=0:\n",
    "            print(open(f\"/content/v38_{nm}_{tag}.log\").read()[-2000:]); break\n",
    "        gc.collect()\n",
    "else:\n",
    "    print(\"bỏ qua ASVspoof19. Đặt CHAM_ASV19 = True nếu còn thời gian (~60-75 phút nữa).\")\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Đọc kết quả\n",
    "\n",
    "**Cổng trước tiên.** `aasist` trên DFADD phải ra khoảng **39,05** (số ta đã đo; bảng công bố là 41,86, lệch do protocol). Lệch xa hơn nghĩa là đường ống sai, đừng tin dòng nào phía dưới.\n",
    "\n",
    "**Mốc đã biết:** `jfrz_*` phải tái lập **8,47 ± 0,35** với ba seed 8,74 · 8,59 · 8,07. Nếu ra khác thì có gì đó đã đổi giữa v36 và lần chạy này — dừng lại tìm nguyên nhân trước khi đọc `mfrz_` và `rfrz_`.\n",
    "\n",
    "| kết quả | nghĩa là gì cho bài |\n",
    "|---|---|\n",
    "| MAE tệ hơn JEPA rõ rệt, dải seed không chồng | ablation chuyển hẳn về DFADD; abstract nói được \"cùng cấu hình, cùng tập\" |\n",
    "| MAE ngang JEPA | mục tiêu SSL không phải yếu tố quyết định **trên DFADD**; kết quả chủ lực phải quy cho thứ khác, và ablation ASVspoof19 giữ nguyên vai trò hiện tại |\n",
    "| ngẫu nhiên ngang MAE | củng cố câu \"tiền huấn luyện MAE gần như không đóng góp\", lần này với n=3 |\n",
    "| ngẫu nhiên cũng tốt | nghiêm trọng: DFADD có thể tách được bằng đặc trưng mel bất kỳ. Phải kiểm tra lại trước khi nộp |\n",
    "\n",
    "Dòng cuối là rủi ro thật, không phải giả định — chuyện tương tự đã xảy ra với F03 của CodecFake, nơi một encoder chưa huấn luyện đạt 0,60 %. Xem `chung/lich-su-thay-doi-va-tuyen-bo-da-rut.md` mục A1.\n",
    "\n",
    "## Sau khi chạy xong\n",
    "\n",
    "Trên Drive có `res_v38_ablation_frozen_dfadd.npz` và các `layerw_*.npy`. Mang số về hội thoại viết bài để cập nhật Bảng 2 và quyết định câu ablation trong abstract.\n"
   ]
  }
 ]
}