{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# v41 — Preflight nhánh tiền huấn luyện tiếng nói\n",
    "\n",
    "**Chạy hết notebook này trước khi cam kết 33 giờ GPU.** Bảy cổng, tổng cộng dưới 20 phút trên T4.\n",
    "Không cổng nào huấn luyện gì cả — chúng chỉ kiểm rằng đường ống đúng, số liệu đọc được, và\n",
    "ước lượng thời gian là thật chứ không phải phỏng đoán.\n",
    "\n",
    "| cổng | kiểm gì | hỏng thì mất gì |\n",
    "|---|---|---|\n",
    "| G1 | hình học lưới, bảng nạp trọng số, `pos_embed` sinh lại | mã hoá vị trí sai lặng lẽ → cả lượt chạy vô nghĩa |\n",
    "| G2 | learning rate và weight decay **thật** đang chạy | lr 1e-3 xoá biểu diễn AudioSet trong epoch đầu |\n",
    "| G3 | mel đúng shape, hai hàng đệm, thống kê mặt nạ | đầu vào lệch khỏi nhánh phát hiện |\n",
    "| G4 | mốc \"khoẻ mạnh\" của bốn chỉ số sụp đổ | không có mốc thì mọi số về sau không đọc được |\n",
    "| G5 | một bước tiền huấn luyện thật + đo thông lượng | ước lượng 9,4 h/lượt có thể sai gấp đôi |\n",
    "| G6 | lưu và khôi phục checkpoint qua Drive | Colab đã đứt ba đêm; mất mười tiếng |\n",
    "| G7 | encoder của đường ống mới trùng **từng bit** với v38/v40 | nhánh phát hiện bị hỏng mà không ai biết |\n",
    "\n",
    "Tài liệu quyết định: `tien-huan-luyen-tieng-noi/01-ba-quyet-dinh-thiet-ke.md`.\n",
    "Bối cảnh nhánh: `tien-huan-luyen-tieng-noi/00-DOC-TRUOC-tien-huan-luyen.md`.\n",
    "\n",
    "Điều kiện đi tiếp nằm ở ô markdown cuối. **Bảy cổng phải xanh hết.**\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 1 · Môi trường và ba nhánh trọng số\n",
    "import os, sys, glob, json, time, math, random, shutil, io, types, subprocess, gc\n",
    "subprocess.run([sys.executable,\"-m\",\"pip\",\"install\",\"-q\",\"rich\",\"librosa\",\"soundfile\",\n",
    "  \"pandas\",\"scikit-learn\",\"datasets>=2.19,<4.0\",\"huggingface_hub\"],check=False)\n",
    "import numpy as np, 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/pre41\"; os.makedirs(BASE,exist_ok=True); os.chdir(BASE)\n",
    "REPO=\"/content/audio-jepa\"\n",
    "if not os.path.isdir(REPO):\n",
    "    subprocess.run([\"git\",\"clone\",\"--depth\",\"1\",\"-q\",\n",
    "                    \"https://github.com/LudovicTuncay/Audio-JEPA.git\",REPO],check=False)\n",
    "\n",
    "# --- shim flash_attn + namespace src, y nguyên v38/v40 ---\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\",\"src.optimizers\"):\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, VisionTransformerPredictor\n",
    "from src.models.components.positional_embedding import get_2d_sincos_pos_embed\n",
    "from src.masks.components.utils import apply_masks\n",
    "from src.models.components.loss import Loss\n",
    "from src.optimizers.warmup_cosine import WarmupCosineScheduler\n",
    "from src.optimizers.cosine_wd import CosineWDScheduler\n",
    "from huggingface_hub import hf_hub_download\n",
    "\n",
    "# --- hằng số: QUY CÁCH PHÁT HIỆN, quyết định 5.1 phương án (a) ---\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,HEADS  = 768,12,12\n",
    "PATCH                  = (8,32)          # (thời gian, tần số) -> lưới 32 x 4 = 128 token\n",
    "PRED_DIM,PRED_DEPTH    = 384,6\n",
    "# --- siêu tham số tiền huấn luyện tiếp, đã chốt (xem bẫy 3 trong tài liệu) ---\n",
    "PRE_REF_LR   = 1e-4      # KHÔNG dùng 1e-3 mặc định của repo\n",
    "PRE_START_LR = 1e-6\n",
    "PRE_FINAL_LR = 0.0\n",
    "PRE_WARMUP   = 500\n",
    "PRE_WD       = 1e-6\n",
    "PRE_BATCH    = 32\n",
    "EMA_TAU0     = 0.999     # thay cho 0.996 mặc định\n",
    "MASK_RATIO   = (0.4,0.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",
    "def _branch(p):\n",
    "    return {k[len(p):]:v for k,v in _CK.items() if k.startswith(p)}\n",
    "_ENC={k[8:]:v for k,v in _CK.items() if k.startswith(\"encoder.\") and not k.startswith(\"encoder_\")}\n",
    "_TGT=_branch(\"target_encoder.\")\n",
    "_PRD=_branch(\"predictor.\")\n",
    "print(f\"checkpoint: encoder {len(_ENC)} tensor | target_encoder {len(_TGT)} | predictor {len(_PRD)}\",flush=True)\n",
    "assert len(_ENC)==149 and len(_TGT)==149 and len(_PRD)==80, (len(_ENC),len(_TGT),len(_PRD))\n",
    "\n",
    "GRID_H,GRID_W = T_BINS//PATCH[0], N_MELS//PATCH[1]\n",
    "N_TOK = GRID_H*GRID_W\n",
    "\n",
    "def build_encoder(patch=PATCH, src=None):\n",
    "    \"\"\"Giống hệt build_encoder của v38/v40. pos_embed BỊ BỎ -> sinh lại theo lưới mới.\"\"\"\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=HEADS,mlp_ratio=4.0,\n",
    "                        use_flash_attn=False)\n",
    "    sd=dict(_ENC if src is None else src); 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)                      # <-- bắt buộc, xem bẫy 1\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",
    "def build_predictor():\n",
    "    p=VisionTransformerPredictor(num_patches_h=GRID_H,num_patches_w=GRID_W,\n",
    "                                 embed_dim=EMBED_DIM,predictor_embed_dim=PRED_DIM,\n",
    "                                 depth=PRED_DEPTH,num_heads=HEADS,mlp_ratio=4.0,\n",
    "                                 use_flash_attn=False)\n",
    "    sd=dict(_PRD); sd.pop(\"predictor_pos_embed\",None)   # <-- bắt buộc, xem bẫy 1\n",
    "    miss,unexp=p.load_state_dict(sd,strict=False)\n",
    "    assert unexp==[] and miss==[\"predictor_pos_embed\"], (miss,unexp[:5])\n",
    "    return p\n",
    "\n",
    "print(f\"lưới {GRID_H} x {GRID_W} = {N_TOK} token | {PATCH[0]*HOP_MS:.0f} ms/token\"\n",
    "      f\" | {N_MELS//PATCH[1]} bin tần số/token\",flush=True)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 2 · G1 — hình học, bảng nạp trọng số, pos_embed sinh lại\n",
    "ok=[]\n",
    "enc = build_encoder()\n",
    "tgt = build_encoder(src=_TGT)     # nhánh target nạp trọng số target_encoder của checkpoint\n",
    "prd = build_predictor()\n",
    "\n",
    "# --- G1a hình học ---\n",
    "assert (enc.patch_embed.num_patches_h,enc.patch_embed.num_patches_w)==(GRID_H,GRID_W)\n",
    "assert enc.patch_embed.num_patches==N_TOK==128\n",
    "ok.append(f\"G1a lưới {GRID_H}x{GRID_W}={N_TOK} token\")\n",
    "\n",
    "# --- G1b pos_embed phải KHÁC tensor trong checkpoint (nếu bằng => đã nạp nhầm) ---\n",
    "pe_ck = _ENC[\"pos_embed\"]                       # lưới 16 x 8 của bản gốc\n",
    "pe_new = enc.pos_embed.detach()\n",
    "assert pe_ck.shape==pe_new.shape, (pe_ck.shape,pe_new.shape)   # CÙNG shape - đây là cái bẫy\n",
    "d = (pe_ck-pe_new).abs().max().item()\n",
    "cs = F.cosine_similarity(pe_ck[0],pe_new[0],dim=-1).mean().item()\n",
    "assert d > 1e-3, \"pos_embed TRÙNG checkpoint -> đã nạp nhầm lưới 16x8!\"\n",
    "ok.append(f\"G1b pos_embed sinh lại: cùng shape {tuple(pe_new.shape)}, max|hiệu|={d:.3f}, cos={cs:.3f}\")\n",
    "\n",
    "pp_ck = _PRD[\"predictor_pos_embed\"]; pp_new = prd.predictor_pos_embed.detach()\n",
    "assert pp_ck.shape==pp_new.shape and (pp_ck-pp_new).abs().max().item() > 1e-3\n",
    "ok.append(f\"G1c predictor_pos_embed sinh lại: max|hiệu|={(pp_ck-pp_new).abs().max().item():.3f}\")\n",
    "\n",
    "# --- G1d chứng minh cái bẫy: nạp pos_embed sai lưới KHÔNG hề báo lỗi ---\n",
    "_bad = VisionTransformer(input_size=(T_BINS,N_MELS),patch_size=PATCH,in_chans=1,\n",
    "                         embed_dim=EMBED_DIM,depth=DEPTH,num_heads=HEADS,mlp_ratio=4.0,\n",
    "                         use_flash_attn=False)\n",
    "_sd=dict(_ENC); _w=_sd[\"patch_embed.proj.weight\"]\n",
    "_sd[\"patch_embed.proj.weight\"]=F.interpolate(_w.float(),size=PATCH,mode=\"bicubic\",align_corners=False)\n",
    "_m,_u=_bad.load_state_dict(_sd,strict=False)          # GIỮ nguyên pos_embed sai lưới\n",
    "print(f\"  [bẫy 1] nạp pos_embed lưới 16x8 vào mô hình lưới 32x4 -> missing={_m}, unexpected={_u}\")\n",
    "print(f\"  [bẫy 1] không lỗi, không cảnh báo. interpolate_pos_encoding trả về nguyên xi vì \"\n",
    "      f\"{GRID_H}*{GRID_W} == {pe_ck.shape[1]}\")\n",
    "del _bad; gc.collect()\n",
    "ok.append(\"G1d đã tái lập cái bẫy im lặng (chỉ để chứng minh, không dùng)\")\n",
    "\n",
    "# --- G1e KHÔNG HỒI QUY: encoder phải trùng từng bit với cách dựng của v38/v40 ---\n",
    "ref = build_encoder()\n",
    "a,b = enc.state_dict(), ref.state_dict()\n",
    "assert a.keys()==b.keys()\n",
    "bad=[k for k in a if not torch.equal(a[k],b[k])]\n",
    "assert not bad, bad[:5]\n",
    "ok.append(f\"G1e encoder trùng từng bit với v38/v40 trên cả {len(a)} tensor\")\n",
    "del ref; gc.collect()\n",
    "\n",
    "print(\"\\n\".join(\"  ✓ \"+s for s in ok))\n",
    "print(\"\\nG1 XANH\" if len(ok)==5 else \"\\nG1 ĐỎ\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 3 · G2 — learning rate và weight decay THẬT đang chạy\n",
    "STEPS_TOTAL = 2_812_500//PRE_BATCH        # 2,81 M mẫu / batch  (xem tài liệu mục 2.5)\n",
    "\n",
    "def peak_lr(ref_lr, warmup, total, probe=None):\n",
    "    p=[nn.Parameter(torch.zeros(2))]\n",
    "    o=torch.optim.AdamW(p,lr=3e-4,weight_decay=0.05,betas=(0.9,0.95))\n",
    "    s=WarmupCosineScheduler(o,warmup_steps=warmup,start_lr=PRE_START_LR,ref_lr=ref_lr,\n",
    "                            T_max=total,final_lr=PRE_FINAL_LR)\n",
    "    lrs=[]\n",
    "    for _ in range(min(total,8000)):\n",
    "        lrs.append(o.param_groups[0]['lr']); o.step(); s.step()\n",
    "    return max(lrs), lrs\n",
    "\n",
    "pk_repo,_   = peak_lr(1e-3, 1000, STEPS_TOTAL)          # mặc định của repo\n",
    "pk_chon,lrs = peak_lr(PRE_REF_LR, PRE_WARMUP, STEPS_TOTAL)\n",
    "print(f\"  tổng bước cho 2,81 M mẫu, batch {PRE_BATCH}: {STEPS_TOTAL:,}\")\n",
    "print(f\"  lr khai trong configs/model/jepa.yaml (AdamW) : 3.000e-04\")\n",
    "print(f\"  ĐỈNH THẬT với ref_lr mặc định của repo        : {pk_repo:.3e}   <-- gấp {pk_repo/3e-4:.2f} lần\")\n",
    "print(f\"  ĐỈNH THẬT với ref_lr ta chọn                  : {pk_chon:.3e}\")\n",
    "for s in (0,1,PRE_WARMUP//2,PRE_WARMUP-1,PRE_WARMUP,PRE_WARMUP+1,2000,5000):\n",
    "    if s < len(lrs): print(f\"      bước {s:5d}  lr = {lrs[s]:.3e}\")\n",
    "assert abs(pk_chon-PRE_REF_LR)/PRE_REF_LR < 1e-3, pk_chon\n",
    "assert pk_repo > 3*PRE_REF_LR\n",
    "\n",
    "# weight decay\n",
    "p2=[nn.Parameter(torch.zeros(2))]\n",
    "o2=torch.optim.AdamW(p2,lr=PRE_REF_LR,weight_decay=0.05)\n",
    "w2=CosineWDScheduler(o2,ref_wd=PRE_WD,T_max=STEPS_TOTAL,final_wd=PRE_WD)\n",
    "for _ in range(50): o2.step(); w2.step()\n",
    "print(f\"  wd khai trong config: 0.05 | wd THẬT sau 50 bước: {o2.param_groups[0]['weight_decay']}\")\n",
    "assert abs(o2.param_groups[0]['weight_decay']-PRE_WD) < 1e-12\n",
    "print(\"\\nG2 XANH\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 4 · G3 — mel đúng quy cách, và thống kê mặt nạ\n",
    "_RS16 = torchaudio.transforms.Resample(16_000,SR_MODEL).to(DEVICE)\n",
    "def mel_batch(w16):\n",
    "    \"\"\"w16: (B, L) waveform 16 kHz trên DEVICE -> (B,1,T_BINS,N_MELS). Y hệt mel_old của v38/v40.\"\"\"\n",
    "    w=_RS16(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",
    "\n",
    "_w = (0.1*torch.randn(4,int(16_000*CLIP_S))).to(DEVICE)\n",
    "_m = mel_batch(_w)\n",
    "assert tuple(_m.shape)==(4,1,T_BINS,N_MELS), _m.shape\n",
    "_zero = int((_m[0,0].abs().sum(1)==0).sum())\n",
    "print(f\"  mel shape {tuple(_m.shape)} | kaldi trả 254 khung, đệm lên {T_BINS}\"\n",
    "      f\" -> {_zero} hàng toàn 0 ở cuối\")\n",
    "assert _zero==2, _zero\n",
    "print(f\"  hai hàng đệm nằm trong {GRID_W} token cuối của lưới {GRID_H}x{GRID_W}, chiếm 2/8 số hàng của chúng\")\n",
    "\n",
    "# --- mặt nạ: bản sao nguyên văn logic của src/masks/components/random_block.py ---\n",
    "def make_masks(B, ratio_range=MASK_RATIO, seed=None):\n",
    "    g=torch.Generator()\n",
    "    if seed is not None: g.manual_seed(seed)\n",
    "    r=ratio_range[0]+torch.rand(1,generator=g).item()*(ratio_range[1]-ratio_range[0])\n",
    "    keep=int(N_TOK*(1.-r))\n",
    "    perms=[torch.randperm(N_TOK,generator=g) for _ in range(B)]\n",
    "    ctx=[torch.stack([p[:keep] for p in perms])]\n",
    "    tgt=[torch.stack([p[keep:] for p in perms])]\n",
    "    return ctx,tgt,r,keep\n",
    "\n",
    "_c,_t,_r,_keep = make_masks(8,seed=0)\n",
    "print(f\"  tỉ lệ che {_r:.3f} -> giữ {_keep}/{N_TOK} token ngữ cảnh, dự đoán {N_TOK-_keep}\")\n",
    "assert _c[0].shape==(8,_keep) and _t[0].shape==(8,N_TOK-_keep)\n",
    "\n",
    "def neighbour_visibility(trials=300):\n",
    "    a=[];f=[]\n",
    "    for s in range(trials):\n",
    "        c,t,_,_=make_masks(1,seed=1000+s)\n",
    "        cs=set(c[0][0].tolist())\n",
    "        for tok in t[0][0].tolist():\n",
    "            h,w=divmod(tok,GRID_W); nb=[]\n",
    "            for dh,dw in ((1,0),(-1,0),(0,1),(0,-1)):\n",
    "                H2,W2=h+dh,w+dw\n",
    "                if 0<=H2<GRID_H and 0<=W2<GRID_W: nb.append(H2*GRID_W+W2)\n",
    "            v=sum(1 for x in nb if x in cs); a.append(v>0); f.append(v/len(nb))\n",
    "    return float(np.mean(a)),float(np.mean(f))\n",
    "_a,_f = neighbour_visibility()\n",
    "print(f\"  [bẫy 5] rắc ngẫu nhiên: {_a*100:.1f}% token đích còn ≥1 hàng xóm nhìn thấy được,\"\n",
    "      f\" trung bình {_f*100:.1f}% hàng xóm\")\n",
    "print(f\"           -> nhiệm vụ dự đoán dễ. Đây là đặc điểm của bản cài đặt, KHÔNG sửa ở nhánh này.\")\n",
    "print(\"\\nG3 XANH\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "%%writefile /content/collapse_metrics.py\n",
    "\"\"\"Chi so phat hien sup do bieu dien cho Audio-JEPA.\n",
    "\n",
    "Dung trong lop tien huan luyen tiep: goi moi epoch tren MOT batch co dinh,\n",
    "in ra va luu cung checkpoint. Xem `tien-huan-luyen-tieng-noi/01-ba-quyet-dinh-thiet-ke.md` muc 3.\n",
    "\n",
    "Diem cot loi: ham loss cua Audio-JEPA la `2 - 2*cos` tren target da chuan hoa\n",
    "tung token, nen loss GIAM khi bieu dien sup do. Khong duoc dung loss de canh.\n",
    "\n",
    "Bon chi so bo tro nhau; da kiem bang tensor tong hop (chay `python collapse_metrics.py`):\n",
    "\n",
    "    truong hop                              std_token  rank_eff  cos_cross  std_utt\n",
    "    khoe (ngau nhien day du hang)              0.0361     758.7     0.0003   0.0360\n",
    "    sup do hoan toan (moi token 1 huong)       0.0000     758.7     1.0000   0.0000\n",
    "    sup do mot phan (hang 8)                   0.0350       8.0    -0.0023   0.0350\n",
    "    moi utterance giong nhau, token khac       0.0243     159.9     0.5468   0.0001\n",
    "\n",
    "Doc bang nay ky truoc khi dat nguong:\n",
    "  - `rank_eff` (tren ma tran DA TRU TRUNG BINH) **khong** bat duoc sup do hoan toan\n",
    "    — no van 758.7 — vi sau khi tru trung binh chi con nhieu, va nhieu thi day hang.\n",
    "    Sup do hoan toan do `std_token` va `cos_cross` bat.\n",
    "  - Nguoc lai, sup do mot phan (hang thap) **chi** co `rank_eff` bat duoc; ba chi so\n",
    "    kia deu trong nhu binh thuong.\n",
    "  - `std_utt` la chi so duy nhat bat duoc \"moi utterance giong nhau\".\n",
    "  => Phai theo doi ca bon. Bo mot cai la mu mot che do hong.\n",
    "\"\"\"\n",
    "from __future__ import annotations\n",
    "import math\n",
    "import torch\n",
    "import torch.nn.functional as F\n",
    "\n",
    "__all__ = [\"collapse_metrics\", \"check_gate\", \"format_row\", \"HEADER\"]\n",
    "\n",
    "HEADER = f\"{'epoch':>5s} {'std_token':>10s} {'rank_eff':>9s} {'rank_raw':>9s} {'cos_cross':>10s} {'std_utt':>8s}\"\n",
    "\n",
    "\n",
    "def _effective_rank(X: torch.Tensor, eps: float = 1e-12) -> float:\n",
    "    \"\"\"exp(entropy cua pho gia tri ky di da chuan hoa). Bang 1 khi hang 1.\"\"\"\n",
    "    sv = torch.linalg.svdvals(X.double())\n",
    "    p = sv / (sv.sum() + eps)\n",
    "    p = p[p > 0]\n",
    "    return float(torch.exp(-(p * torch.log(p)).sum()))\n",
    "\n",
    "\n",
    "@torch.no_grad()\n",
    "def collapse_metrics(H: torch.Tensor, n_pairs: int = 4096, seed: int = 0) -> dict:\n",
    "    \"\"\"H: (B, N, D) dau ra encoder tren mot batch co dinh. B >= 2, N >= 2.\n",
    "\n",
    "    Tra ve dict 4 chi so + `rank_raw` (hang hieu dung khi KHONG tru trung binh,\n",
    "    bat duoc sup do ve mot huong chung).\n",
    "    \"\"\"\n",
    "    assert H.dim() == 3, f\"can (B,N,D), nhan duoc {tuple(H.shape)}\"\n",
    "    B, N, D = H.shape\n",
    "    assert B >= 2 and N >= 2, \"can it nhat 2 utterance va 2 token\"\n",
    "    H = H.detach().float().cpu()\n",
    "\n",
    "    Hn3 = F.normalize(H, dim=-1)                 # (B,N,D) da chuan hoa\n",
    "    Hn = Hn3.reshape(B * N, D)\n",
    "\n",
    "    std_token = float(Hn.std(dim=0).mean())\n",
    "\n",
    "    Hf = H.reshape(B * N, D)\n",
    "    rank_eff = _effective_rank(Hf - Hf.mean(0, keepdim=True))   # sup do chieu\n",
    "    rank_raw = _effective_rank(Hn)                              # sup do ve 1 huong\n",
    "\n",
    "    g = torch.Generator().manual_seed(seed)\n",
    "    b1 = torch.randint(0, B, (n_pairs,), generator=g)\n",
    "    b2 = (b1 + torch.randint(1, B, (n_pairs,), generator=g)) % B   # bao dam b2 != b1\n",
    "    n1 = torch.randint(0, N, (n_pairs,), generator=g)\n",
    "    n2 = torch.randint(0, N, (n_pairs,), generator=g)\n",
    "    cos_cross = float((Hn3[b1, n1] * Hn3[b2, n2]).sum(-1).mean())\n",
    "\n",
    "    std_utt = float(F.normalize(H.mean(1), dim=-1).std(dim=0).mean())\n",
    "\n",
    "    return dict(std_token=std_token, rank_eff=rank_eff, rank_raw=rank_raw,\n",
    "                cos_cross=cos_cross, std_utt=std_utt)\n",
    "\n",
    "\n",
    "def format_row(epoch, m: dict) -> str:\n",
    "    return (f\"{epoch:>5} {m['std_token']:10.4f} {m['rank_eff']:9.1f} {m['rank_raw']:9.1f} \"\n",
    "            f\"{m['cos_cross']:10.4f} {m['std_utt']:8.4f}\")\n",
    "\n",
    "\n",
    "def check_gate(m: dict, base: dict) -> list[str]:\n",
    "    \"\"\"Tra ve danh sach ly do phai DUNG. Rong = di tiep.\n",
    "\n",
    "    `base` la bo chi so do tren checkpoint goc (epoch 0). Quy tac dung chi\n",
    "    kich hoat khi vi pham o HAI checkpoint lien tiep — nguoi goi tu giu trang thai.\n",
    "    \"\"\"\n",
    "    bad = []\n",
    "    if m[\"rank_eff\"] < 0.50 * base[\"rank_eff\"]:\n",
    "        bad.append(f\"rank_eff {m['rank_eff']:.1f} < 50% moc {base['rank_eff']:.1f} (sup do chieu)\")\n",
    "    if m[\"rank_raw\"] < 0.50 * base[\"rank_raw\"]:\n",
    "        bad.append(f\"rank_raw {m['rank_raw']:.1f} < 50% moc {base['rank_raw']:.1f} (sup do ve 1 huong)\")\n",
    "    if m[\"cos_cross\"] > 0.9 or m[\"cos_cross\"] > base[\"cos_cross\"] + 0.3:\n",
    "        bad.append(f\"cos_cross {m['cos_cross']:.3f} (moc {base['cos_cross']:.3f})\")\n",
    "    if m[\"std_token\"] < 0.25 * base[\"std_token\"]:\n",
    "        bad.append(f\"std_token {m['std_token']:.4f} < 25% moc {base['std_token']:.4f}\")\n",
    "    if m[\"std_utt\"] < 0.25 * base[\"std_utt\"]:\n",
    "        bad.append(f\"std_utt {m['std_utt']:.4f} < 25% moc {base['std_utt']:.4f} (moi utterance giong nhau)\")\n",
    "    return bad\n",
    "\n",
    "\n",
    "if __name__ == \"__main__\":\n",
    "    torch.manual_seed(0)\n",
    "    B, N, D = 64, 128, 768\n",
    "    v = torch.randn(1, 1, D)\n",
    "    u = torch.randn(1, 1, D)\n",
    "    cases = {\n",
    "        \"khoe (ngau nhien day du hang)\":       torch.randn(B, N, D),\n",
    "        \"sup do hoan toan (moi token 1 huong)\": v.repeat(B, N, 1) + 0.001 * torch.randn(B, N, D),\n",
    "        \"sup do mot phan (hang 8)\":             torch.randn(B, N, 8) @ torch.randn(8, D),\n",
    "        \"moi utterance giong nhau, token khac\": u + 0.9 * torch.randn(1, N, D).repeat(B, 1, 1)\n",
    "                                                  + 0.02 * torch.randn(B, N, D),\n",
    "    }\n",
    "    base = collapse_metrics(cases[\"khoe (ngau nhien day du hang)\"])\n",
    "    print(f\"{'truong hop':40s} {'std_token':>10s} {'rank_eff':>9s} {'rank_raw':>9s} {'cos_cross':>10s} {'std_utt':>8s}  cong\")\n",
    "    for k, H in cases.items():\n",
    "        m = collapse_metrics(H)\n",
    "        bad = check_gate(m, base)\n",
    "        verdict = \"DI TIEP\" if not bad else \"DUNG: \" + \"; \".join(x.split(\" (\")[0] for x in bad)\n",
    "        print(f\"{k:40s} {m['std_token']:10.4f} {m['rank_eff']:9.1f} {m['rank_raw']:9.1f} \"\n",
    "              f\"{m['cos_cross']:10.4f} {m['std_utt']:8.4f}  {verdict}\")\n",
    "    print(f\"\\nmoc ly thuyet trai deu tren mat cau D={D}: std_token = 1/sqrt(D) = {1/math.sqrt(D):.4f}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 5b · Batch cố định từ kho tiền huấn luyện + mốc khoẻ mạnh\n",
    "sys.path.insert(0,\"/content\")\n",
    "import collapse_metrics as CM; import importlib; importlib.reload(CM)\n",
    "from datasets import load_dataset, Audio\n",
    "\n",
    "N_FIX = 64            # batch cố định để đo chỉ số, dùng lại y hệt ở MỌI epoch\n",
    "FIX_SEED = 20260909\n",
    "CORPUS = \"voxpopuli\"  # kho của nhánh N; đổi sang \"librispeech\" cho nhánh T\n",
    "\n",
    "def _take_clips(n=N_FIX):\n",
    "    L16=int(16_000*CLIP_S); out=[]\n",
    "    if CORPUS==\"voxpopuli\":\n",
    "        ds=load_dataset(\"facebook/voxpopuli\",\"en\",split=\"train\",streaming=True)\n",
    "    else:\n",
    "        ds=load_dataset(\"openslr/librispeech_asr\",\"clean\",split=\"train.100\",streaming=True)\n",
    "    ds=ds.cast_column(\"audio\",Audio(sampling_rate=16_000))\n",
    "    for ex in ds:\n",
    "        x=np.asarray(ex[\"audio\"][\"array\"],dtype=np.float32)\n",
    "        if len(x)<L16: continue\n",
    "        out.append(x[:L16])\n",
    "        if len(out)>=n: break\n",
    "    assert len(out)==n, f\"chỉ lấy được {len(out)}/{n} clip\"\n",
    "    return torch.from_numpy(np.stack(out))\n",
    "\n",
    "t0=time.perf_counter(); FIX_W = _take_clips(); print(f\"  lấy {N_FIX} clip từ {CORPUS}: {time.perf_counter()-t0:.0f}s\")\n",
    "with torch.no_grad():\n",
    "    FIX_MEL = mel_batch(FIX_W.to(DEVICE)).cpu()\n",
    "torch.save({\"w\":FIX_W,\"mel\":FIX_MEL,\"corpus\":CORPUS,\"seed\":FIX_SEED},\n",
    "           f\"{DRIVE}/preflight_fixed_batch_{CORPUS}.pt\")\n",
    "print(f\"  batch cố định đã lưu Drive: preflight_fixed_batch_{CORPUS}.pt\")\n",
    "\n",
    "@torch.no_grad()\n",
    "def encode_fixed(model):\n",
    "    model=model.to(DEVICE).eval(); hs=[]\n",
    "    for i in range(0,N_FIX,8):\n",
    "        hs.append(model(FIX_MEL[i:i+8].to(DEVICE)).cpu())\n",
    "    return torch.cat(hs)\n",
    "\n",
    "BASE_ENC = CM.collapse_metrics(encode_fixed(enc))\n",
    "BASE_TGT = CM.collapse_metrics(encode_fixed(tgt))\n",
    "print(\"\\n  MỐC KHOẺ MẠNH — đo trên chính JEPA.ckpt, quy cách mel của nhánh phát hiện\")\n",
    "print(\"  \" + CM.HEADER)\n",
    "print(\"  \" + CM.format_row(\"enc\", BASE_ENC))\n",
    "print(\"  \" + CM.format_row(\"tgt\", BASE_TGT))\n",
    "print(f\"\\n  tham chiếu lý thuyết trải đều trên mặt cầu D={EMBED_DIM}: std = {1/math.sqrt(EMBED_DIM):.4f}\")\n",
    "json.dump({\"encoder\":BASE_ENC,\"target_encoder\":BASE_TGT,\"corpus\":CORPUS,\"n_fixed\":N_FIX},\n",
    "          open(f\"{DRIVE}/preflight_baseline_{CORPUS}.json\",\"w\"),indent=1)\n",
    "print(f\"  mốc đã lưu Drive: preflight_baseline_{CORPUS}.json\")\n",
    "assert BASE_ENC[\"rank_eff\"]>10 and BASE_ENC[\"std_token\"]>0.005, \"checkpoint gốc đã trông như sụp đổ — dừng lại và xem lại mel\"\n",
    "print(\"\\nG4 XANH\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 5c · G4b — biểu diễn dị hướng: do checkpoint hay do ta đổi quy cách mel?\n",
    "# G4 đo trên quy cách PHÁT HIỆN cho cos_cross ~0,72 và std_utt ~0,006 (mốc trải đều 0,036).\n",
    "# Ô này đo LẠI cùng chỉ số nhưng bằng ĐÚNG cấu hình mà checkpoint được huấn luyện:\n",
    "# clip 10 s, hop 39,06 ms, cửa sổ 97,66 ms, f_max 16 kHz, patch (16,16), pos_embed GỐC.\n",
    "# Nếu cos_cross vẫn ~0,7  -> dị hướng là thuộc tính của checkpoint (kết quả cho bài).\n",
    "# Nếu cos_cross tụt hẳn   -> chính việc ta đổi quy cách mel gây ra, phải viết rõ trong bài.\n",
    "CLIP_S_O, HOP_O, FRAME_O, FMAX_O, PATCH_O = 10.0, 39.0625, 97.65625, 16000, (16,16)\n",
    "\n",
    "def mel_orig(w16):\n",
    "    w=_RS16(w16)[:, :int(CLIP_S_O*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_O,frame_shift=HOP_O,num_mel_bins=N_MELS,\n",
    "            low_freq=F_MIN,high_freq=FMAX_O,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 _take_clips_long(n=N_FIX, secs=CLIP_S_O):\n",
    "    L=int(16_000*secs); out=[]\n",
    "    if CORPUS==\"voxpopuli\":\n",
    "        ds=load_dataset(\"facebook/voxpopuli\",\"en\",split=\"train\",streaming=True)\n",
    "    else:\n",
    "        ds=load_dataset(\"openslr/librispeech_asr\",\"clean\",split=\"train.100\",streaming=True)\n",
    "    ds=ds.cast_column(\"audio\",Audio(sampling_rate=16_000))\n",
    "    for ex in ds:\n",
    "        x=np.asarray(ex[\"audio\"][\"array\"],dtype=np.float32)\n",
    "        if len(x)<L: continue\n",
    "        out.append(x[:L])\n",
    "        if len(out)>=n: break\n",
    "    assert len(out)==n, f\"chỉ lấy được {len(out)}/{n} clip {secs}s\"\n",
    "    return torch.from_numpy(np.stack(out))\n",
    "\n",
    "t0=time.perf_counter(); W10=_take_clips_long(); print(f\"  lấy {N_FIX} clip 10 s: {time.perf_counter()-t0:.0f}s\")\n",
    "with torch.no_grad(): MEL10 = mel_orig(W10.to(DEVICE)).cpu()\n",
    "\n",
    "enc_o = VisionTransformer(input_size=(T_BINS,N_MELS),patch_size=PATCH_O,in_chans=1,\n",
    "                          embed_dim=EMBED_DIM,depth=DEPTH,num_heads=HEADS,mlp_ratio=4.0,\n",
    "                          use_flash_attn=False)\n",
    "_m,_u = enc_o.load_state_dict(dict(_ENC), strict=False)      # GIỮ pos_embed gốc: ở đây nó ĐÚNG\n",
    "assert _m==[] and _u==[], (_m,_u)\n",
    "print(f\"  encoder gốc: lưới {enc_o.patch_embed.num_patches_h}x{enc_o.patch_embed.num_patches_w}\"\n",
    "      f\" = {enc_o.patch_embed.num_patches} token, pos_embed nạp nguyên từ checkpoint\")\n",
    "\n",
    "@torch.no_grad()\n",
    "def _enc_all(model, mel):\n",
    "    model=model.to(DEVICE).eval(); hs=[]\n",
    "    for i in range(0,mel.shape[0],8): hs.append(model(mel[i:i+8].to(DEVICE)).cpu())\n",
    "    return torch.cat(hs)\n",
    "\n",
    "M_ORIG = CM.collapse_metrics(_enc_all(enc_o, MEL10))\n",
    "enc_o = enc_o.cpu(); del enc_o; gc.collect(); torch.cuda.empty_cache()\n",
    "\n",
    "print(\"\\n  \" + CM.HEADER)\n",
    "print(\"  \" + CM.format_row(\"gốc\", M_ORIG))\n",
    "print(\"  \" + CM.format_row(\"p.hiện\", BASE_ENC))\n",
    "print(f\"\\n  tham chiếu trải đều D={EMBED_DIM}: {1/math.sqrt(EMBED_DIM):.4f}\")\n",
    "dc = M_ORIG[\"cos_cross\"] - BASE_ENC[\"cos_cross\"]\n",
    "print(f\"  Δcos_cross = {dc:+.3f}\")\n",
    "if abs(dc) < 0.15:\n",
    "    print(\"  => dị hướng có sẵn trong checkpoint, KHÔNG do ta đổi quy cách mel.\")\n",
    "    print(\"     Đây là số đo trực tiếp đầu tiên giải thích 'đóng băng ≈ khởi tạo ngẫu nhiên' (mục 3 bàn giao).\")\n",
    "else:\n",
    "    print(\"  => phần lớn dị hướng do ĐỔI QUY CÁCH MEL sinh ra. Phải viết rõ trong bài,\")\n",
    "    print(\"     và ngưỡng cổng của lượt tiền huấn luyện phải lấy theo mốc quy cách phát hiện.\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 6 · G5 — một bước tiền huấn luyện THẬT, và đo lại ước lượng giờ\n",
    "criterion = Loss(loss_type=\"norm_mse\", norm_pix_loss=True)\n",
    "enc_d, tgt_d, prd_d = enc.to(DEVICE), tgt.to(DEVICE), prd.to(DEVICE)\n",
    "for p in tgt_d.parameters(): p.requires_grad_(False)\n",
    "\n",
    "params = list(enc_d.parameters())+list(prd_d.parameters())\n",
    "opt = torch.optim.AdamW(params, lr=PRE_REF_LR, weight_decay=PRE_WD, betas=(0.9,0.95))\n",
    "sch = WarmupCosineScheduler(opt,warmup_steps=PRE_WARMUP,start_lr=PRE_START_LR,\n",
    "                            ref_lr=PRE_REF_LR,T_max=STEPS_TOTAL,final_lr=PRE_FINAL_LR)\n",
    "\n",
    "@torch.no_grad()\n",
    "def ema_update(student, teacher, tau):\n",
    "    for (_,sp),(_,tp) in zip(student.named_parameters(), teacher.named_parameters()):\n",
    "        tp.data.mul_(tau).add_(sp.data, alpha=1-tau)\n",
    "\n",
    "def train_step(mel, step_seed):\n",
    "    ctx,tgtm,_,_ = make_masks(mel.shape[0], seed=step_seed)\n",
    "    ctx=[m.to(DEVICE) for m in ctx]; tgtm=[m.to(DEVICE) for m in tgtm]\n",
    "    h = enc_d(mel, ctx)\n",
    "    z = prd_d(h, ctx, tgtm)\n",
    "    with torch.no_grad():\n",
    "        ht = apply_masks(tgt_d(mel), tgtm)\n",
    "    loss = criterion(z, ht)\n",
    "    opt.zero_grad(set_to_none=True); loss.backward(); opt.step(); sch.step()\n",
    "    ema_update(enc_d, tgt_d, EMA_TAU0)\n",
    "    return float(loss)\n",
    "\n",
    "# --- một bước, kiểm gradient chảy ---\n",
    "_mel = FIX_MEL[:PRE_BATCH].to(DEVICE)\n",
    "_before = enc_d.blocks[0].mlp.fc1.weight.detach().clone()\n",
    "_l = train_step(_mel, 0)\n",
    "_moved = (enc_d.blocks[0].mlp.fc1.weight.detach()-_before).abs().max().item()\n",
    "print(f\"  loss bước đầu = {_l:.4f}  (ngẫu nhiên hoàn toàn ≈ 2,0; sụp đổ hoàn toàn ≈ 0,0)\")\n",
    "assert math.isfinite(_l), \"loss NaN/inf ngay bước đầu\"\n",
    "assert _moved > 0, \"trọng số encoder KHÔNG đổi -> gradient không chảy\"\n",
    "print(f\"  encoder đã cập nhật: max|Δw| = {_moved:.3e}\")\n",
    "\n",
    "# --- thông lượng ---\n",
    "torch.cuda.synchronize(); t0=time.perf_counter()\n",
    "NW=12\n",
    "for s in range(NW): train_step(_mel, 100+s)\n",
    "torch.cuda.synchronize(); step_s=(time.perf_counter()-t0)/NW\n",
    "ms_sample = step_s/PRE_BATCH*1000\n",
    "\n",
    "torch.cuda.synchronize(); t0=time.perf_counter()\n",
    "with torch.no_grad(): mel_batch(FIX_W[:32].to(DEVICE))\n",
    "torch.cuda.synchronize(); mel_ms = (time.perf_counter()-t0)/32*1000\n",
    "\n",
    "N_SAMPLES = 2_812_500\n",
    "est_h = N_SAMPLES*ms_sample/1000/3600\n",
    "print(f\"\\n  ĐO THẬT trên {torch.cuda.get_device_name(0)}:\")\n",
    "print(f\"    bước tiền huấn luyện : {step_s*1000:7.1f} ms/bước  = {ms_sample:5.2f} ms/mẫu (batch {PRE_BATCH})\")\n",
    "print(f\"    dựng mel             : {mel_ms:7.2f} ms/mẫu\")\n",
    "print(f\"    -> một lượt {N_SAMPLES:,} mẫu = {est_h:.1f} giờ GPU   (tài liệu ước 9,4 h)\")\n",
    "print(f\"    -> mel cần {mel_ms/ms_sample:.2f} worker để theo kịp GPU; Colab có {os.cpu_count()} vCPU\"\n",
    "      f\" -> {'ĐỦ, tính mel tại chỗ được' if mel_ms/ms_sample < os.cpu_count() else 'KHÔNG ĐỦ, phải cache mel'}\")\n",
    "print(f\"\\n  ngân sách lại: N + T (2 lượt) = {2*(est_h+1.5+4):.1f} h; thêm mốc 300 h = {3*(est_h+1.5+4):.1f} h\")\n",
    "if est_h > 14: print(\"  ⚠ lệch quá xa ước lượng — xem lại kế hoạch TRƯỚC khi chạy\")\n",
    "print(\"\\nG5 XANH\" if math.isfinite(_l) else \"\\nG5 ĐỎ\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 6b · G5b — đo lại thông lượng với AMP fp16 (T4 có tensor core fp16)\n",
    "# G5 đo 640 ms/bước = 20,0 ms/mẫu = 15,6 h/lượt, ở fp32. Tài liệu ước 9,4 h.\n",
    "# T4 chạy fp16 nhanh hơn fp32 nhiều lần. Ô này đo thật rồi mới quyết.\n",
    "from torch.amp import autocast, GradScaler\n",
    "scaler = GradScaler(\"cuda\")\n",
    "\n",
    "def train_step_amp(mel, step_seed):\n",
    "    ctx,tgtm,_,_ = make_masks(mel.shape[0], seed=step_seed)\n",
    "    ctx=[m.to(DEVICE) for m in ctx]; tgtm=[m.to(DEVICE) for m in tgtm]\n",
    "    with autocast(\"cuda\", dtype=torch.float16):\n",
    "        h = enc_d(mel, ctx)\n",
    "        z = prd_d(h, ctx, tgtm)\n",
    "        with torch.no_grad(): ht = apply_masks(tgt_d(mel), tgtm)\n",
    "        loss = criterion(z, ht)\n",
    "    opt.zero_grad(set_to_none=True)\n",
    "    scaler.scale(loss).backward(); scaler.step(opt); scaler.update(); sch.step()\n",
    "    ema_update(enc_d, tgt_d, EMA_TAU0)\n",
    "    return float(loss.detach())\n",
    "\n",
    "def bench(fn, bs, n=10):\n",
    "    mel = FIX_MEL[:bs].to(DEVICE) if bs<=FIX_MEL.shape[0] else \\\n",
    "          FIX_MEL.repeat((bs//FIX_MEL.shape[0])+1,1,1,1)[:bs].to(DEVICE)\n",
    "    for s in range(3): fn(mel, 900+s)                       # hâm nóng\n",
    "    torch.cuda.synchronize(); t0=time.perf_counter()\n",
    "    for s in range(n): l=fn(mel, 1000+s)\n",
    "    torch.cuda.synchronize()\n",
    "    per = (time.perf_counter()-t0)/n\n",
    "    return per*1000, per/bs*1000, l\n",
    "\n",
    "N_SAMPLES=2_812_500\n",
    "rows=[]\n",
    "for name, fn, bs in [(\"fp32\",train_step,PRE_BATCH),(\"amp16\",train_step_amp,PRE_BATCH),\n",
    "                     (\"amp16\",train_step_amp,64)]:\n",
    "    try:\n",
    "        ms_step, ms_smp, l = bench(fn, bs)\n",
    "        rows.append((name,bs,ms_step,ms_smp,N_SAMPLES*ms_smp/1000/3600,l))\n",
    "    except RuntimeError as e:\n",
    "        print(f\"  {name} batch {bs}: {type(e).__name__} {str(e)[:80]}\")\n",
    "        torch.cuda.empty_cache()\n",
    "\n",
    "print(f\"\\n  {'chế độ':7s} {'batch':>5s} {'ms/bước':>9s} {'ms/mẫu':>8s} {'giờ/lượt':>9s} {'loss':>8s}\")\n",
    "for n_,b_,a_,c_,h_,l_ in rows:\n",
    "    print(f\"  {n_:7s} {b_:5d} {a_:9.1f} {c_:8.2f} {h_:9.1f} {l_:8.4f}\")\n",
    "if len(rows)>=2:\n",
    "    sp = rows[0][3]/min(r[3] for r in rows[1:])\n",
    "    best = min(rows[1:], key=lambda r:r[3])\n",
    "    print(f\"\\n  AMP nhanh hơn {sp:.2f} lần -> {best[4]:.1f} h/lượt thay vì {rows[0][4]:.1f} h\")\n",
    "    print(f\"  ngân sách hai nhánh N+T = {2*(best[4]+1.5+4):.1f} h (fp32 là {2*(rows[0][4]+1.5+4):.1f} h)\")\n",
    "    print(f\"  chênh loss fp32 {rows[0][5]:.4f} vs amp {best[5]:.4f} — phải cùng cỡ, lệch lớn là AMP có vấn đề\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 7 · G6 — lưu/khôi phục checkpoint (đĩa cục bộ; KHÔNG ghi thẳng Drive)\n",
    "import shutil\n",
    "\n",
    "def _free_gb(p):\n",
    "    try: return shutil.disk_usage(p).free/2**30\n",
    "    except Exception as e: return float(\"nan\")\n",
    "\n",
    "def cpu_sd(m): return {k: v.detach().cpu() for k, v in m.state_dict().items()}\n",
    "\n",
    "LOCAL_CKPT = f\"{BASE}/pre_last.pt\"\n",
    "t0 = time.perf_counter()\n",
    "torch.save({\"epoch\":0,\"step\":int(sch.last_epoch),\n",
    "            \"encoder\":cpu_sd(enc_d),\"target_encoder\":cpu_sd(tgt_d),\"predictor\":cpu_sd(prd_d),\n",
    "            \"opt\":opt.state_dict(),\"sch\":sch.state_dict(),\n",
    "            \"metrics\":{\"base_enc\":BASE_ENC,\"base_tgt\":BASE_TGT},\n",
    "            \"cfg\":{\"patch\":PATCH,\"grid\":(GRID_H,GRID_W),\"ref_lr\":PRE_REF_LR,\"wd\":PRE_WD,\n",
    "                   \"warmup\":PRE_WARMUP,\"tau0\":EMA_TAU0,\"corpus\":CORPUS}},\n",
    "           LOCAL_CKPT)\n",
    "save_s = time.perf_counter()-t0\n",
    "size_mb = os.path.getsize(LOCAL_CKPT)/1e6\n",
    "print(f\"  lưu {size_mb:.0f} MB ra đĩa cục bộ trong {save_s:.1f}s -> {size_mb/save_s:.0f} MB/s\")\n",
    "\n",
    "# --- khôi phục: so từng bit, KHÔNG dựng lại mô hình (đỡ 344 MB) ---\n",
    "blob = torch.load(LOCAL_CKPT, map_location=\"cpu\", weights_only=False)\n",
    "cur = enc_d.state_dict()\n",
    "bad = [k for k, v in cur.items() if not torch.equal(v.detach().cpu(), blob[\"encoder\"][k])]\n",
    "assert not bad, bad[:5]\n",
    "# --- SỬA LỖI: opt bao cả encoder LẪN predictor, nên optimizer khôi phục phải bao đúng ngần ấy.\n",
    "#     Bản cũ dựng o2 chỉ trên encoder -> ValueError \"parameter group that doesn't match\".\n",
    "opt2 = torch.optim.AdamW(list(enc_d.parameters())+list(prd_d.parameters()),\n",
    "                         lr=PRE_REF_LR, weight_decay=PRE_WD, betas=(0.9,0.95))\n",
    "opt2.load_state_dict(blob[\"opt\"])\n",
    "n_state = len(opt2.state_dict()[\"state\"])\n",
    "assert blob[\"step\"] == int(sch.last_epoch)\n",
    "print(f\"  khôi phục: encoder trùng từng bit trên {len(blob['encoder'])} tensor; \"\n",
    "      f\"optimizer {n_state} tensor có moment; bước {blob['step']}\")\n",
    "del blob; gc.collect()\n",
    "\n",
    "# --- Drive: chỉ THĂM DÒ, không ghi checkpoint vào đó nữa ---\n",
    "probe = f\"{DRIVE}/_probe.bin\"\n",
    "drive_ok, drive_msg = False, \"\"\n",
    "try:\n",
    "    with open(probe,\"wb\") as fh: fh.write(b\"\\0\"*(8*1024*1024))   # 8 MB\n",
    "    assert os.path.getsize(probe)==8*1024*1024\n",
    "    os.remove(probe); drive_ok=True; drive_msg=\"ghi/đọc/xoá 8 MB OK\"\n",
    "except Exception as e:\n",
    "    drive_msg=f\"{type(e).__name__}: {e}\"\n",
    "print(f\"\\n  Drive: {'DÙNG ĐƯỢC' if drive_ok else 'HỎNG'} — {drive_msg}\")\n",
    "print(f\"  chỗ trống: đĩa Colab {_free_gb('/content'):.0f} GB | mount Drive {_free_gb(DRIVE):.1f} GB\")\n",
    "\n",
    "# --- ngân sách lưu trữ cho lượt chạy thật ---\n",
    "roll = 2*size_mb/1000            # giữ 2 bản luân phiên\n",
    "mile = 3*(size_mb*0.42)/2/1000   # 3 mốc, chỉ enc+tgt, fp16\n",
    "print(f\"\\n  NGÂN SÁCH LƯU TRỮ một nhánh:\")\n",
    "print(f\"    lưu MỖI epoch x 20 epoch, mỗi bản {size_mb:.0f} MB  = {20*size_mb/1000:.1f} GB  <-- kế hoạch cũ, KHÔNG khả thi\")\n",
    "print(f\"    luân phiên 2 bản + 3 mốc enc/tgt fp16              = {roll+mile:.1f} GB  <-- kế hoạch mới\")\n",
    "print(f\"    hai nhánh N và T                                   = {2*(roll+mile):.1f} GB\")\n",
    "print(\"\\nG6 XANH\" if not bad else \"\\nG6 ĐỎ\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#@title 10 · Tổng kết chín cổng\n",
    "print(f\"\"\"\n",
    "G1 hình học + pos_embed sinh lại + không hồi quy encoder  ... ô 2\n",
    "G2 lr thật = {PRE_REF_LR:.0e} (không phải 1e-3 của repo)  ... ô 3\n",
    "G3 mel {T_BINS}x{N_MELS}, 2 hàng đệm, mặt nạ rắc ngẫu nhiên ... ô 4\n",
    "G4 mốc sụp đổ trên JEPA.ckpt                              ... ô 5b\n",
    "G4b dị hướng: quy cách gốc so với quy cách phát hiện       ... ô 5c\n",
    "G5 một bước thật chạy được (fp32)                          ... ô 6\n",
    "G5b thông lượng với AMP fp16 -> CHỌN CHẾ ĐỘ CHẠY           ... ô 6b\n",
    "G6 lưu/khôi phục checkpoint ĐĨA CỤC BỘ + ngân sách lưu trữ ... ô 7\n",
    "G7 = G1e (encoder trùng từng bit với v38/v40)\n",
    "\n",
    "Ba việc KHÔNG nằm trong notebook này, phải xong trước khi chạy lượt thật:\n",
    "  * chọn nơi lưu checkpoint bền — Drive KHÔNG đủ chỗ (0,0 GB trống)\n",
    "  * seed 2 và 3 cho khởi tạo ngẫu nhiên fine-tune (~1,5 h)\n",
    "  * cổng aasist: ASVspoof19 0,83 · DFADD 39,05 · LibriSeVoc 37,95\n",
    "\"\"\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Đọc kết quả\n",
    "\n",
    "**Đi tiếp khi và chỉ khi cả bảy cổng xanh.** Cổng đỏ nào cũng dừng — không cổng nào ở đây tốn quá vài phút để sửa, còn một lượt chạy hỏng tốn mười tiếng.\n",
    "\n",
    "Ba con số cần nhìn kỹ, kể cả khi mọi thứ xanh:\n",
    "\n",
    "**Đỉnh lr ở ô 3.** Phải là `1e-04`. Nếu thấy `1e-03` nghĩa là `ref_lr` chưa được truyền và lượt chạy sẽ xoá biểu diễn AudioSet trong epoch đầu — đây là cách rẻ nhất để mất mười tiếng.\n",
    "\n",
    "**`est_h` ở ô 6.** Tài liệu ước 9,4 giờ cho một lượt 2,81 M mẫu, dựa trên phép ngoại suy từ chi phí fine-tune. Ô 6 đo thật. Lệch dưới 20 % thì kế hoạch giữ nguyên; lệch trên 50 % thì phải tính lại ngân sách **trước** khi bắt đầu, không phải sau.\n",
    "\n",
    "**Mốc bốn chỉ số ở ô 5b.** Đây là số duy nhất làm cho mọi con số về sau đọc được. `std_token` của một encoder khoẻ nên ở gần `1/sqrt(768) = 0,036`. Nếu checkpoint gốc đã cho chỉ số trông như sụp đổ thì vấn đề nằm ở quy cách mel, không phải ở mô hình — dừng và xem lại ô 4.\n",
    "\n",
    "Trong lúc tiền huấn luyện, in bốn chỉ số **mỗi epoch** và so với mốc bằng `CM.check_gate`. Nhắc lại điều dễ quên nhất: **loss của công thức này giảm khi biểu diễn sụp đổ**, nên đường loss đẹp không phải bằng chứng gì cả.\n"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  },
  "colab": {
   "provenance": [],
   "gpuType": "T4"
  },
  "accelerator": "GPU"
 },
 "nbformat": 4,
 "nbformat_minor": 0
}