"""Audit the frozen four-worker AWC sweep without clipping and export all measurements."""

import csv
import hashlib
import itertools
import json
import math
import statistics
import subprocess
from collections import defaultdict
from pathlib import Path

import yaml

ROOT = Path(__file__).resolve().parent
REPO = ROOT.parent.parent.parent
ARCHIVE = REPO / "runs/packed4-awc-noclip-20m-batch128k-20260912"


def sha(path):
    return hashlib.sha256(path.read_bytes()).hexdigest()


def write_csv(name, rows):
    with (ROOT / name).open("w") as handle:
        writer = csv.DictWriter(handle, fieldnames=list(rows[0]), lineterminator="\n")
        writer.writeheader()
        writer.writerows(rows)


def main():
    manifest = json.loads((ARCHIVE / "manifest.json").read_text())
    for name, expected in manifest["hashes"].items():
        assert sha(ARCHIVE / name) == expected, name
    package = ARCHIVE / "source/tiny_llm"
    digest = hashlib.sha256()
    for path in sorted(package.rglob("*.py")):
        digest.update(path.relative_to(package).as_posix().encode())
        digest.update(path.read_bytes())
    source_hash = digest.hexdigest()
    job = json.loads((ARCHIVE / "submission.json").read_text())["sweep"]
    accounting = subprocess.check_output(
        [
            "sacct",
            "-X",
            "-j",
            job,
            "--noheader",
            "--parsable2",
            "--format=JobID,State,ExitCode,Start,End",
        ],
        text=True,
    )
    jobs = {line.split("|")[0]: line.split("|")[1:3] for line in accounting.splitlines()}
    assert len(jobs) == 240
    assert all(value == ["COMPLETED", "0:0"] for value in jobs.values())
    (ROOT / "slurm-accounting.txt").write_text(accounting)
    rows, groups, environments, artifacts = [], defaultdict(list), [], {}
    for row in manifest["runs"]:
        directory = Path(row["directory"])
        expected = yaml.safe_load(Path(row["config"]).read_text())
        actual = yaml.safe_load((directory / "resolved.yaml").read_text())
        assert actual == expected, row["run_id"]
        assert actual["optimizer"]["grad_clip"] is None
        assert tuple(actual["optimizer"][k] for k in ("lr", "beta1", "beta2")) == tuple(
            row[k] for k in ("lr", "beta1", "beta2")
        )
        assert actual["runtime"]["seed"] == row["seed"]
        assert actual["training"]["checkpoint_policy"] == "none"
        assert actual["training"]["batch_tokens"] == 131072
        assert actual["training"]["micro_batch_size"] == 32
        assert actual["decentralized"] == dict(
            num_models=4, topology="one_peer_exponential", scheme="awc", adaptive_consensus=None
        )
        assert not list(directory.rglob("*.pt"))
        assert not list(directory.rglob("*.safetensors"))
        result = json.loads((directory / "result.json").read_text())
        assert result == json.loads((directory / "status.json").read_text())
        assert result["status"] == "complete"
        for key, value in dict(
            parameters=20403520,
            total_parameters=81614080,
            num_models=4,
            topology="one_peer_exponential",
            scheme="awc",
            tokens=408068096,
            tokens_per_model=102017024,
            step=3119,
            epochs=40,
        ).items():
            assert result[key] == value, (row["run_id"], key)
        val = result["final_validation"]
        assert val["split"] == "full" and val["validation_complete"]
        assert val["tokens"] == 197411295 and math.isfinite(val["loss"])
        env = json.loads((directory / "environment.json").read_text())
        assert env["source_hash"] == source_hash
        assert env["scheme"] == "awc" and env["num_models"] == 4
        assert env["loader"]["cache_identity"] == manifest["cache_identity"]
        assert env["loader"]["seed"] == row["seed"]
        assert jobs[f"{job}_{row['index']}"] == ["COMPLETED", "0:0"]
        environments.append(
            {k: env[k] for k in ["python", "versions", "cuda", "gpu", "cpu_threads"]}
        )
        for name in ["result.json", "resolved.yaml", "environment.json"]:
            artifacts[str((directory / name).relative_to(REPO))] = sha(directory / name)
        output = dict(
            index=row["index"],
            run_id=row["run_id"],
            job_id=f"{job}_{row['index']}",
            lr=row["lr"],
            beta1=row["beta1"],
            beta2=row["beta2"],
            seed=row["seed"],
            loss=val["loss"],
            seconds=result["seconds_this_session"],
            training_seconds=result["training_seconds"],
            peak_memory_bytes=result["peak_memory_bytes"],
            directory=str(directory.relative_to(REPO)),
        )
        rows.append(output)
        groups[(row["lr"], row["beta1"], row["beta2"])].append(output)
    assert all(env == environments[0] for env in environments)
    assert len(rows) == 240 and len(groups) == 80
    assert set(groups) == set(
        itertools.product(manifest["lrs"], manifest["beta1s"], manifest["beta2s"])
    )
    summary = []
    for (lr, beta1, beta2), members in groups.items():
        assert sorted(r["seed"] for r in members) == [42, 43, 44]
        losses = [r["loss"] for r in members]
        summary.append(
            dict(
                lr=lr,
                beta1=beta1,
                beta2=beta2,
                seeds=3,
                mean_loss=statistics.mean(losses),
                std_loss=statistics.stdev(losses),
            )
        )
    summary.sort(key=lambda r: (r["mean_loss"], r["lr"], r["beta1"], r["beta2"]))
    reference_path = REPO / "doc/data/recipe_sweep_packed4_20m_128k/results.json"
    four = json.loads(reference_path.read_text())
    for key in [
        "parameters",
        "num_models",
        "topology",
        "batch_tokens",
        "micro_batch_size",
        "training_tokens",
        "optimizer_steps",
        "epochs",
        "validation_tokens",
        "cache_identity",
    ]:
        assert manifest[key] == four["protocol"][key], key
    for campaign in four["campaigns"].values():
        assert all(
            campaign["environment"][key] == environments[0][key]
            for key in ["python", "versions", "cuda", "gpu"]
        )
    four_groups = {(r["lr"], r["beta1"], r["beta2"]): r for r in four["groups"]}
    baseline_runs = {(r["lr"], r["beta1"], r["beta2"], r["seed"]): r for r in four["runs"]}
    matched_runs = []
    for row in rows:
        key = tuple(row[k] for k in ("lr", "beta1", "beta2", "seed"))
        if key not in baseline_runs:
            continue
        baseline = baseline_runs[key]
        directory = REPO / baseline["archive"] / baseline["directory"]
        archived_only = not (directory / "result.json").exists()
        assert baseline["result"]["final_validation"]["loss"] == baseline["loss"]
        if not archived_only:
            baseline_result = json.loads((directory / "result.json").read_text())
            assert baseline_result["final_validation"]["loss"] == baseline["loss"]
            clipped_config = yaml.safe_load((directory / "resolved.yaml").read_text())
            unclipped_config = yaml.safe_load(
                (REPO / row["directory"] / "resolved.yaml").read_text()
            )
            assert clipped_config["optimizer"]["grad_clip"] == 1.0
            for config in (clipped_config, unclipped_config):
                config["optimizer"].pop("grad_clip")
                config["training"].pop("checkpoint_policy")
                config["runtime"].pop("output_dir")
                config["decentralized"].setdefault("scheme", "awc")
            assert clipped_config == unclipped_config, key
            for name in ["result.json", "resolved.yaml"]:
                artifacts[str((directory / name).relative_to(REPO))] = sha(directory / name)
        matched_runs.append(
            dict(
                lr=row["lr"],
                beta1=row["beta1"],
                beta2=row["beta2"],
                seed=row["seed"],
                unclipped_loss=row["loss"],
                clipped_loss=baseline["loss"],
                clipped_evidence="published_per_run_record"
                if archived_only
                else "original_files_reaudited",
                difference=row["loss"] - baseline["loss"],
                unclipped_directory=row["directory"],
                clipped_directory=str(directory.relative_to(REPO)),
            )
        )
    matched = []
    for r in summary:
        key = (r["lr"], r["beta1"], r["beta2"])
        if key in four_groups:
            f = four_groups[key]
            matched.append(
                dict(
                    paired_std=statistics.stdev(
                        x["difference"]
                        for x in matched_runs
                        if (x["lr"], x["beta1"], x["beta2"]) == key
                    ),
                    lr=r["lr"],
                    beta1=r["beta1"],
                    beta2=r["beta2"],
                    unclipped_mean=r["mean_loss"],
                    unclipped_std=r["std_loss"],
                    clipped_mean=f["mean_loss"],
                    clipped_std=f["std_loss"],
                    difference_unclipped_minus_clipped=r["mean_loss"] - f["mean_loss"],
                )
            )
    assert len(matched) == 28 and len(matched_runs) == 84
    for row in matched:
        pairs = [x for x in matched_runs if all(x[k] == row[k] for k in ("lr", "beta1", "beta2"))]
        assert sorted(x["seed"] for x in pairs) == [42, 43, 44]
        assert abs(statistics.mean(x["clipped_loss"] for x in pairs) - row["clipped_mean"]) < 1e-12
        assert abs(statistics.stdev(x["clipped_loss"] for x in pairs) - row["clipped_std"]) < 1e-12
        assert (
            abs(
                statistics.mean(x["difference"] for x in pairs)
                - row["difference_unclipped_minus_clipped"]
            )
            < 1e-12
        )
    write_csv("matched_runs.csv", matched_runs)
    write_csv("runs.csv", rows)
    write_csv("summary.csv", summary)
    write_csv(
        "matched_clipped.csv", sorted(matched, key=lambda r: (r["lr"], r["beta1"], r["beta2"]))
    )
    results = dict(
        schema_version=1,
        measured_dates=sorted(
            {line.split("|")[i].split("T")[0] for line in accounting.splitlines() for i in [3, 4]}
        ),
        run_count=240,
        group_count=80,
        selection="Mean final full-validation loss; ties: LR, beta1, beta2",
        winner=summary[0],
        groups=summary,
        runs=rows,
        protocol={k: v for k, v in manifest.items() if k not in ["hashes", "runs"]},
        archive=str(ARCHIVE.relative_to(REPO)),
        source_hash=source_hash,
        source_and_config_hashes=manifest["hashes"],
        artifact_hashes=artifacts,
        environment=environments[0],
        job_id=job,
        audit=dict(successful_jobs=240, complete_three_seed_groups=80, checkpoint_files=0),
        clipped_reference=dict(
            path=str(reference_path.relative_to(REPO)),
            sha256=sha(reference_path),
            winner=four["winner"],
            run_count=four["run_count"],
        ),
        matched_clipped=matched,
        matched_runs=matched_runs,
    )
    (ROOT / "results.json").write_text(json.dumps(results, indent=2, allow_nan=False) + "\n")
    print(
        json.dumps(dict(winner=summary[0], top5=summary[:5], environment=environments[0]), indent=2)
    )


if __name__ == "__main__":
    main()
