"""Tính lại ví dụ lớp 1; không chạy checkpoint/dataset.

Python 3 + NumPy + Matplotlib. Output tương đối với chính file script.
"""
from pathlib import Path
from datetime import datetime, timezone
import json
import math
import platform
import re
import sys
from urllib.parse import unquote

import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

ROOT = Path(__file__).resolve().parents[1]
CHECKS = []


def close(name, actual, expected, atol=1e-10):
    actual, expected = np.asarray(actual), np.asarray(expected)
    error = float(np.max(np.abs(actual - expected)))
    assert np.allclose(actual, expected, atol=atol, rtol=0), (name, error)
    CHECKS.append({"name": name, "passed": True, "max_abs_error": error,
                   "absolute_tolerance": atol})


def erank(spectrum):
    spectrum = np.asarray(spectrum, dtype=float)
    if spectrum.sum() == 0:
        raise ValueError("All-zero spectrum: undefined")
    p = spectrum[spectrum > 0] / spectrum.sum()
    return np.exp(-np.sum(p * np.log(p)))


def cka(x, y):
    x, y = np.asarray(x, float), np.asarray(y, float)
    x, y = x - x.mean(0), y - y.mean(0)
    return (np.linalg.norm(x.T @ y, "fro") ** 2 /
            (np.linalg.norm(x.T @ x, "fro") * np.linalg.norm(y.T @ y, "fro")))


def softmax(scores):
    scores = scores - scores.max(axis=-1, keepdims=True)
    values = np.exp(scores)
    return values / values.sum(axis=-1, keepdims=True)


close("B1 F0: 160 samples at 16 kHz", 16000 / 160, 100)
close("B1 source-filter convolution", np.convolve([1, 0, 1], [1, .5]),
      [1, .5, 1, .5])
close("B1 exercise convolution", np.convolve([2, 0, 1], [1, .5]),
      [2, 1, 1, .5])
close("B2 convolution", np.convolve([1, 2, -1], [1, .5]), [1, 2.5, 0, -.5])
x = np.array([1, 0, -1, 0])
X = np.fft.fft(x)
close("B2 DFT", X, [0, 2, 0, 2])
close("B2 inverse DFT", np.fft.ifft(X), x)
Y = np.fft.fft(np.roll(x, 1))
close("B2 shift phase", Y, [0, -2j, 0, 2j])
close("B2 same magnitude after circular shift", np.abs(Y), np.abs(X))
close("B2 circular convolution N=2",
      np.fft.ifft(np.fft.fft([1, 2]) * np.fft.fft([1, 1])), [3, 3])
close("B2 padded linear convolution N=3",
      np.fft.ifft(np.fft.fft([1, 2], 3) * np.fft.fft([1, 1], 3)), [1, 3, 2])

n = np.arange(64)
close("B3 alias 10/6 kHz at 16 kHz",
      np.cos(2*np.pi*10000*n/16000), np.cos(2*np.pi*6000*n/16000))
close("B3 cosine table", np.cos(2*np.pi*6000*np.arange(5)/16000),
      [1, -1/math.sqrt(2), 0, 1/math.sqrt(2), -1])
close("B3 Nyquist sine disappears", np.sin(np.pi*n), np.zeros(64))
close("B3 sample counts", [2.56*16000, 2.56*32000], [40960, 81920])
close("B3 linear interpolation", np.interp(np.arange(0, 2.01, .5),
      [0, 1, 2], [0, 1, 0]), [0, .5, 1, .5, 0])

frames = np.array([x, -x])
spectra = np.fft.fft(frames, axis=1)
close("B4 two STFT frames", spectra, [[0, 2, 0, 2], [0, -2, 0, -2]])
close("B4 Parseval", np.sum(np.abs(spectra)**2, axis=1)/4, [2, 2])
one_sided = np.abs(np.fft.rfft(x))**2
close("B4 one-sided energy", (one_sided[0]+2*one_sided[1]+one_sided[2])/4, 2)
close("B4 frame counts by convention",
      [1+(81920-800)//320, 1+(81920-1024)//320, 1+81920//320], [254, 253, 257])
close("B4 eight-frame spans in ms", [8*10, 7*10, 7*10+25], [80, 70, 95])
close("B4 four-frame spans in ms", [4*10, 3*10, 3*10+25], [40, 30, 55])
close("B4 power dB", 10*np.log10(4), 6.020599913279624)
close("B4 complex example", [abs(3+4j), abs(3+4j)**2, np.angle(3+4j)],
      [5, 25, .9272952180016122])

W = np.array([[1, .5, 0], [0, .5, 1]])
close("B5 filterbank", W @ [2, 4, 6], [4, 8])
close("B5 non-injective filterbank", W @ [3, 2, 7], [4, 8])
close("B5 filterbank nullspace", W @ [-.5, 1, -.5], [0, 0])
S = np.log([4, 8])
D = np.array([[1, 1], [1, -1]]) / np.sqrt(2)
c = D @ S
close("B5 orthonormal DCT", c, [2.450645, -.490129], atol=5e-7)
close("B5 full DCT inverse", D.T @ c, S)
close("B5 truncated DCT inverse", D.T @ [c[0], 0], [np.log(np.sqrt(32))]*2)
matrix = np.array([[1., 2.], [3., 4.]])
close("B5 temporal CMVN", (matrix-matrix.mean(0))/matrix.std(0),
      [[-1, -1], [1, 1]])
close("B5 feature LayerNorm toy",
      (matrix-matrix.mean(1, keepdims=True))/matrix.std(1, keepdims=True),
      [[-1, 1], [-1, 1]])
close("B5 padding changes mean", np.vstack([matrix, [0, 0]]).mean(0), [4/3, 2])
close("B5 patch grid", [256//8, 128//32, (256//8)*(128//32), 8*32],
      [32, 4, 128, 256])

posterior_train = .6*.1/(.6*.1+.2*.9)
posterior_deploy = .6*.99/(.6*.99+.2*.01)
close("B6 Bayes train", posterior_train, .25)
close("B6 Bayes deploy", posterior_deploy, .9966442953020135)
close("B6 logits", [1/(1+np.exp(-2)), 1/(1+np.exp(-3))],
      [.8807970779778823, .9525741268224334])
close("B6 weighted optimum", 9*.25/(9*.25+.75), .75)
close("B6 weighted logit decomposition", np.log(.75/.25),
      np.log(3)+np.log(.1/.9)+np.log(9))
close("B6 weighted exercise", 4*.2/(4*.2+.8), .5)
close("B6 batch-size-one weight cancellation", 9*(-np.log(.8))/9, -np.log(.8))
target, prob = np.array([-2., 2.]), np.array([.25, .75])
mean = np.dot(prob, target)
close("B6 target mean", mean, 1)
close("B6 MSE predictions", [np.dot(prob, (target-a)**2) for a in [1, 2, 0]],
      [3, 4, 4])
close("B6 random independent target MSE",
      np.sum(prob[:, None]*prob[None, :]*(target[:, None]-target[None, :])**2), 6)
close("B6 MSE decomposition at a=3",
      np.dot(prob, (target-3)**2), np.dot(prob, (target-mean)**2)+(3-mean)**2)
close("B6 cost threshold", 100*.01/.99, 1.0101010101010102)

close("B7 norm/cosine", [np.linalg.norm([3, 4]), np.linalg.norm([6, 8]),
      np.dot([3, 4], [6, 8])/(5*10)], [5, 10, 1])
close("B7 high cosine counterexample",
      np.dot([10, 1], [10, -1])/101, .9801980198019802)
H = np.array([[-2, 0], [-1, 0], [1, 0], [2, 0]], dtype=float)
C = np.cov(H, rowvar=False, ddof=1)
close("B7 rank-one covariance", C, [[10/3, 0], [0, 0]])
singular = np.linalg.svd(H-H.mean(0), compute_uv=False)
close("B7 covariance/SVD eigenvalues", np.linalg.eigvalsh(C)[::-1],
      singular**2/(len(H)-1))
close("B7 center two-row covariance",
      np.cov([[1, 2], [3, 4]], rowvar=False, ddof=1), [[2, 2], [2, 2]])
nuisance = np.array([[10*u, v] for u in [-1, 1] for v in [-1, 1]])
close("B7 nuisance PCA covariance", np.cov(nuisance, rowvar=False),
      [[400/3, 0], [0, 4/3]])
close("B7 singular/eigen entropy difference", [erank([3, 1]), erank([9, 1])],
      [1.754765, 1.384145], atol=5e-7)
close("B7 entropy scale invariance", erank(np.array([3., 1.])*1e-8), erank([3, 1]))
try:
    erank([0, 0])
    raise AssertionError("zero-spectrum must be undefined")
except ValueError:
    CHECKS.append({"name": "B7 zero-spectrum rejected", "passed": True})
cx = np.array([-1, 0, 1])[:, None]
close("B7 CKA scale", cka(cx, 2*cx), 1)
close("B7 CKA zero", cka(cx, np.array([1, -2, 1])[:, None]), 0)
xor = np.array([[-1, -1], [-1, 1], [1, -1], [1, 1]])
close("B7 XOR covariance", np.cov(xor, rowvar=False), [[4/3, 0], [0, 4/3]])
close("B7 centered sample rank bound", np.linalg.matrix_rank(H-H.mean(0)), 1)

params = np.array([.5, 0, 1, 0])  # w, b, v, c


def loss(theta, input_x=2):
    w, b, v, c = theta
    d = v*max(0, w*input_x+b)+c
    return np.logaddexp(0, -d)


g = 1/(1+math.exp(-1))-1
analytic = np.array([2*g, g, g, g])
numeric = []
for index in range(4):
    step = np.zeros(4)
    step[index] = 1e-6
    numeric.append((loss(params+step)-loss(params-step))/(2e-6))
close("B8 backprop checked by finite differences", numeric, analytic, atol=1e-8)
close("B8 input gradient finite difference",
      (loss(params, 2+1e-6)-loss(params, 2-1e-6))/(2e-6), .5*g, atol=1e-8)
new_params = params-.1*analytic
close("B8 simultaneous SGD update", new_params,
      [.5537882842739991, .026894142136999512, 1.0268941421369995, .026894142136999512])
close("B8 loss after SGD", loss(new_params), .2651689745906373)
input_conv = np.array([1, 2, 0, 1])
conv_windows = np.lib.stride_tricks.sliding_window_view(input_conv, 2)
conv_out = conv_windows @ [1, -1]
close("B8 CNN cross-correlation", conv_out, [-1, 2, -1])
close("B8 shared kernel gradient", conv_windows.T @ conv_out, [3, -3])
Q, K, V = np.array([[1.], [0.]]), np.array([[np.log(3)], [0.]]), np.array([[2.], [6.]])
A = softmax(Q @ K.T)
close("B8 attention weights", A, [[.75, .25], [.5, .5]])
close("B8 attention output", A @ V, [[3], [4]])
close("B8 heads and pair counts", [768//12, 128**2, 256**2], [64, 16384, 65536])
pool_mean = np.dot([.25, .75], [0., 2.])
pool_var = np.dot([.25, .75], np.array([0., 2.])**2)-pool_mean**2
close("B8 attentive statistics", [pool_mean, pool_var, np.sqrt(pool_var)],
      [1.5, .75, .8660254037844386])
close("B8 reported head count from layer dimensions",
      [768*128+128+128+1, 1536*2+1536*256+256+256*2+2,
       13+(768*128+128+128+1)+(1536*2+1536*256+256+256*2+2)],
      [98561, 397058, 495632])
m1, v1 = (1-.9)*.5, (1-.999)*.5**2
mh, vh = m1/(1-.9), v1/(1-.999)
close("B8 AdamW first step", (1-.1*.01)*2-.1*mh/np.sqrt(vh), 1.898)
close("B8 Adam L2 vs AdamW zero-loss-gradient", [2-.1*.2/(.2+1e-8), 2-.1*.1*2],
      [1.9, 1.98], atol=1e-7)
close("B8 LoRA parameter counts", [8*(768+768), 768**2], [12288, 589824])

# Figures use generated signals only.
assets = ROOT / "assets"
assets.mkdir(exist_ok=True)
plt.rcParams.update({"font.family": "DejaVu Sans", "font.size": 11})
fig, ax = plt.subplots(figsize=(10, 5), layout="constrained")
t = np.linspace(0, .0005, 1500)
sample_t = np.arange(9)/16000
ax.plot(t*1000, np.cos(2*np.pi*6000*t), label="Cosine 6 kHz", lw=2, color="#126782")
ax.plot(t*1000, np.cos(2*np.pi*10000*t), label="Cosine 10 kHz", lw=1.8, color="#d97706")
ax.scatter(sample_t*1000, np.cos(2*np.pi*6000*sample_t),
           color="#111827", s=38, zorder=5, label="Mẫu tại 16 kHz — hai sóng trùng nhau")
ax.set(xlabel="Thời gian (ms)", ylabel="Biên độ quy ước",
       title="Aliasing: hai tần số khác nhau cho cùng dãy mẫu",
       xlim=(0, .5), ylim=(-1.2, 1.2))
ax.grid(alpha=.2)
ax.legend(loc="upper center", bbox_to_anchor=(.5, -.22), ncol=2, fontsize=9)
fig.savefig(assets/"aliasing-16khz.png", dpi=160)
plt.close(fig)

fs, length = 32000, 800
tone_t = np.arange(length)/fs
tone = (np.cos(2*np.pi*1000*tone_t)+np.cos(2*np.pi*1020*tone_t))*np.hanning(length)
close("B4 zero-padding samples same spectrum on shared grid",
      np.fft.rfft(tone, 8192)[::8], np.fft.rfft(tone, 1024))
fig, ax = plt.subplots(figsize=(10, 4.5), layout="constrained")
for nfft, style in [(8192, "-"), (1024, "o")]:
    values = abs(np.fft.rfft(tone, n=nfft))
    frequencies = np.fft.rfftfreq(nfft, 1/fs)
    ax.plot(frequencies, values, style, markersize=5, lw=2,
            label=f"n_fft={nfft:,}; cùng 800 mẫu Hann (25 ms)")
ax.axvline(1000, ls=":", color="#6b7280")
ax.axvline(1020, ls=":", color="#6b7280")
ax.set(xlim=(850, 1170), xlabel="Tần số (Hz)", ylabel="Magnitude DFT (chưa normalized)",
       title="Zero-padding làm lưới mịn hơn; hai tone vẫn chung một vùng đỉnh")
ax.grid(alpha=.2)
ax.legend(fontsize=9)
fig.savefig(assets/"zero-padding-resolution.png", dpi=160)
plt.close(fig)

# Verify local Markdown links and table columns; don't fetch external URLs.
markdown_files = sorted(ROOT.rglob("*.md")) + [ROOT.parent/"00-BAT-DAU.md",
                                               ROOT.parent/"01-NEN-TANG.md"]
link_count, broken, table_errors = 0, [], []
for md in markdown_files:
    raw = md.read_text(encoding="utf-8-sig")
    for target in re.findall(r"!?\[[^\]]*\]\((<[^>]+>|[^)\s]+)\)", raw):
        target = target.strip("<>")
        if re.match(r"^(https?://|mailto:|app:|codex:)", target):
            continue
        target = unquote(target.split("#", 1)[0])
        if not target:
            continue
        link_count += 1
        path = Path(target) if re.match(r"^[A-Za-z]:[/\\]", target) else md.parent/target
        if not path.exists():
            broken.append({"file": md.name, "target": target})
    expected_columns = None
    for line_number, line in enumerate(raw.splitlines(), 1):
        if line.startswith("|") and line.endswith("|"):
            columns = len(re.split(r"(?<!\\)\|", line))-2
            if expected_columns is None:
                expected_columns = columns
            if columns != expected_columns:
                table_errors.append({"file": md.name, "line": line_number,
                                     "columns": columns, "expected": expected_columns})
        else:
            expected_columns = None
assert not broken, broken
assert not table_errors, table_errors
report = {
    "status": "passed",
    "checked_at_utc": datetime.now(timezone.utc).isoformat(),
    "scope": "Educational toy calculations and local Markdown; no model/dataset run",
    "runtime": {"python": sys.version.split()[0], "executable": sys.executable,
                "platform": platform.platform(), "numpy": np.__version__,
                "matplotlib": matplotlib.__version__},
    "numerical_checks_passed": len(CHECKS), "checks": CHECKS,
    "markdown_files_checked": len(markdown_files), "local_links_checked": link_count,
    "broken_local_links": broken, "table_errors": table_errors,
    "figures": ["assets/aliasing-16khz.png", "assets/zero-padding-resolution.png"],
}
(ROOT/"research"/"ket-qua-kiem-tra.json").write_text(
    json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8")
print(json.dumps({k: report[k] for k in ["numerical_checks_passed",
      "markdown_files_checked", "local_links_checked", "broken_local_links", "table_errors"]},
      ensure_ascii=False, indent=2))
