# =====================================================================================
# BA Ô THAY THẾ / BỔ SUNG cho audio_jepa_v41_preflight_tien_huan_luyen.ipynb
# Dán thẳng vào phiên Colab ĐANG CHẠY — không cần khởi động lại, không mất G1..G5.
# Lý do: xem "02-preflight-ket-qua.md" mục "Lượt chạy thật 09/09".
#   1) Ô 7 cũ ghi 1,65 GB thẳng vào Drive -> Drive đầy (14,99/15 GB) -> chết.
#      Ô 7 mới ghi ra ĐĨA CỤC BỘ (còn 186 GB), và sửa lỗi khôi phục optimizer.
#   2) Ô G4b: tách "biểu diễn dị hướng" là do checkpoint hay do ta đổi quy cách mel.
#   3) Ô G5b: đo lại thông lượng với AMP fp16 — 640 ms/bước hiện tại là fp32.
# =====================================================================================


# ══════════════════════════════ Ô 7 THAY THẾ — G6 ══════════════════════════════
#@title 7 · G6 — lưu/khôi phục checkpoint (đĩa cục bộ; KHÔNG ghi thẳng Drive)
import shutil

def _free_gb(p):
    try: return shutil.disk_usage(p).free/2**30
    except Exception as e: return float("nan")

def cpu_sd(m): return {k: v.detach().cpu() for k, v in m.state_dict().items()}

LOCAL_CKPT = f"{BASE}/pre_last.pt"
t0 = time.perf_counter()
torch.save({"epoch":0,"step":int(sch.last_epoch),
            "encoder":cpu_sd(enc_d),"target_encoder":cpu_sd(tgt_d),"predictor":cpu_sd(prd_d),
            "opt":opt.state_dict(),"sch":sch.state_dict(),
            "metrics":{"base_enc":BASE_ENC,"base_tgt":BASE_TGT},
            "cfg":{"patch":PATCH,"grid":(GRID_H,GRID_W),"ref_lr":PRE_REF_LR,"wd":PRE_WD,
                   "warmup":PRE_WARMUP,"tau0":EMA_TAU0,"corpus":CORPUS}},
           LOCAL_CKPT)
save_s = time.perf_counter()-t0
size_mb = os.path.getsize(LOCAL_CKPT)/1e6
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")

# --- khôi phục: so từng bit, KHÔNG dựng lại mô hình (đỡ 344 MB) ---
blob = torch.load(LOCAL_CKPT, map_location="cpu", weights_only=False)
cur = enc_d.state_dict()
bad = [k for k, v in cur.items() if not torch.equal(v.detach().cpu(), blob["encoder"][k])]
assert not bad, bad[:5]
# --- 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.
#     Bản cũ dựng o2 chỉ trên encoder -> ValueError "parameter group that doesn't match".
opt2 = torch.optim.AdamW(list(enc_d.parameters())+list(prd_d.parameters()),
                         lr=PRE_REF_LR, weight_decay=PRE_WD, betas=(0.9,0.95))
opt2.load_state_dict(blob["opt"])
n_state = len(opt2.state_dict()["state"])
assert blob["step"] == int(sch.last_epoch)
print(f"  khôi phục: encoder trùng từng bit trên {len(blob['encoder'])} tensor; "
      f"optimizer {n_state} tensor có moment; bước {blob['step']}")
del blob; gc.collect()

# --- Drive: chỉ THĂM DÒ, không ghi checkpoint vào đó nữa ---
probe = f"{DRIVE}/_probe.bin"
drive_ok, drive_msg = False, ""
try:
    with open(probe,"wb") as fh: fh.write(b"\0"*(8*1024*1024))   # 8 MB
    assert os.path.getsize(probe)==8*1024*1024
    os.remove(probe); drive_ok=True; drive_msg="ghi/đọc/xoá 8 MB OK"
except Exception as e:
    drive_msg=f"{type(e).__name__}: {e}"
print(f"\n  Drive: {'DÙNG ĐƯỢC' if drive_ok else 'HỎNG'} — {drive_msg}")
print(f"  chỗ trống: đĩa Colab {_free_gb('/content'):.0f} GB | mount Drive {_free_gb(DRIVE):.1f} GB")

# --- ngân sách lưu trữ cho lượt chạy thật ---
roll = 2*size_mb/1000            # giữ 2 bản luân phiên
mile = 3*(size_mb*0.42)/2/1000   # 3 mốc, chỉ enc+tgt, fp16
print(f"\n  NGÂN SÁCH LƯU TRỮ một nhánh:")
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")
print(f"    luân phiên 2 bản + 3 mốc enc/tgt fp16              = {roll+mile:.1f} GB  <-- kế hoạch mới")
print(f"    hai nhánh N và T                                   = {2*(roll+mile):.1f} GB")
print("\nG6 XANH" if not bad else "\nG6 ĐỎ")


# ══════════════════════════ Ô G4b MỚI — dị hướng do đâu ══════════════════════════
#@title 5c · G4b — biểu diễn dị hướng: do checkpoint hay do ta đổi quy cách mel?
# 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ày đo LẠI cùng chỉ số nhưng bằng ĐÚNG cấu hình mà checkpoint được huấn luyệ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ế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ế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.
CLIP_S_O, HOP_O, FRAME_O, FMAX_O, PATCH_O = 10.0, 39.0625, 97.65625, 16000, (16,16)

def mel_orig(w16):
    w=_RS16(w16)[:, :int(CLIP_S_O*SR_MODEL)]; out=[]
    for i in range(w.shape[0]):
        v=w[i:i+1]
        f=torchaudio.compliance.kaldi.fbank(v-v.mean(),sample_frequency=SR_MODEL,
            frame_length=FRAME_O,frame_shift=HOP_O,num_mel_bins=N_MELS,
            low_freq=F_MIN,high_freq=FMAX_O,use_log_fbank=True,window_type="hanning")
        n=f.shape[0]
        if n<T_BINS: f=torch.cat([f,torch.zeros(T_BINS-n,N_MELS,device=f.device)],0)
        out.append(f[:T_BINS])
    return torch.stack(out).unsqueeze(1)

def _take_clips_long(n=N_FIX, secs=CLIP_S_O):
    L=int(16_000*secs); out=[]
    if CORPUS=="voxpopuli":
        ds=load_dataset("facebook/voxpopuli","en",split="train",streaming=True)
    else:
        ds=load_dataset("openslr/librispeech_asr","clean",split="train.100",streaming=True)
    ds=ds.cast_column("audio",Audio(sampling_rate=16_000))
    for ex in ds:
        x=np.asarray(ex["audio"]["array"],dtype=np.float32)
        if len(x)<L: continue
        out.append(x[:L])
        if len(out)>=n: break
    assert len(out)==n, f"chỉ lấy được {len(out)}/{n} clip {secs}s"
    return torch.from_numpy(np.stack(out))

t0=time.perf_counter(); W10=_take_clips_long(); print(f"  lấy {N_FIX} clip 10 s: {time.perf_counter()-t0:.0f}s")
with torch.no_grad(): MEL10 = mel_orig(W10.to(DEVICE)).cpu()

enc_o = VisionTransformer(input_size=(T_BINS,N_MELS),patch_size=PATCH_O,in_chans=1,
                          embed_dim=EMBED_DIM,depth=DEPTH,num_heads=HEADS,mlp_ratio=4.0,
                          use_flash_attn=False)
_m,_u = enc_o.load_state_dict(dict(_ENC), strict=False)      # GIỮ pos_embed gốc: ở đây nó ĐÚNG
assert _m==[] and _u==[], (_m,_u)
print(f"  encoder gốc: lưới {enc_o.patch_embed.num_patches_h}x{enc_o.patch_embed.num_patches_w}"
      f" = {enc_o.patch_embed.num_patches} token, pos_embed nạp nguyên từ checkpoint")

@torch.no_grad()
def _enc_all(model, mel):
    model=model.to(DEVICE).eval(); hs=[]
    for i in range(0,mel.shape[0],8): hs.append(model(mel[i:i+8].to(DEVICE)).cpu())
    return torch.cat(hs)

M_ORIG = CM.collapse_metrics(_enc_all(enc_o, MEL10))
enc_o = enc_o.cpu(); del enc_o; gc.collect(); torch.cuda.empty_cache()

print("\n  " + CM.HEADER)
print("  " + CM.format_row("gốc", M_ORIG))
print("  " + CM.format_row("p.hiện", BASE_ENC))
print(f"\n  tham chiếu trải đều D={EMBED_DIM}: {1/math.sqrt(EMBED_DIM):.4f}")
dc = M_ORIG["cos_cross"] - BASE_ENC["cos_cross"]
print(f"  Δcos_cross = {dc:+.3f}")
if abs(dc) < 0.15:
    print("  => dị hướng có sẵn trong checkpoint, KHÔNG do ta đổi quy cách mel.")
    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).")
else:
    print("  => phần lớn dị hướng do ĐỔI QUY CÁCH MEL sinh ra. Phải viết rõ trong bài,")
    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.")


# ══════════════════════ Ô G5b MỚI — thông lượng với AMP ══════════════════════
#@title 6b · G5b — đo lại thông lượng với AMP fp16 (T4 có tensor core fp16)
# 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.
# T4 chạy fp16 nhanh hơn fp32 nhiều lần. Ô này đo thật rồi mới quyết.
from torch.amp import autocast, GradScaler
scaler = GradScaler("cuda")

def train_step_amp(mel, step_seed):
    ctx,tgtm,_,_ = make_masks(mel.shape[0], seed=step_seed)
    ctx=[m.to(DEVICE) for m in ctx]; tgtm=[m.to(DEVICE) for m in tgtm]
    with autocast("cuda", dtype=torch.float16):
        h = enc_d(mel, ctx)
        z = prd_d(h, ctx, tgtm)
        with torch.no_grad(): ht = apply_masks(tgt_d(mel), tgtm)
        loss = criterion(z, ht)
    opt.zero_grad(set_to_none=True)
    scaler.scale(loss).backward(); scaler.step(opt); scaler.update(); sch.step()
    ema_update(enc_d, tgt_d, EMA_TAU0)
    return float(loss.detach())

def bench(fn, bs, n=10):
    mel = FIX_MEL[:bs].to(DEVICE) if bs<=FIX_MEL.shape[0] else \
          FIX_MEL.repeat((bs//FIX_MEL.shape[0])+1,1,1,1)[:bs].to(DEVICE)
    for s in range(3): fn(mel, 900+s)                       # hâm nóng
    torch.cuda.synchronize(); t0=time.perf_counter()
    for s in range(n): l=fn(mel, 1000+s)
    torch.cuda.synchronize()
    per = (time.perf_counter()-t0)/n
    return per*1000, per/bs*1000, l

N_SAMPLES=2_812_500
rows=[]
for name, fn, bs in [("fp32",train_step,PRE_BATCH),("amp16",train_step_amp,PRE_BATCH),
                     ("amp16",train_step_amp,64)]:
    try:
        ms_step, ms_smp, l = bench(fn, bs)
        rows.append((name,bs,ms_step,ms_smp,N_SAMPLES*ms_smp/1000/3600,l))
    except RuntimeError as e:
        print(f"  {name} batch {bs}: {type(e).__name__} {str(e)[:80]}")
        torch.cuda.empty_cache()

print(f"\n  {'chế độ':7s} {'batch':>5s} {'ms/bước':>9s} {'ms/mẫu':>8s} {'giờ/lượt':>9s} {'loss':>8s}")
for n_,b_,a_,c_,h_,l_ in rows:
    print(f"  {n_:7s} {b_:5d} {a_:9.1f} {c_:8.2f} {h_:9.1f} {l_:8.4f}")
if len(rows)>=2:
    sp = rows[0][3]/min(r[3] for r in rows[1:])
    best = min(rows[1:], key=lambda r:r[3])
    print(f"\n  AMP nhanh hơn {sp:.2f} lần -> {best[4]:.1f} h/lượt thay vì {rows[0][4]:.1f} h")
    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)")
    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 đề")
