"""Verify pedagogical counterexamples and study-file boundaries, not model results.

Run with the existing lop01-runtime Python (NumPy/Matplotlib already available).
No downloads, training, installation, or edits outside this layer's artifacts.
"""
from pathlib import Path
import hashlib
import json
import math
import re
from datetime import datetime, timezone
import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

LAYER = Path(__file__).resolve().parents[1]
STUDY = LAYER.parent
checks = []


def check(name, passed, details):
    checks.append({"name": name, "passed": bool(passed), "details": details})


def near(name, actual, expected, atol=1e-8):
    check(name, np.allclose(actual, expected, atol=atol, rtol=1e-8),
          {"actual": np.asarray(actual).tolist(), "expected": np.asarray(expected).tolist()})


def finite_diff(fn, x, eps=1e-6):
    x = np.asarray(x, dtype=float)
    out = np.zeros_like(x)
    for i in range(x.size):
        delta = np.zeros_like(x)
        delta.flat[i] = eps
        out.flat[i] = np.asarray((fn(x + delta) - fn(x - delta)) / (2 * eps)).item()
    return out


def erank(h):
    s = np.linalg.svd(h, compute_uv=False)
    # Zero case is a display convention, not entropy of an all-zero distribution.
    if np.sum(s) == 0:
        return 0.0
    p = s[s > 1e-15 * np.max(s)] / np.sum(s)
    return float(np.exp(-np.sum(p * np.log(p))))


def adjacent_visible(mask):
    rows, cols = mask.shape
    hits = []
    for f, t in zip(*np.where(mask)):
        hits.append(any(0 <= ff < rows and 0 <= tt < cols and not mask[ff, tt]
                        for ff, tt in ((f-1, t), (f+1, t), (f, t-1), (f, t+1))))
    return float(np.mean(hits))


# Information preserved under an invertible XOR transform, but not linearly separable.
pairs = [(u, v) for u in (-1, 1) for v in (-1, 1)]
h = np.array([(u*v, v) for u, v in pairs])
near("xor_inverse_recovers_label", h[:, 0]*h[:, 1], [u for u, v in pairs])
class_means = np.array([h[np.array([u for u, v in pairs]) == label].mean(axis=0) for label in (-1, 1)])
near("xor_class_convex_hulls_intersect", class_means, np.zeros((2, 2)))

err = np.array([1.0, -2.0])
near("coordinate_mse_vs_vector_squared_norm", [np.mean(err**2), np.sum(err**2)], [2.5, 5.0])
p = np.array([.6, .3, .1]); q = np.array([.5, .4, .1])
soft_ce = -np.sum(q*np.log(p)); entropy = -np.sum(q*np.log(q))
kl = np.sum(q*np.log(q/p))
near("soft_ce_entropy_kl_identity", [soft_ce, entropy, kl],
     [.9672604429127745, .9433483923290391, .02391205058373512])
near("soft_ce_minus_entropy", soft_ce-entropy, kl)
scores = np.array([2., 0., 0.]); probs = np.exp(scores)/np.exp(scores).sum()
near("infonce_three_candidates", [probs[0], -math.log(probs[0])],
     [.7869860421615985, .2395447662218845])
grad = probs - np.array([1., 0., 0.])
near("infonce_gradient_finite_difference", finite_diff(lambda s: -s[0]+np.log(np.exp(s).sum()), scores), grad)
near("infonce_equal_logits", -math.log(1/3), math.log(3))
near("hubert_masked_mean_ce", -np.mean(np.log([.8, .6])), .3669845875401002)
u = np.array([3., 4.])/5; codes = np.eye(2)
near("random_quantizer_squared_distances", np.sum((codes-u)**2, axis=1), [.8, .4])
near("audiomae_toy_masked_mse", np.mean(np.array([[0., -2.], [0., -4.]])**2), 5.)
near("contextual_target_two_positions", .75*2+.25*10, 4.)
layers = np.array([[0., 2., 4.], [10., 10., 14.]])
norm_layers = (layers-layers.mean(axis=1, keepdims=True))/layers.std(axis=1, keepdims=True)
near("normalized_top_layer_average", norm_layers[:, 2].mean(), 1.319479216882342)

# Detach means cache target at the current step while differentiating the online path.
cached_target = 2.0
near("detached_target_partial_gradient", finite_diff(lambda a: .5*(a-cached_target)**2, np.array([1.])), [-1.])
near("ordinary_shared_parameter_gradient", finite_diff(lambda a: .5*(a-2*a)**2, np.array([1.])), [1.])
near("ema_example_and_halflives", [.9*10+.1*14, math.log(.5)/math.log(.9), math.log(.5)/math.log(.996)],
     [10.4, 6.578813478960585, 172.93999003737392])
z = np.array([[-1., -1.], [1., 1.]])
cov = np.cov(z, rowvar=False, ddof=1)
near("vicreg_covariance_penalty", (np.sum(cov**2)-np.sum(np.diag(cov)**2))/2, 4.)
near("barlow_duplicate_coordinates_offdiag", np.corrcoef(z, rowvar=False), np.ones((2, 2)))

def code_norm_loss(pred, target):
    target = (target-target.mean()) / math.sqrt(float(target.var(ddof=1))+1e-6)
    pred = pred/max(float(np.linalg.norm(pred)), 1e-12)
    target = target/max(float(np.linalg.norm(target)), 1e-12)
    return 2-2*float(np.dot(pred, target))

near("audiojepa_code_loss_vs_raw", [code_norm_loss(np.array([-1., 1.]), np.array([1., 3.])),
     np.sum((np.array([-1., 1.])-np.array([1., 3.]))**2)], [0., 8.])
near("audiojepa_prediction_not_centered", code_norm_loss(np.array([1., 2.]), np.array([1., 3.])), 1.367544467966324)
step_grad = finite_diff(lambda ab: .5*(ab[0]*ab[1]-7.)**2, np.array([2., 3.]))
near("scalar_jepa_context_predictor_gradient", step_grad, [-3., -2.])
new_ab = np.array([2., 3.])-.1*step_grad
teacher = .9*3.5+.1*new_ab[0]
near("scalar_jepa_optimizer_then_ema", [*new_ab, teacher, new_ab.prod(), 2*teacher,
     .5*(new_ab.prod()-2*teacher)**2], [2.3, 3.2, 3.38, 7.36, 6.76, .18])

grid = np.indices((4, 8))
mask_a = (grid[0]+grid[1]) % 2 == 0
mask_b = np.zeros((4, 8), dtype=bool); mask_b[:, 2:6] = True
mask_d = np.zeros((4, 8), dtype=bool); mask_d[1:3, :] = True
near("mask_equal_coverage", [m.mean() for m in (mask_a, mask_b, mask_d)], [.5, .5, .5])
near("mask_local_context_adjacency", [adjacent_visible(m) for m in (mask_a, mask_b, mask_d)], [1., .5, 1.])
near("mean_dominated_cosine", np.dot([100., 1.], [100., -1.])/10001., 9999/10001.)
constant = np.full((4, 2), 7.); centered = constant-constant.mean(axis=0)
near("constant_uncentered_vs_centered_rank", [np.linalg.matrix_rank(constant), np.linalg.matrix_rank(centered)], [1, 0])
tiny = 1e-6*(np.eye(4)-np.ones((4, 4))/4)
near("tiny_scale_effective_rank", erank(tiny), 3.)
near("tiny_scale_sample_variance", tiny.var(axis=0, ddof=1), np.full(4, 2.5e-13), atol=1e-20)
position = np.array([[-1., -1.], [-1., 1.], [1., -1.], [1., 1.]])
tokens = np.tile(position, (4, 1)); pooled = np.tile(position.mean(axis=0), (4, 1))
near("position_only_token_vs_utterance", [erank(tokens), np.linalg.matrix_rank(pooled)], [2., 0.])
near("position_only_pooled_token_variance", tokens.var(axis=0, ddof=1), [16/15, 16/15])
near("rotated_encoder_predictor_error", np.sum((np.array([1., 0.])-np.array([0., 1.]))**2), 2.)
near("weighted_ce_population_shift", [9.6/10.6, math.log(9.6)], [.9056603773584906, 2.2617630984737906])

# Visuals are independent plots of toy geometry; no measured model data.
assets = LAYER/"assets"; assets.mkdir(exist_ok=True)
fig, axes = plt.subplots(1, 3, figsize=(11, 3.1), constrained_layout=True)
for ax, mask, title in zip(axes, (mask_a, mask_b, mask_d),
                          ("A: checkerboard", "B: time block", "D: frequency band")):
    ax.imshow(mask, cmap=matplotlib.colors.ListedColormap(["#e1f3f4", "#d06b3f"]), vmin=0, vmax=1)
    ax.set_title(f"{title}\nvisible neighbor: {adjacent_visible(mask):.0%}", fontsize=11)
    ax.set_xlabel("Time patch"); ax.set_ylabel("Mel patch")
    ax.set_xticks(range(8)); ax.set_yticks(range(4))
    ax.set_xticks(np.arange(-.5, 8, 1), minor=True); ax.set_yticks(np.arange(-.5, 4, 1), minor=True)
    ax.grid(which="minor", color="white", linewidth=1.5); ax.tick_params(which="minor", length=0)
fig.suptitle("Toy only: same 50% mask, different context. Orange = target", fontsize=12)
fig.savefig(assets/"masking-toy.png", dpi=180); plt.close(fig)
fig, axes = plt.subplots(1, 2, figsize=(9, 3.7), constrained_layout=True)
axes[0].scatter(position[:, 0], position[:, 1], s=170, color="#237981")
for i, point in enumerate(position):
    axes[0].annotate(f"position {i+1}", point, xytext=(5, 6), textcoords="offset points", fontsize=9)
axes[0].set_title("Tokens vary by position\nSame four tokens for every utterance")
axes[1].scatter([0], [0], s=200, color="#d06b3f")
axes[1].annotate("all 4 utterances", (0, 0), xytext=(8, 10), textcoords="offset points")
axes[1].set_title("Mean-pooled utterances coincide\nNo between-utterance variation")
for ax in axes:
    ax.set_xlim(-1.5, 1.8); ax.set_ylim(-1.5, 1.5); ax.set_aspect("equal")
    ax.set_xlabel("Coordinate 1"); ax.set_ylabel("Coordinate 2"); ax.grid(alpha=.2)
fig.suptitle("Toy only: token rank 2 does not imply utterance information", fontsize=12)
fig.savefig(assets/"collapse-toy.png", dpi=180); plt.close(fig)

# Baseline is captured before changes; only two pre-existing Markdown files may change.
baseline = json.loads((LAYER/"research/baseline-hashes.json").read_text(encoding="utf-8-sig"))
allowed = {"00-BAT-DAU.md", "03-SSL-VA-JEPA.md"}
changed, forbidden, missing = [], [], []
for name, previous in baseline.items():
    path = STUDY/name
    if not path.exists():
        missing.append(name); continue
    actual = hashlib.sha256(path.read_bytes()).hexdigest()
    if actual != previous:
        changed.append(name)
        if name not in allowed:
            forbidden.append(name)
check("preexisting_markdown_scope", not forbidden and not missing,
      {"baseline_files": len(baseline), "changed": changed, "forbidden_changes": forbidden, "missing": missing})

documents = sorted(LAYER.rglob("*.md")) + [STUDY/n for n in sorted(allowed)]
# The final JSON is emitted below; permit its first-run link after creating a placeholder.
if not (LAYER/"research/validation.json").exists():
    (LAYER/"research/validation.json").write_text("{}", encoding="utf-8")
broken = []; checked_links = 0; formatting = []; replacement = []
link_pattern = re.compile(r"!?\[[^\]\n]*\]\((<[^>]+>|[^)\n]+)\)")
for path in documents:
    content = path.read_text(encoding="utf-8-sig")
    if "\ufffd" in content:
        replacement.append(str(path.relative_to(STUDY)))
    if len(re.findall(r"^```", content, flags=re.M)) % 2:
        formatting.append(f"Unpaired code fence: {path.name}")
    if len(re.findall(r"^~~~", content, flags=re.M)) % 2:
        formatting.append(f"Unpaired tilde code fence: {path.name}")
    if content.count("$$") % 2:
        formatting.append(f"Unpaired display-math delimiters: {path.name}")
    if content.count("<details>") != content.count("</details>"):
        formatting.append(f"Unpaired details: {path.name}")
    table_width = None
    for line_number, line in enumerate(content.splitlines(), start=1):
        if line.startswith("|"):
            width = len(re.findall(r"(?<!\\)\|", line))
            if table_width is None:
                table_width = width
            elif width != table_width:
                formatting.append(f"Table columns: {path.name}:{line_number}")
        else:
            table_width = None
    for match in link_pattern.finditer(content):
        target = match.group(1).strip().strip("<>")
        if re.match(r"(?:https?://|mailto:|#)", target):
            continue
        target = target.split("#", 1)[0]
        if not target:
            continue
        resolved = Path(target) if re.match(r"[A-Za-z]:[/\\]", target) else path.parent/target
        checked_links += 1
        if not resolved.exists():
            broken.append({"file": str(path.relative_to(STUDY)), "target": target})
check("local_markdown_links", not broken, {"links_checked": checked_links, "broken": broken})
check("markdown_structure_and_encoding", not formatting and not replacement,
      {"issues": formatting, "replacement_char_files": replacement})

lessons = sorted(p for p in LAYER.glob("*.md") if re.match(r"(?:0[1-9]|10)-", p.name))
question_counts = {}
for path in lessons:
    content = path.read_text(encoding="utf-8")
    parts = re.split(r"^## .*?(?:Tự kiểm tra|Bài tập).*?$", content, flags=re.M | re.I)
    if len(parts) > 1:
        questions = parts[-1].split("<details>", 1)[0]
        question_counts[path.name] = len(re.findall(r"^\d+\. ", questions, flags=re.M))
check("lesson_answer_and_assessment_structure", len(lessons) == 10 and
      all("<details>" in p.read_text(encoding="utf-8") for p in lessons) and
      all(n >= 4 for n in question_counts.values()) and len(question_counts) == 10,
      {"lessons": len(lessons), "questions_per_lesson": question_counts, "total_questions": sum(question_counts.values())})

result = {"generated_at_utc": datetime.now(timezone.utc).isoformat(),
          "scope": "toy math, geometry, local links and pre-existing Markdown hashes; no model evaluation",
          "python_numpy": np.__version__, "matplotlib": matplotlib.__version__,
          "checks": checks, "passed": sum(c["passed"] for c in checks),
          "total": len(checks), "all_passed": all(c["passed"] for c in checks)}
(LAYER/"research/validation.json").write_text(json.dumps(result, ensure_ascii=False, indent=2), encoding="utf-8")
print(json.dumps({"passed": result["passed"], "total": result["total"],
                  "failed": [c for c in checks if not c["passed"]]}, ensure_ascii=False, indent=2))
raise SystemExit(0 if result["all_passed"] else 1)
