"""Numerical toys, figure, local-link/structure QA and scope hashes only.

Run with the existing tmp/lop01-runtime Python (NumPy/Matplotlib).
No network, model loading, corpus access, notebook execution or training.
"""
from pathlib import Path
import hashlib
import json
import re
from datetime import datetime, timezone
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]
STUDY = ROOT.parent
WORKSPACE = ROOT.parents[2]
REPORT = ROOT/"research/validation.json"
checks = {}
errors = []


def close(name, value, expected, atol=1e-6):
    value = np.asarray(value)
    expected = np.asarray(expected)
    ok = bool(np.allclose(value, expected, atol=atol, rtol=0))
    checks[name] = {"pass": ok, "value": value.tolist(), "expected": expected.tolist()}
    if not ok:
        errors.append(name)


def sigmoid(d):
    return 1 / (1 + np.exp(-np.asarray(d)))


def masked_stats(h, valid, scores=None):
    h = np.asarray(h, dtype=float)
    valid = np.asarray(valid, dtype=bool)
    if not valid.any():
        raise ValueError("All-padding observation")
    s = np.zeros(len(h)) if scores is None else np.asarray(scores, dtype=float)
    weights = np.zeros(len(h))
    e = np.exp(s[valid] - s[valid].max())
    weights[valid] = e / e.sum()
    mean = (weights[:, None] * h).sum(0)
    var = (weights[:, None] * (h - mean) ** 2).sum(0)
    return mean, np.sqrt(var), weights


# Pooling: mask, population rather than sample variance, axes, order.
mu, sd, w = masked_stats([[0], [0], [6], [0]], [1, 1, 1, 0])
close("masked_mean_std", np.r_[mu, sd], [2, np.sqrt(8)])
um, us, _ = masked_stats([[0], [0], [6], [0]], [1, 1, 1, 1])
close("unmasked_mean_std", np.r_[um, us], [1.5, np.sqrt(6.75)])
am, ast, aw = masked_stats([[0], [2], [4]], [1, 1, 1], np.log([.1, .2, .7]))
close("attentive_population_stats", np.r_[am, ast], [3.2, np.sqrt(1.76)])
try:
    masked_stats([[0], [0]], [0, 0])
    errors.append("all_padding_guard")
    checks["all_padding_guard"] = {"pass": False}
except ValueError:
    checks["all_padding_guard"] = {"pass": True, "policy": "raise before normalization"}
seq = np.array([0., 0., 6.])
close("permutation_summaries", [seq.mean(), seq.std(), seq.max()],
      [seq[::-1].mean(), seq[::-1].std(), seq[::-1].max()])
close("sample_population_variance", [np.var([1, 3]), np.var([1, 3], ddof=1)], [1, 2])
close("time_frequency_mapping", np.array([[1, 10], [2, 20]]).mean(1), [5.5, 11])
close("layer_weight_scale", [.9 + .1 * 10, .1 * 10], [1.9, 1])

# Evaluate both Gaussian densities, not only the simplified ratio.
v = np.array([0., 1., 2.])
lpb = -.5 * np.log(2 * np.pi) - .5 * v ** 2
lps = -.5 * np.log(2 * np.pi) - .5 * (v - 2) ** 2
close("gaussian_llr_frames", lpb - lps, [2, 0, -2])
close("gaussian_llr_utterance", (lpb - lps).mean(), 0)

# Loss reductions and gradients. NumPy arithmetic follows API formulas;
# this does not test an installed PyTorch runtime/version.
y = np.array([1, 0, 0])
q = np.array([.8, .6, .2])
d = np.log(q / (1 - q))
cw = np.array([4, 1, 1])
point = np.logaddexp(0, d) - y * d
numerator = (cw * point).sum()
close("weighted_ce_numerator", numerator, 2.032009)
close("weighted_hard_ce_mean", numerator / cw.sum(), .338668)
close("weighted_soft_ce_bce_mean", numerator / len(y), .677336)
close("weighted_ce_gradient", cw / cw.sum() * (q - y), [-2/15, .1, 1/30])
close("singleton_hard_ce_weight_cancels", 9.6 * point[0] / 9.6, point[0])

def weighted_ce(logits):
    return np.sum(cw * (np.logaddexp(0, logits) - y * logits)) / cw.sum()

eps = 1e-5
fd = []
for i in range(3):
    delta = np.eye(3)[i] * eps
    fd.append((weighted_ce(d + delta) - weighted_ce(d - delta)) / (2 * eps))
close("ce_gradient_finite_difference", fd, cw / cw.sum() * (q - y), atol=1e-8)

def focal_positive(logit):
    p = sigmoid(logit)
    return -(1 - p) ** 2 * np.log(p)

fq = np.array([.9, .1])
fdlogit = np.log(fq / (1 - fq))
floss = focal_positive(fdlogit)
fg = (1 - fq) ** 2 * (2 * fq * np.log(fq) - (1 - fq))
close("focal_values", floss, [.001053605156578263, 1.8650939253251773])
close("focal_derivative", fg, [-.00289648928184087, -1.10201878506643])
close("focal_gradient_finite_difference",
      (focal_positive(fdlogit + eps) - focal_positive(fdlogit - eps)) / (2 * eps), fg,
      atol=1e-8)
close("population_prior_weights", [(9*.1)/(9*.1+.9), (9*.5)/(9*.5+.5)], [.5, .9])

# LoRA forward/rank and initial data-loss gradient, no fitting.
a = np.array([[1., -1.]])
bl = np.array([[.5], [1.]])
x = np.array([3., 1.])
dw = 2 * bl @ a
close("lora_forward", (np.eye(2) + dw) @ x, [5, 5])
close("lora_rank", np.linalg.matrix_rank(dw), 1)
close("lora_full_effective_rank_counterexample", np.linalg.matrix_rank(np.diag([2, 1])), 2)
bl0 = np.zeros((2, 1))
g = np.array([1., 2.])
analytic_b = 2 * g[:, None] @ (a @ x)[None, :]
numeric_b = np.zeros_like(bl0)
for i in range(2):
    delta = np.zeros_like(bl0)
    delta[i, 0] = eps
    numeric_b[i, 0] = (g @ (2*(bl0+delta) @ a @ x)
                       - g @ (2*(bl0-delta) @ a @ x)) / (2*eps)
close("lora_B_initial_gradient", numeric_b, analytic_b)
close("lora_A_initial_gradient", 2 * (bl0.T @ g)[:, None] @ x[None, :], [[0, 0]])
close("lora_parameter_count", 8*(768+768), 12288)

# Label/coverage/order, source-support and domain objectives.
overlap = lambda lo, hi: max(0, min(hi, 3.6) - max(lo, 3.2))
close("partial_crop_overlap_seconds", [overlap(0, 2.56), overlap(1.44, 4)], [0, .4])
close("partial_fake_duration_fraction", overlap(1.44, 4)/2.56, .15625)
raw = np.array([1, 0, 0, 0])
filt = lambda z: np.convolve(z, [1, 1])[:len(z)]
close("filter_then_crop", filt(raw)[1:4], [1, 0, 0])
close("crop_then_filter", filt(raw[1:4]), [0, 0, 0])
close("mask_band_union", 16/256 + 16/128 - (16/256)*(16/128), .1796875)
close("source_shortcut_accuracy", [180/200, 100/200, 1780/1800], [.9, .5, .988888888889])
sizes = np.array([900., 100.])
domain_weights = np.minimum(sizes, 100) ** (1/2)
prob = domain_weights/domain_weights.sum()
close("doss_domain_probability", prob, [.5, .5])
close("doss_per_file_exposure_ratio", (prob[1]/sizes[1])/(prob[0]/sizes[0]), 9)
risks = np.array([[.1, .9], [.2, .3]])
close("erm_worst_group_tradeoff", np.c_[risks @ [.9, .1], risks.max(1)], [[.18, .9], [.21, .3]])
dro_q = np.exp([.1, .9]); dro_q /= dro_q.sum()
close("group_dro_weight_update", dro_q, [.310025518872, .689974481128])

# Score/ranking, windows, latency arithmetic and diagnostics toys.
zb, zs = np.array([7, 4]), np.array([6, 0])
close("two_logit_posterior", sigmoid(zb-zs), [.731058578630, .982013790038])
close("fusion_ranking_example", .5*np.array([4, 1])+.5*np.array([-8, 2]), [-2, 1.5])
close("window_logits_vs_probability", [sigmoid(np.mean([10, -2])), sigmoid([10, -2]).mean()],
      [.982013790038, .559578762077])
close("spoof_window_aggregations", [np.mean([.1, .1, .9, .1]), .9, np.mean([.9, .1])], [.3, .9, .5])
close("naive_final_end_tail_start", [1.28+2.56, 4-2.56], [3.84, 1.44])
close("independent_max_false_alarm", 1-(1-.01)**100, .633967658727)
close("sigmoid_gradient_saturation", 10*sigmoid(10)*(1-sigmoid(10)), .000453958077)
c = np.array([0, 0, 1, 1]); n = np.array([0, 1, 0, 1])
close("probe_without_reliance", 2*c-1, 2*c-1 + 0*n)
acc = lambda scores: ((scores > 0) == c).mean()
close("inference_vs_retrain_ablation", [acc(c+c-1), acc(c-1), acc(2*c-1)], [1, .5, 1])
close("token_geometry", [256//8, 128//32, (256//8)*(128//32), 2.56*32000], [32, 4, 128, 81920])
p_asp = 768*128 + 128 + 128 + 1
p_head = 2*1536 + 1536*256 + 256 + 256*2 + 2
close("backend_parameter_count", [p_asp, p_head, 13+p_asp+p_head], [98561, 397058, 495632])
close("readout_cache_bytes_MiB", [13*128*768*4, 13*128*768*4/2**20], [5111808, 4.875])

# A standalone explanatory figure, using the exact example values.
plt.rcParams.update({"font.family": "DejaVu Sans", "font.size": 11})
fig, ax = plt.subplots(1, 2, figsize=(12, 4.5), constrained_layout=True)
fig.suptitle("Ví dụ pooling — số tự biên soạn, không phải model outputs", fontsize=15)
positions = np.arange(3)
ax[0].bar(positions-.17, [2, np.sqrt(8), 6], .34, label="Valid [0, 0, 6]", color="#2563eb")
ax[0].bar(positions+.17, [1.5, np.sqrt(6.75), 6], .34, label="Tính cả padding 0", color="#ea580c")
ax[0].set_xticks(positions, ["Mean", "Population std", "Max"])
ax[0].set_ylim(0, 7); ax[0].set_ylabel("Giá trị feature toy")
ax[0].set_title("Padding đổi summary dù có giá trị 0")
ax[0].legend(loc="upper left", fontsize=9); ax[0].grid(axis="y", alpha=.2)
ax[1].plot([1, 2, 3], [0, 0, 6], "o-", color="#2563eb", label="[0, 0, 6]")
ax[1].plot([1, 2, 3], [6, 0, 0], "s--", color="#7c3aed", label="[6, 0, 0]")
ax[1].set_xticks([1, 2, 3]); ax[1].set_ylim(-.3, 7)
ax[1].set_xlabel("Chỉ số trong cùng tập feature vectors")
ax[1].set_title("Đổi thứ tự, mean/std/max vẫn giống")
ax[1].legend(); ax[1].grid(alpha=.2)
ax[1].text(1.05, 3.1, "Cả hai: mean = 2\nstd = √8 ≈ 2,828\nmax = 6", fontsize=12,
           bbox={"facecolor": "white", "edgecolor": "#e2e8f0", "alpha": .95, "pad": 6})
(ROOT/"assets").mkdir(exist_ok=True)
fig.savefig(ROOT/"assets/pooling-counterexamples.png", dpi=160)
plt.close(fig)

# Markdown QA covers delivered material plus the two allowed existing files.
files = sorted(ROOT.rglob("*.md")) + [STUDY/"00-BAT-DAU.md", STUDY/"04-KY-THUAT-DETECTOR.md"]
link_errors = []; structure_errors = []; link_count = 0
for path in files:
    content = path.read_text(encoding="utf-8-sig")
    if "\ufffd" in content or any(ord(c)<32 and c not in "\n\r\t" for c in content):
        structure_errors.append(f"{path.name}: encoding/control character")
    if sum(line.startswith("```") for line in content.splitlines()) % 2:
        structure_errors.append(f"{path.name}: unclosed code fence")
    if content.count("$$") % 2:
        structure_errors.append(f"{path.name}: unclosed display math")
    if content.count("<details>") != content.count("</details>"):
        structure_errors.append(f"{path.name}: unclosed details")
    in_code = False; table_width = None
    for ln, line in enumerate(content.splitlines(), 1):
        if line.startswith("```"):
            in_code = not in_code
        if in_code:
            continue
        if line.startswith("|") and line.endswith("|"):
            count = len(re.split(r"(?<!\\)\|", line))-2
            if table_width is None:
                table_width = count
            elif table_width != count:
                structure_errors.append(f"{path.name}:{ln}: table width {count} != {table_width}")
        else:
            table_width = None
    for match in re.finditer(r"!?\[[^\]\n]*\]\((<[^>]*>|[^)\n]*)\)", content):
        target = match.group(1).strip().strip("<>")
        if re.match(r"https?://|mailto:|codex:", target) or target.startswith("#"):
            continue
        target = unquote(target.split("#", 1)[0])
        destination = Path(target) if Path(target).is_absolute() else path.parent/target
        link_count += 1
        # The report links point to the output written below in this same run.
        if not destination.exists() and destination.resolve() != REPORT:
            line = content[:match.start()].count("\n") + 1
            link_errors.append(f"{path.relative_to(WORKSPACE)}:{line}: {target}")

# Verify all pre-existing baseline files. Only the two specified files may change.
baseline = json.loads((ROOT/"research/baseline-hashes.json").read_text(encoding="utf-8-sig"))
allowed = {STUDY/"00-BAT-DAU.md", STUDY/"04-KY-THUAT-DETECTOR.md"}
changed = []; missing = []; protected_changed = []
for key, expected in baseline.items():
    p = Path(key) if Path(key).is_absolute() else WORKSPACE/key
    if not p.exists():
        missing.append(key)
        continue
    current = hashlib.sha256(p.read_bytes()).hexdigest()
    if current != expected:
        changed.append(key)
        if p.resolve() not in {a.resolve() for a in allowed}:
            protected_changed.append(key)
errors += link_errors + structure_errors + protected_changed + missing
report = {
    "generated_at": datetime.now(timezone.utc).isoformat(),
    "scope": "teaching toys and files; no model/corpus/notebook execution",
    "runtime": {"numpy": np.__version__, "matplotlib": matplotlib.__version__},
    "numeric_checks": checks,
    "markdown": {"files": len(files), "local_links": link_count,
                 "link_errors": link_errors, "structure_errors": structure_errors},
    "baseline": {"files": len(baseline), "changed": changed,
                 "protected_changed": protected_changed, "missing": missing},
    "pass": not errors,
    "errors": errors,
}
REPORT.write_text(json.dumps(report, ensure_ascii=False, indent=2)+"\n", encoding="utf-8")
print(json.dumps({"pass": report["pass"], "numeric_cases": len(checks),
                  "markdown": report["markdown"], "baseline": report["baseline"]}, ensure_ascii=True))
raise SystemExit(0 if report["pass"] else 1)
