#!/usr/bin/env python3
"""Per-submission access-distance plots.

For each `../../submissions/*.ir` file, walk the IR and collect the v0 read
distance ⌈√addr⌉ for every operand read (binary-op sources, copy src,
final output reads). Save a 2-panel PNG into this directory:

    [left]  histogram: how many reads happen at each distance
    [right] CDF:       cumulative cost share vs distance
                       (so you can see how much of the total cost
                        comes from the long-distance tail)

Also emits a single combined CDF (`combined_cdf.png`) overlaying the 16×16
record-table submissions on one axis.

Run:
    python3 matmul/doc/access_distance/plot_access_distance.py
"""
from __future__ import annotations

import sys
from pathlib import Path
from typing import List

import matplotlib.pyplot as plt
import numpy as np

HERE = Path(__file__).parent
MATMUL = HERE.parent.parent
SUBMISSIONS = MATMUL / "submissions"

# Make `import matmul` resolve to the parent package's matmul.py.
sys.path.insert(0, str(MATMUL))
import matmul as mm  # noqa: E402

COMBINED = [
    "baseline_16x16.ir",
    "recursive_16x16.ir",
    "tiled_16x16.ir",
    "tiled_16x16_opt1.ir",
    "hierarchical_16x16.ir",
    "sa_cache_16x16.ir",
    "redirect_16x16.ir",
    "sc_outputs_16x16.ir",
    "dead_input_outputs_packed_16x16.ir",
    "aliased_16x16.ir",
    "colmajor_fused_16x16.ir",
    "output_repacked_tail_16x16.ir",
    "output_repacked_tail_deferred_value_colored_live_b_16x16.ir",
    "output_repacked_tail_deferred_value_colored_live_b_tiny_a_endpoint_16x16.ir",
    "weighted_lifetime_copyelim_66707.ir",
    "macro_b_staging_66633.ir",
    "cheap_capture_66524.ir",
    "motif_bundle_66400.ir",
    "best_66300.ir",
]


def collect_read_distances(ir: str) -> List[int]:
    """Replay the IR exactly the way matmul._simulate does, recording
    ⌈√addr⌉ for every operand read."""
    input_addrs, ops, output_addrs = mm._parse(ir)
    distances: List[int] = []
    for op, oprs in ops:
        if op == "copy":
            _, src = oprs
            distances.append(mm._cost(src))
            continue
        # add/sub/mul
        if len(oprs) == 3:
            _, s1, s2 = oprs
        else:
            dest, s2 = oprs
            s1 = dest
        distances.append(mm._cost(s1))
        distances.append(mm._cost(s2))
    for a in output_addrs:
        distances.append(mm._cost(a))
    return distances


def plot_one(ir_path: Path, out_path: Path) -> int:
    distances = np.array(collect_read_distances(ir_path.read_text()))
    total_cost = int(distances.sum())
    n_reads = len(distances)

    fig, (ax_h, ax_c) = plt.subplots(1, 2, figsize=(11, 4))

    edges = np.arange(distances.min(), distances.max() + 2) - 0.5
    ax_h.hist(distances, bins=edges, color="#3b78b4",
              edgecolor="white", linewidth=0.5)
    ax_h.set_xlabel("distance")
    ax_h.set_ylabel("count")
    ax_h.set_title(f"{ir_path.name}\n{n_reads:,} reads, total cost {total_cost:,}")
    ax_h.grid(axis="y", alpha=0.3)

    sorted_d = np.sort(distances)
    cumulative_count = np.arange(1, len(sorted_d) + 1)
    ax_c.plot(sorted_d, cumulative_count, color="#cc4c4c", linewidth=1.5)
    ax_c.set_xlabel("distance")
    ax_c.set_ylabel("count")
    ax_c.set_title("CDF (reads at distance ≤ x)")
    ax_c.grid(alpha=0.3)

    plt.tight_layout()
    plt.savefig(out_path, dpi=120)
    plt.close(fig)
    return total_cost


def plot_combined_cdf(ir_names: List[str], out_path: Path) -> None:
    fig, ax = plt.subplots(figsize=(9, 6))
    cmap = plt.get_cmap("turbo")
    rows = []
    for name in ir_names:
        ir_path = SUBMISSIONS / name
        distances = np.sort(collect_read_distances(ir_path.read_text()))
        rows.append((name, distances, int(distances.sum())))
    # Sort legend by cost (cheapest first reads bottom→top in stacked order).
    rows.sort(key=lambda r: r[2])
    n = len(rows)
    for i, (name, distances, total) in enumerate(rows):
        color = cmap(0.05 + 0.9 * i / max(n - 1, 1))
        cumulative = np.arange(1, len(distances) + 1)
        ax.plot(distances, cumulative, label=f"{name}  (cost {total:,})",
                color=color, linewidth=1.5)
    ax.set_xlabel("distance")
    ax.set_ylabel("count")
    ax.set_title(f"CDF — reads at distance ≤ x  ({n} submissions)")
    ax.grid(alpha=0.3)
    ax.legend(loc="lower right", fontsize=8)
    plt.tight_layout()
    plt.savefig(out_path, dpi=120)
    plt.close(fig)


def main() -> None:
    print(f"{'submission':<40}{'reads':>10}{'total_cost':>14}  out")
    print("-" * 86)
    for ir_path in sorted(SUBMISSIONS.glob("*.ir")):
        if ir_path.name.endswith(".raw.ir"):
            continue
        out_path = HERE / (ir_path.stem + ".png")
        total_cost = plot_one(ir_path, out_path)
        n_reads = len(collect_read_distances(ir_path.read_text()))
        print(f"{ir_path.name:<40}{n_reads:>10,}{total_cost:>14,}  "
              f"{out_path.name}")

    combined_path = HERE / "combined_cdf.png"
    plot_combined_cdf(COMBINED, combined_path)
    print(f"\nCombined CDF: {combined_path.name}  ({len(COMBINED)} submissions)")


if __name__ == "__main__":
    main()
