{
 "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": [
    "# v40 — Độ phân giải thời gian: giả thuyết giải thích cả DFADD lẫn LibriSeVoc\n",
    "\n",
    "Tải lên Colab → T4 GPU → `Run all`. ~3,5 giờ. Kết quả lưu lên Drive **sau mỗi hệ**, không đợi tới cuối.\n",
    "\n",
    "## Giả thuyết\n",
    "\n",
    "Mô hình cắt phổ thành mảnh 8 khung × 32 dải mel. Tám khung × hop 10 ms = **80 ms mỗi token**, lưới 32 bước thời gian. Thô hơn wav2vec2 bốn lần, thô hơn AASIST vài bậc.\n",
    "\n",
    "| loại tấn công | dấu vết ở đâu | ta thấy? |\n",
    "|---|---|---|\n",
    "| **vocoder** (WaveNet, MelGAN, DiffWave…) sinh **dạng sóng** | mất kết hợp pha, bất thường chu kỳ cao độ — **thang mili giây** | không, bị trung bình hoá |\n",
    "| **TTS khuếch tán / flow-matching** (GradTTS, Matcha, pflow…) sinh **phổ mel** | kết cấu phổ, bao hình tần cao — **ổn định theo thời gian** | có |\n",
    "\n",
    "Khớp với toàn bộ dữ liệu: DFADD toàn tấn công tầng phổ → ta thắng đậm (8,47 vs 41,86). LibriSeVoc toàn vocoder → ta thua (41,96 vs 37,65).\n",
    "\n",
    "Bằng chứng đã có từ v5: `patch 16x16 -> 8x32` (160 ms -> 80 ms) là can thiệp **duy nhất** chạm được A17 — tấn công chuyển giọng dựa trên lọc dạng sóng — kéo 52,63 xuống 39,58, trong khi đổi hop hay f_max không làm được.\n",
    "\n",
    "## Thí nghiệm\n",
    "\n",
    "| | lưới | token | thời gian/token | tần số |\n",
    "|---|---|---|---|---|\n",
    "| mốc (8,32) — checkpoint `audio_jepa_s*` đã có | 32x4 | 128 | 80 ms | 4 bin |\n",
    "| **biến (4,32)** — huấn luyện mới | **64x4** | **256** | **40 ms** | **4 bin, giữ nguyên** |\n",
    "\n",
    "Chỉ đổi trục thời gian. v6 ngày xưa giữ 128 token cố định nên (4,64) làm tần số tụt còn 2 bin — lẫn hai biến. Ở đây chấp nhận gấp đôi token để tách sạch.\n",
    "\n",
    "Mọi thứ khác giữ y hệt cấu hình chính: mel 2,56 s / hop 10 ms / fmax 8000, fine-tune đầy đủ, lr encoder 3e-5, head 1e-3, 6 epoch, ba seed.\n",
    "\n",
    "## Dự đoán có thể sai được\n",
    "\n",
    "**LibriSeVoc cải thiện rõ. DFADD gần như không đổi.**\n",
    "\n",
    "Nếu cả hai cùng cải thiện thì đó chỉ là \"nhiều token hơn thì tốt hơn\", không phải cơ chế. Nếu LibriSeVoc không đổi thì giả thuyết sai và ta bỏ nó trước khi in ra giấy.\n",
    "\n",
    "Chấm lại cả `audio_jepa_s*` trong cùng lượt để so cùng điều kiện, và `aasist` làm cổng.\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 1 · Môi trường và encoder hai kích thước patch\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\",\"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/arena40\"; 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",
    "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=768,12\n",
    "PATCH_BASE=(8,32)      # moc, da co checkpoint audio_jepa_s*\n",
    "PATCH_NEW =(4,32)      # bien: min gap doi theo thoi gian, tan so giu nguyen\n",
    "BATCH,LR_ENC,LR_HEAD,WD=32,3e-5,1e-3,0.05\n",
    "EPOCHS=6\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",
    "_ENC={k[8:]:v for k,v in _CK.items() if k.startswith(\"encoder.\") and not k.startswith(\"encoder_\")}\n",
    "\n",
    "def build_encoder(patch):\n",
    "    e=VisionTransformer(input_size=(T_BINS,N_MELS),patch_size=tuple(patch),in_chans=1,\n",
    "                        embed_dim=EMBED_DIM,depth=DEPTH,num_heads=12,mlp_ratio=4.0,\n",
    "                        use_flash_attn=False)\n",
    "    sd=dict(_ENC); 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=tuple(patch),mode=\"bicubic\",\n",
    "                                                    align_corners=False)\n",
    "    sd.pop(\"pos_embed\",None)\n",
    "    miss,unexp=e.load_state_dict(sd,strict=False)\n",
    "    assert unexp==[] and miss==[\"pos_embed\"], (patch,miss,unexp[:5])\n",
    "    return e\n",
    "\n",
    "for _p in (PATCH_BASE,PATCH_NEW):\n",
    "    _e=build_encoder(_p)\n",
    "    print(f\"patch {_p}: lưới {_e.patch_embed.num_patches_h}x{_e.patch_embed.num_patches_w}\"\n",
    "          f\" = {_e.patch_embed.num_patches} token | {_p[0]*HOP_MS:.0f} ms/token\"\n",
    "          f\" | {N_MELS//_p[1]} bin tần số\",flush=True)\n",
    "    del _e\n",
    "gc.collect()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 2 · Mel ASVspoof train (giống hệt v36/v38, 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 · Fine-tune ba seed với patch (4,32)\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):\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):\n",
    "        o=_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)))))\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)} | A05 CHỈ để theo dõi\",flush=True)\n",
    "\n",
    "def train_ft(patch,seed,tag):\n",
    "    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)\n",
    "    mdl=Net(build_encoder(patch)).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_ENC},\n",
    "                          {\"params\":hp,\"lr\":LR_HEAD}],weight_decay=WD)\n",
    "    sc=torch.optim.lr_scheduler.OneCycleLR(op,max_lr=[LR_ENC,LR_HEAD],\n",
    "                                           total_steps=EPOCHS*len(dl_tr),pct_start=0.1)\n",
    "    sca=torch.amp.GradScaler(\"cuda\")\n",
    "    ntok=mdl.enc.patch_embed.num_patches\n",
    "    print(f\"\\n=== {tag} | patch {patch} | {ntok} token | seed {seed} ===\",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\"{time.perf_counter()-t0:.0f}s\",flush=True)\n",
    "    torch.save(mdl.state_dict(),f\"{DRIVE}/{tag}.pth\"); print(f\"  [lưu] {tag}.pth -> Drive\")\n",
    "    del mdl; gc.collect(); torch.cuda.empty_cache()\n",
    "\n",
    "for sd_,tag in [(1234,\"p432_s1\"),(2,\"p432_s2\"),(3,\"p432_s3\")]:\n",
    "    if os.path.exists(f\"{DRIVE}/{tag}.pth\"): print(f\"[bỏ qua] {tag}\"); continue\n",
    "    train_ft(PATCH_NEW,sd_,tag)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 4 · File model cho arena (tự nhận patch từ checkpoint)\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, HOP_MS, FRAME_MS = 2.56, 10.0, 25.0\n",
    "F_MIN, F_MAX = 20, 8000\n",
    "EMBED_DIM, DEPTH = 768, 12\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, patch, 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=tuple(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",
    "    sd=torch.load(model_path,map_location=\"cpu\"); sd=sd.get(\"state_dict\",sd)\n",
    "    patch=tuple(sd[\"enc.patch_embed.proj.weight\"].shape[-2:])\n",
    "    print(f\"[nhan dien] patch={patch}\")\n",
    "    m=Net(patch,out_score_file_name)\n",
    "    miss,unexp=m.load_state_dict(sd,strict=False)\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",
    "WANT=[\"p432_s1\",\"p432_s2\",\"p432_s3\",\"audio_jepa_s1\",\"audio_jepa_s2\",\"audio_jepa_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",
    "miss=[t for t in WANT if t not in TAGS]\n",
    "if miss: print(\"[thiếu]\",miss)\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 DFADD và LibriSeVoc, lưu Drive sau MỖI hệ\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%4000==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",
    "def harvest():\n",
    "    r={}\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_,dd in v.items(): r[f\"{ds_}|{k}\"]=dd[\"EER (%)\"]\n",
    "    return r\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",
    "EXPECT={\"DFADD\":3755,\"librisevoc\":18487}\n",
    "for name in [\"DFADD\",\"librisevoc\"]:\n",
    "    build_hf(name)\n",
    "    d=pd.read_csv(f\"{PROTO}/{name}.csv\")\n",
    "    assert len(d)>=0.95*EXPECT[name], f\"{name} chỉ có {len(d)} dòng, cần ~{EXPECT[name]}\"\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/v40_{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",
    "        cur=harvest()\n",
    "        got=cur.get(f\"{name}|{tag}\")\n",
    "        print(f\"   {name:12s} {tag:16s} mã {r.returncode} | \"\n",
    "              f\"EER {'?' if got is None else round(got,2)} | {time.perf_counter()-t0:.0f}s\",flush=True)\n",
    "        json.dump(cur,open(f\"{DRIVE}/res_v40_partial.json\",\"w\"))   # LUU SAU MOI HE\n",
    "        if r.returncode!=0:\n",
    "            print(open(f\"/content/v40_{name}_{tag}.log\").read()[-2000:]); break\n",
    "        gc.collect()\n",
    "print(\"\\n[lưu] res_v40_partial.json -> Drive sau mỗi hệ\",flush=True)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "cellView": "form"
   },
   "outputs": [],
   "source": [
    "#@title 6 · Bảng: 80 ms so với 40 ms mỗi token\n",
    "res=harvest()\n",
    "GATE={\"DFADD\":(39.05,41.86),\"librisevoc\":(37.95,37.65)}\n",
    "BASE_REF={\"DFADD\":(10.75,3.48),\"librisevoc\":(41.96,2.93)}\n",
    "for name in [\"DFADD\",\"librisevoc\"]:\n",
    "    g=res.get(f\"{name}|aasist\")\n",
    "    lo,pub=GATE[name]\n",
    "    ok=g is not None and abs(g-lo)<1.5\n",
    "    print(\"=\"*64); print(f\"  {name}\"); print(\"=\"*64)\n",
    "    print(f\"   aasist (cổng)   {('?' if g is None else round(g,2))}   [v36/v38 đo {lo} · công bố {pub}]\"\n",
    "          f\"   -> {'ĐẠT' if ok else 'LỆCH'}\")\n",
    "    if not ok: print(\"   cổng lệch, không đọc số dưới\\n\"); continue\n",
    "    for label,pref,ref in [(\"patch (8,32) 80ms\",\"audio_jepa_s\",BASE_REF[name]),\n",
    "                           (\"patch (4,32) 40ms\",\"p432_s\",None)]:\n",
    "        vals=[v for k,v in res.items() if k.startswith(f\"{name}|{pref}\")]\n",
    "        if not vals: print(f\"   {label:20s} (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",
    "        seeds=\" \".join(f\"{v:.2f}\" for v in vals)\n",
    "        extra=f\"   [lần trước {ref[0]:.2f} ± {ref[1]:.2f}]\" if ref else \"\"\n",
    "        print(f\"   {label:20s} {m:7.2f} ± {sd:4.2f}  n={len(vals)}  [{seeds}]{extra}\")\n",
    "    a=[v for k,v in res.items() if k.startswith(f\"{name}|audio_jepa_s\")]\n",
    "    b=[v for k,v in res.items() if k.startswith(f\"{name}|p432_s\")]\n",
    "    if a and b:\n",
    "        d=float(np.mean(a))-float(np.mean(b))\n",
    "        print(f\"\\n   40ms so với 80ms: {d:+.2f} điểm  ({'TỐT HƠN' if d>0 else 'tệ hơn'})\")\n",
    "    print()\n",
    "np.savez(f\"{DRIVE}/res_v40_patch.npz\",**res); print(\"[lưu] res_v40_patch.npz -> Drive\")\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Đọc kết quả\n",
    "\n",
    "**Cổng trước.** `aasist` phải ra ~39,05 trên DFADD và ~37,95 trên LibriSeVoc — số v36/v38 đã đo bằng chính bản audio HuggingFace này, không phải số công bố. Lệch quá 1,5 thì ô 6 tự chặn.\n",
    "\n",
    "**Rồi so hai dòng.**\n",
    "\n",
    "| kết quả | nghĩa là gì |\n",
    "|---|---|\n",
    "| **LibriSeVoc tốt lên rõ, DFADD gần như không đổi** | giả thuyết đứng. Dấu vết vocoder ở thang mili giây, mảnh 80 ms bỏ lỡ; dấu vết TTS khuếch tán ở kết cấu phổ, không phụ thuộc độ mịn thời gian. Có cơ chế, có bằng chứng, có hướng khắc phục — đủ cho một mục Discussion mạnh |\n",
    "| cả hai cùng tốt lên | chỉ là \"nhiều token hơn thì tốt hơn\", không phải cơ chế. Vẫn dùng được như cải tiến kỹ thuật nhưng không viết thành lời giải thích |\n",
    "| LibriSeVoc không đổi | giả thuyết **sai**, bỏ trước khi in. Quay lại khung cũ và viết bài với năm tập hiện có |\n",
    "| cả hai tệ đi | 256 token làm loãng chú ý, hoặc nội suy trọng số patch xuống (4,32) làm hỏng biểu diễn tiền huấn luyện |\n",
    "\n",
    "**Vòng tiếp theo nếu vòng này thắng:** back-end đồ thị AASIST từng thất bại vì lưới thời gian chỉ 32 bước nên đồ thị gần rỗng. Với 64 bước thì đáng thử lại — và đó đúng là hướng \"JEPA + AASIST\" mà mục 4.1 của nhánh tối ưu đã loại.\n"
   ]
  }
 ]
}