"""CUDA graphs of the training step's forward and backward: one captured pair per StepGraphKey.

What it replaces: running the compiled model eagerly every step, which launches ~1-2 thousand kernels
from Python (inductor's generated code, the Triton/fp8 kernels, FA3) and leaves the GPU waiting on the
host between many of them.

Why it is faster: a replay launches the whole forward (or backward) as one graph, so the host cost per
step drops to a handful of copies and two replays, and the host runs ahead of the device instead of
behind it. Record #360 measured ~3.0 s over the run for this and the other graphs.

What is captured, per key (capture_plan.StepGraphKey: the shapes and the Python values the compiled model
specializes on):
  forward  the model's training loss: (per-token loss).sum()
  backward torch.autograd.grad of it w.r.t. every parameter and the n-gram gradient sink
Between the two replays the host runs the eager row-pull id exchange (perf/row_prefetch.py), so the
all_to_all hides under the backward exactly as it does without graphs.

What must stay address-stable (a replay reads and writes the addresses the capture saw):
  - the per-key static inputs this module owns: inputs, targets, cu_seqlens, n-gram cache slots, the
    MTP and prefix weights; each step copies its batch into them before the forward replay
  - every parameter's .data, every model buffer (fp8 caches and scales, YaRN tables, the n-gram row cache,
    the prefix table), the n-gram gradient sink leaves (NgramTable.grad_sink), value_embeds' gradient
    buffer (ValueEmbedPull.grad_accum), the sampled-softmax candidate buffers (SampledSoftmax.loss_inputs)
baked_addresses.py censuses all of them before and after the post-warmup reset.
Python values the capture froze: the key, and the YaRN attention scales (re-checked at every replay).

Gradients: the captured backward writes each gradient into its own pool buffer, rewritten by the next
replay, and no AccumulateGrad runs, so p.grad is set here the way loss.backward() would: ANVIL banks take
the graph's buffer directly (their reduce-scatter reads it before the next step's backward); Adam
parameters, whose gradient accumulates over two steps, go through one persistent buffer per parameter
(copy on the first step, add on the second). The n-gram sink gradient is held by the table until its
update event, up to MAX_CYCLE_STEPS steps, so each non-event step's copy goes to a private ring slot.

Capture protocol (capture_plan.py): the first warmup visit of a key captures (both graphs), then
replays the forward; the second visit replays and compares the loss with an eager forward on the same
buffers. seal() then makes a missing capture fatal. Captures run only in the untimed warmup.

Provenance: record #360 (ANVIL2): `_CGraphRunner` (fwd/bwd, `_acc`, `_sink_flat` ring, `_selfcheck_now`,
`seal`). Deviations: no host run-ahead throttle (every pinned ring this run rewrites is event-guarded);
the value_embeds gradient buffer needs no ring (it is one persistent accumulator).
"""
import dataclasses
from dataclasses import dataclass

import torch
from torch import Tensor, nn

from track_1_short.data import Batch
from track_1_short.model.gpt import ForwardScheduleConfig
from track_1_short.ngram_table import MAX_CYCLE_STEPS, NGRAM_DIM, NgramTable
from track_1_short.optim.anvil import AnvilAndAdam
from track_1_short.perf.cuda_graphs.capture_plan import StepGraphKey
from track_1_short.perf.cuda_graphs.capture_support import WARM_ITERATIONS
from track_1_short.perf.value_embed_pull import ValueEmbedPull

# The self-check compares scalar losses of two runs of the same kernels; atomics in the loss kernels
# make them not bitwise equal (record #360's tolerance).
SELF_CHECK_ABS, SELF_CHECK_REL = 1e-3, 5e-4


@dataclass(slots=True)
class StaticInputs:
    """One key's input buffers: the graphs read these addresses, each step copies its batch in."""
    inputs: Tensor
    targets: Tensor
    cum_seqlens: Tensor
    ngram_slots: Tensor
    mtp_weights: Tensor
    prefix_weight: Tensor

    def buffers(self):
        return (self.inputs, self.targets, self.cum_seqlens, self.ngram_slots, self.mtp_weights, self.prefix_weight)


def step_inputs(batch: Batch, ngram_slots: Tensor, cfg: ForwardScheduleConfig):
    """This step's tensors, in StaticInputs order."""
    return (batch.inputs, batch.targets, batch.cum_seqlens, ngram_slots, cfg.mtp_weights, cfg.prefix_weight)


@dataclass(slots=True)
class CapturedStep:
    key: StepGraphKey
    static: StaticInputs
    cfg: ForwardScheduleConfig          # the key's config, reading static.mtp_weights / prefix_weight
    ngram_sink: Tensor                  # the zeros leaf the forward adds to the n-gram rows
    attn_scales: tuple[float, ...]      # YaRN attention scales the compiled forward baked
    forward: torch.cuda.CUDAGraph
    backward: torch.cuda.CUDAGraph
    loss: Tensor                        # forward output (pool memory)
    anvil_grads: list[tuple[nn.Parameter, Tensor]]
    adam_grads: list[tuple[nn.Parameter, Tensor, Tensor]]  # (param, graph output, persistent accumulator)
    ngram_grad: Tensor                  # backward output: the sink's gradient [2T, NGRAM_DIM]


def live_key(batch: Batch, cfg: ForwardScheduleConfig) -> StepGraphKey:
    return StepGraphKey(
        tokens=batch.inputs.numel(),
        cu_seqlens_rows=batch.cum_seqlens.numel(),
        mtp_targets=cfg.mtp_weights.numel(),
        candidates=0 if cfg.sampled_loss is None else cfg.sampled_loss.rows.shape[0],
        ws_short=cfg.ws_short,
        ws_long=cfg.ws_long,
        max_seq_len=cfg.train_max_seq_len,
    )


class StepGraphs:
    """The training step's forward and backward graphs. Created in main(); used by train_step for every
    training step, warmup and timed: forward(), then backward(), then hold_ngram_grad()."""

    def __init__(self, model: nn.Module, optimizer: AnvilAndAdam, step_keys: list[StepGraphKey],
                 ngram_table: NgramTable, value_embeds: ValueEmbedPull, max_step_tokens: int,
                 device: torch.device):
        self.model = model  # the compiled GPT
        self.params = list(optimizer.param_cfgs)
        self.anvil_params = {p for p, cfg in optimizer.param_cfgs.items() if cfg.optim == "anvil"}
        self.step_keys = step_keys  # step -> key, for every timed training step
        self.ngram_table = ngram_table
        self.value_embeds = value_embeds
        self.stream = torch.cuda.Stream(device)
        # One pool for every key's two graphs. Replays never interleave across keys (forward and
        # backward of one key are adjacent), so blocks freed after one capture can serve another's.
        self.pool = torch.cuda.graph_pool_handle()
        self.captured: dict[StepGraphKey, CapturedStep] = {}
        self.unchecked: set[StepGraphKey] = set()
        self.sealed = False
        self.adam_accumulators: dict[nn.Parameter, Tensor] = {}
        # n-gram sink gradients held until the update event: slot i holds pending entry i of the cycle.
        # Flat, sized for the largest step; each key views its own [2T, NGRAM_DIM] prefix.
        self.ngram_grad_ring = [torch.empty(2 * max_step_tokens * NGRAM_DIM, dtype=torch.bfloat16, device=device)
                                for _ in range(MAX_CYCLE_STEPS)]
        self.current: CapturedStep | None = None

    def _attn_scales(self) -> tuple[float, ...]:
        return tuple(yarn.attn_scale for yarn in (self.model.yarn, self.model.yarn_paired_head, self.model.yarn_wide))

    def _loss(self, cs: CapturedStep) -> Tensor:
        """The trainer's forward call on the key's static buffers (the single definition the capture,
        the warm iterations and the self-check all run)."""
        s = cs.static
        return self.model(s.inputs, s.targets, s.cum_seqlens, s.ngram_slots, cs.cfg, ngram_sink=cs.ngram_sink,
                          value_embed_grad=self.value_embeds.grad_accum).sum()

    # ---- per step ----

    def forward(self, step: int, batch: Batch, ngram_slots: Tensor, cfg: ForwardScheduleConfig):
        """Copy this step's batch into its key's buffers and replay the forward (capturing it first on
        the key's first warmup visit)."""
        key = live_key(batch, cfg)
        # The replay cannot re-run dynamo's guards: the live shapes and specialized values must be the
        # ones the schedule says, which is what the capture plan made sure warmup captured.
        assert key == self.step_keys[step], f"step {step}: live key {key} != schedule key {self.step_keys[step]}"
        cs = self.captured.get(key)
        fresh = cs is None
        if fresh:
            # Captures are warmup-only: after seal() a missing graph is fatal, never a silent eager step.
            assert not self.sealed, f"step {step}: no graph for {key}"
            cs = self._capture(key, batch, ngram_slots, cfg)
        # The attention scale is a Python float the compiled forward baked; it follows the window
        # schedule, so it must still be the captured one.
        assert self._attn_scales() == cs.attn_scales, f"step {step}: YaRN attention scale moved"
        # The candidate buffers of one P are allocated once, so the same object means the same addresses.
        assert cfg.sampled_loss is cs.cfg.sampled_loss
        for dst, src in zip(cs.static.buffers(), step_inputs(batch, ngram_slots, cfg)):
            dst.copy_(src)
        cs.forward.replay()
        if not fresh and key in self.unchecked:
            # Not on the capture visit: the buffers would still hold the batch the capture recorded with.
            self.unchecked.discard(key)
            self._self_check(cs)
        self.current = cs

    def backward(self) -> Tensor:
        """Replay the backward and set p.grad as loss.backward() would; returns the n-gram sink gradient."""
        cs = self.current
        cs.backward.replay()
        for p, grad in cs.anvil_grads:
            p.grad = grad
        first, accumulate = [], []
        for p, grad, acc in cs.adam_grads:
            (first if p.grad is None else accumulate).append((acc, grad))
            p.grad = acc
        if first:
            torch._foreach_copy_([acc for acc, _ in first], [grad for _, grad in first])
        if accumulate:
            torch._foreach_add_([acc for acc, _ in accumulate], [grad for _, grad in accumulate])
        return cs.ngram_grad

    def hold_ngram_grad(self, ngram_grad: Tensor, pending: int, event_this_step: bool) -> Tensor:
        """The n-gram gradient the table may keep: the graph's own buffer only if this step's update
        event consumes it before the next replay, else a copy in ring slot `pending`."""
        if event_this_step:
            return ngram_grad
        assert pending < len(self.ngram_grad_ring), f"{pending} held n-gram gradients > ring"
        slot = self.ngram_grad_ring[pending][:ngram_grad.numel()].view_as(ngram_grad)
        slot.copy_(ngram_grad)
        return slot

    # ---- capture and checks (warmup only) ----

    def _capture(self, key: StepGraphKey, batch: Batch, ngram_slots: Tensor, cfg: ForwardScheduleConfig) -> CapturedStep:
        static = StaticInputs(*(t.clone() for t in step_inputs(batch, ngram_slots, cfg)))
        cs = CapturedStep(
            key=key, static=static,
            cfg=dataclasses.replace(cfg, mtp_weights=static.mtp_weights, prefix_weight=static.prefix_weight),
            ngram_sink=self.ngram_table.grad_sink(key.tokens), attn_scales=self._attn_scales(),
            forward=torch.cuda.CUDAGraph(), backward=torch.cuda.CUDAGraph(), loss=None,
            anvil_grads=[], adam_grads=[], ngram_grad=None,
        )
        leaves = self.params + [cs.ngram_sink]
        # The warm iterations add into value_embeds' gradient buffer and the fp8 amax buffers: harmless,
        # the post-warmup reset clears both.
        self.stream.wait_stream(torch.cuda.current_stream())
        with torch.cuda.stream(self.stream):
            for _ in range(WARM_ITERATIONS):
                torch.autograd.grad(self._loss(cs), leaves, allow_unused=True)
        torch.cuda.current_stream().wait_stream(self.stream)
        torch.cuda.synchronize()
        # thread_local: the loader and prep threads may allocate pinned memory during a capture.
        with torch.cuda.graph(cs.forward, pool=self.pool, stream=self.stream, capture_error_mode="thread_local"):
            cs.loss = self._loss(cs)
        with torch.cuda.graph(cs.backward, pool=self.pool, stream=self.stream, capture_error_mode="thread_local"):
            # retain_graph / create_graph must stay False: AOTAutograd compiles this backward with
            # donated buffers and refuses to run it otherwise (record #360).
            grads = torch.autograd.grad(cs.loss, leaves, allow_unused=True)
        torch.cuda.synchronize()
        for p, grad in zip(self.params, grads):
            if grad is None:  # value_embeds: its gradient goes to ValueEmbedPull.grad_accum instead
                continue
            if p in self.anvil_params:
                cs.anvil_grads.append((p, grad))
            else:
                if p not in self.adam_accumulators:  # shared by every key: the gradient's shape is the parameter's
                    self.adam_accumulators[p] = torch.empty_like(grad)
                cs.adam_grads.append((p, grad, self.adam_accumulators[p]))
        cs.ngram_grad = grads[-1]
        self.captured[key] = cs
        self.unchecked.add(key)
        return cs

    @torch.no_grad()
    def _self_check(self, cs: CapturedStep):
        """Replay vs eager forward on the same buffers: catches a graph that baked an input's value or a
        stale address instead of reading the live tensor."""
        replay, eager = cs.loss.item(), self._loss(cs).item()
        assert abs(replay - eager) <= max(SELF_CHECK_ABS, abs(eager) * SELF_CHECK_REL), \
            f"step graph {cs.key}: replay loss {replay} != eager {eager}"

    def seal(self):
        """After warmup: every key the timed run reaches owns a self-checked capture."""
        missing = set(self.step_keys) - set(self.captured)
        assert not missing, f"step graphs never captured: {missing}"
        assert not self.unchecked, f"step graphs never self-checked: {self.unchecked}"
        self.sealed = True

    def baked_tensors(self):
        """(role, tensor) for every tensor the step graphs bake that the model's parameters and buffers do
        not cover. Tensors other objects own are read from their owners, live, so a rebind shows up."""
        yield "value_embeds.grad_accum", self.value_embeds.grad_accum
        for tokens, sink in self.ngram_table.sinks.items():
            yield f"ngram_table.sinks[{tokens}]", sink
        for i, slot in enumerate(self.ngram_grad_ring):
            yield f"ngram_grad_ring[{i}]", slot
        for cs in self.captured.values():
            tag = "step_graph" + str(tuple(cs.key))
            for field in dataclasses.fields(cs.static):
                yield f"{tag}.{field.name}", getattr(cs.static, field.name)
            if cs.cfg.sampled_loss is not None:
                for field in dataclasses.fields(cs.cfg.sampled_loss):
                    yield f"{tag}.sampled_loss.{field.name}", getattr(cs.cfg.sampled_loss, field.name)
