{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# v35 — Quét số epoch (không cần gõ phím, chỉ Run all)\n",
    "\n",
    "Tải lên Colab, `Runtime → Change runtime type → T4 GPU`, rồi `Runtime → Run all`. Không phải sửa gì.\n",
    "\n",
    "## Vì sao\n",
    "\n",
    "Chẩn đoán trên LibriSeVoc cho thấy mô hình **bão hoà**, không phải lật dấu:\n",
    "\n",
    "| hệ | tb điểm thật | tb điểm giả | khoảng cách |\n",
    "|---|---|---|---|\n",
    "| aasist | −1,095 | −1,905 | 0,81 |\n",
    "| audio_jepa_s1 | −14,458 | −14,695 | **0,24** |\n",
    "\n",
    "Điểm dồn quanh −14, tức `softmax` cho xác suất bonafide ~1e-6 cho **mọi** mẫu. Dấu hiệu quá khớp: ta train 6 epoch tới loss 0,002.\n",
    "\n",
    "Đây là thí nghiệm duy nhất nhắm thẳng vào nguyên nhân đó: **giảm số epoch**, giữ nguyên mọi thứ khác (mel cũ 2,56 s / hop 10 ms / fmax 8000, patch (8,32), lr 3e-5, seed 1234).\n",
    "\n",
    "## Mốc so sánh (6 epoch, đã có)\n",
    "\n",
    "| tập | JEPA 6 epoch | aasist |\n",
    "|---|---|---|\n",
    "| DFADD | 10,75 ± 3,48 | 39,05 |\n",
    "| LibriSeVoc | 41,96 ± 2,93 | 37,95 |\n",
    "| ASVspoof19 | 7,39 ± 0,69 | 0,83 |\n",
    "\n",
    "Nếu quá khớp đúng là thủ phạm: **LibriSeVoc phải khá lên rõ**, DFADD không mất nhiều.\n",
    "\n",
    "Thời gian: mel ~6 phút, train 2+3+4 epoch ~25 phút, tải hai tập ~10 phút, chấm ~25 phút. Khoảng 1 giờ 10."
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 1 · Môi trường và 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\",\"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 CŨ, đúng bản đã cho 7,39 trên ASVspoof19 ---\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_ENC,LR_HEAD,WD=32,3e-5,1e-3,0.05\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",
    "def build_encoder():\n",
    "    e=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",
    "    sd=dict(_ENC); w=sd[\"patch_embed.proj.weight\"]\n",
    "    if tuple(w.shape[-2:])!=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=e.load_state_dict(sd,strict=False)\n",
    "    assert unexp==[] and miss==[\"pos_embed\"], (miss,unexp[:5])\n",
    "    return e\n",
    "_e=build_encoder()\n",
    "print(f\"✓ lưới {_e.patch_embed.num_patches_h}x{_e.patch_embed.num_patches_w}\"\n",
    "      f\" = {_e.patch_embed.num_patches} token | mel {CLIP_S}s hop {HOP_MS}ms fmax {F_MAX}\",flush=True)"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 2 · Mel ASVspoof train (cấu hình cũ, 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\n",
    "print(\"bonafide\",int(Y_TR.sum()),\"| spoof\",int((1-Y_TR).sum()),flush=True)"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 3 · Train 2 / 3 / 4 epoch\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)}\",flush=True)\n",
    "\n",
    "def train_one(tag,epochs):\n",
    "    torch.manual_seed(1234); np.random.seed(1234); random.seed(1234)\n",
    "    mdl=Net(build_encoder()).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",
    "    print(f\"\\n=== {tag} | {epochs} epoch ===\",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)\n",
    "        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",
    "    torch.save(mdl.state_dict(),f\"{DRIVE}/{tag}.pth\"); print(f\"  [lưu] {tag}.pth\")\n",
    "    del mdl; gc.collect(); torch.cuda.empty_cache()\n",
    "\n",
    "for tag,ep in [(\"jep_ep2\",2),(\"jep_ep3\",3),(\"jep_ep4\",4)]:\n",
    "    if os.path.exists(f\"{DRIVE}/{tag}.pth\"): print(f\"[bỏ qua] {tag}\"); continue\n",
    "    train_one(tag,ep)\n",
    "print(\"\\nGhi chú: cột |điểm| tb cho biết mức bão hoà. 6 epoch cho ~14 trên dữ liệu ngoài miền.\")"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 4 · File model\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   = 10.0                     # Audio-JEPA tien huan luyen tren clip 10 giay\n",
    "HOP_MS   = 10000.0/256              # 39.0625 ms  (clip_length*1000 / target_time_bins)\n",
    "FRAME_MS = 2.5*HOP_MS               # 97.65625 ms (2.5 * hop)\n",
    "F_MIN, F_MAX = 20, 16000            # f_max = sr//2, KHONG phai 8000\n",
    "EMBED_DIM, DEPTH, PATCH = 768, 12, (16,16)\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",
    "        \"\"\"DUNG y het MelSpecTransform cua repo Audio-JEPA.\"\"\"\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.replace(\"CLIP_S   = 10.0\",\"CLIP_S   = 2.56\")\n",
    "      .replace(\"HOP_MS   = 10000.0/256\",\"HOP_MS   = 10.0\")\n",
    "      .replace(\"FRAME_MS = 2.5*HOP_MS\",\"FRAME_MS = 25.0\")\n",
    "      .replace(\"F_MIN, F_MAX = 20, 16000\",\"F_MIN, F_MAX = 20, 8000\")\n",
    "      .replace(\"PATCH = 768, 12, (16,16)\",\"PATCH = 768, 12, (8,32)\"))\n",
    "ast.parse(CODE)\n",
    "for k in [\"CLIP_S   = 2.56\",\"F_MIN, F_MAX = 20, 8000\",\"(8,32)\"]:\n",
    "    assert k in CODE, k\n",
    "TAGS=[t for t in [\"jep_ep2\",\"jep_ep3\",\"jep_ep4\",\"audio_jepa_s1\"]\n",
    "      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",
    "print(\"sẽ chấm:\",[\"aasist\"]+TAGS,flush=True)"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 5 · Tải hai tập và chấm (DFADD trước, LibriSeVoc sau)\n",
    "from datasets import load_dataset as _ld\n",
    "def build_ds(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 P\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%5000==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",
    "    return P\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",
    "for name in [\"DFADD\",\"LibriSeVoc\"]:\n",
    "    build_ds(name)\n",
    "    d=pd.read_csv(f\"{PROTO}/{name}.csv\"); print(f\"{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/v35_{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:12s} {tag:16s} mã {r.returncode} | {time.perf_counter()-t0:.0f}s\",flush=True)\n",
    "        if r.returncode!=0: print(open(f\"/content/v35_{name}_{tag}.log\").read()[-2000:]); break\n",
    "        gc.collect()"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": [
    "#@title 6 · Bảng: số epoch ảnh hưởng thế nào\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",
    "REF={\"DFADD\":(10.75,39.05),\"LibriSeVoc\":(41.96,37.95)}   # (JEPA 6 epoch, aasist) đã có\n",
    "ORDER=[\"jep_ep2\",\"jep_ep3\",\"jep_ep4\",\"audio_jepa_s1\",\"aasist\"]\n",
    "NAME={\"jep_ep2\":\"JEPA 2 epoch\",\"jep_ep3\":\"JEPA 3 epoch\",\"jep_ep4\":\"JEPA 4 epoch\",\n",
    "      \"audio_jepa_s1\":\"JEPA 6 epoch (cũ)\",\"aasist\":\"aasist (mốc)\"}\n",
    "for ds_ in [\"DFADD\",\"LibriSeVoc\"]:\n",
    "    r=res.get(ds_,{})\n",
    "    if not r: continue\n",
    "    print(\"=\"*56); print(f\"  {ds_}\"); print(\"=\"*56)\n",
    "    for k in ORDER:\n",
    "        if k in r: print(f\"   {NAME[k]:22s} {r[k]:7.2f}\")\n",
    "    g=r.get(\"aasist\"); ref6,refa=REF[ds_]\n",
    "    print(f\"   (tham chiếu: 6 epoch 3 seed = {ref6:.2f}, aasist công bố lần trước = {refa:.2f})\")\n",
    "    best=min([(v,k) for k,v in r.items() if k!=\"aasist\"],default=None)\n",
    "    if best and g:\n",
    "        print(f\"\\n   tốt nhất: {NAME.get(best[1],best[1])} = {best[0]:.2f}\")\n",
    "        print(f\"   so aasist {g:.2f}: {g-best[0]:+.2f}  \"\n",
    "              +(\"-> THẮNG\" if best[0]<g else \"-> thua\"))\n",
    "        print(f\"   so 6 epoch {ref6:.2f}: {ref6-best[0]:+.2f}  \"\n",
    "              +(\"-> ÍT EPOCH TỐT HƠN\" if best[0]<ref6 else \"-> 6 epoch vẫn tốt hơn\"))\n",
    "    print()\n",
    "np.savez(f\"{DRIVE}/res_v35_epoch.npz\",\n",
    "         **{f\"{d}|{k}\":v for d,dd in res.items() for k,v in dd.items()})\n",
    "print(\"[lưu] res_v35_epoch.npz -> Drive\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Đọc kết quả\n",
    "\n",
    "Một câu hỏi: **giảm epoch có cứu được LibriSeVoc không.**\n",
    "\n",
    "| LibriSeVoc ra | nghĩa là |\n",
    "|---|---|\n",
    "| < 37,95 | quá khớp đúng là thủ phạm, và ta **thắng aasist ở tập thứ ba** |\n",
    "| 38 – 42 | có cải thiện nhưng chưa đủ, cần thêm tăng cường dữ liệu |\n",
    "| ≈ 42 hoặc tệ hơn | quá khớp **không** phải nguyên nhân chính, dừng hướng này |\n",
    "\n",
    "Cột `|điểm| tb` in trong lúc train là chỉ số phụ đáng nhìn: 6 epoch cho giá trị tuyệt đối khoảng 14 trên dữ liệu ngoài miền. Nếu 2–3 epoch giữ nó ở mức 3–5 thì mô hình còn giữ được độ bất định, và đó là dấu hiệu tốt kể cả trước khi xem EER.\n",
    "\n",
    "Trên DFADD, mất một chút so với 10,75 là chấp nhận được — miễn LibriSeVoc khá lên nhiều hơn phần mất đó."
   ]
  }
 ],
 "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
}