"""Original CPU teaching reference for FlashAttention's online-softmax tiling.

Unmasked forward only, query-major loop, no CUDA kernel or custom backward.
Autograd is checked for arithmetic correctness, not linear-memory backward.
The paper's original Algorithm 1 uses a key-major loop and fused GPU execution.
"""
import math
import torch
import torch.nn as nn


class DenseAttention(nn.Module):
    def __init__(self):
        super().__init__()
        self.softmax = nn.Softmax(dim=-1)

    def forward(self, q, k, v):
        scores = q @ k.transpose(-2, -1) / math.sqrt(q.shape[-1])
        probabilities = self.softmax(scores)
        return probabilities @ v


class TiledAttention(nn.Module):
    def __init__(self, block=2):
        super().__init__()
        self.block = block

    def forward(self, q, k, v):
        outputs = []
        for start in range(0, q.shape[-2], self.block):
            query = q[:, start : start + self.block]
            maximum = torch.full_like(query[..., :1], float("-inf"))
            denominator = torch.zeros_like(maximum)
            numerator = torch.zeros_like(query)
            for offset in range(0, k.shape[-2], self.block):
                keys = k[:, offset : offset + self.block]
                values = v[:, offset : offset + self.block]
                scores = query @ keys.transpose(-2, -1) / math.sqrt(q.shape[-1])
                next_maximum = torch.maximum(maximum, scores.amax(-1, keepdim=True))
                correction = torch.exp(maximum - next_maximum)
                probabilities = torch.exp(scores - next_maximum)
                denominator = denominator * correction + probabilities.sum(
                    -1, keepdim=True
                )
                numerator = numerator * correction + probabilities @ values
                maximum = next_maximum
            outputs.append(numerator / denominator)
        return torch.cat(outputs, dim=1)


class AttentionSchedules(nn.Module):
    def __init__(self):
        super().__init__()
        self.dense = DenseAttention()
        self.tiled = TiledAttention()

    def forward(self, q, k, v):
        return torch.stack([self.dense(q, k, v), self.tiled(q, k, v)])


def verify():
    torch.manual_seed(4)
    for length in [5, 6]:
        q, k, v = [torch.randn(1, length, 4, dtype=torch.float64) for _ in range(3)]
        reference = DenseAttention()(q, k, v)
        for block in [1, 2, 3, 4, 8]:
            torch.testing.assert_close(
                TiledAttention(block)(q, k, v), reference, atol=1e-12, rtol=1e-12
            )
        torch.testing.assert_close(
            TiledAttention()(q * 100, k * 100, v),
            DenseAttention()(q * 100, k * 100, v),
            atol=1e-10,
            rtol=1e-10,
        )
    q, k, v = [value.requires_grad_() for value in [q, k, v]]
    dense_grad = torch.autograd.grad(
        DenseAttention()(q, k, v).square().sum(), (q, k, v)
    )
    tiled_grad = torch.autograd.grad(
        TiledAttention()(q, k, v).square().sum(), (q, k, v)
    )
    for a, b in zip(dense_grad, tiled_grad):
        torch.testing.assert_close(a, b, atol=1e-11, rtol=1e-11)
    return [
        "Tiled attention matches dense attention across five tile sizes and uneven final tiles",
        "Running-max correction stays finite and correct for large logits",
        "CPU autograd gradients agree with dense attention; this is not the paper's memory-efficient backward",
    ]


def experiment():
    torch.manual_seed(4)
    q, k, v = [torch.randn(1, 6, 4, dtype=torch.float64) for _ in range(3)]
    query = q[:, :2]
    maximum = torch.full((1, 2, 1), float("-inf"), dtype=torch.float64)
    denominator = torch.zeros_like(maximum)
    numerator = torch.zeros_like(query)
    reference = DenseAttention()(query, k, v)
    cases = []
    for offset in [0, 2, 4]:
        scores = query @ k[:, offset : offset + 2].transpose(-2, -1) / 2
        next_maximum = torch.maximum(maximum, scores.amax(-1, keepdim=True))
        correction = torch.exp(maximum - next_maximum)
        probabilities = torch.exp(scores - next_maximum)
        denominator = denominator * correction + probabilities.sum(-1, keepdim=True)
        numerator = numerator * correction + probabilities @ v[:, offset : offset + 2]
        maximum = next_maximum
        partial = numerator / denominator
        cases.append(
            {
                "id": f"tile-{offset//2+1}",
                "label": f"Query rows 0–1 · key columns {offset}–{offset+1}",
                "target": "tiled",
                "note": "The current score tile is 2 × 2. Earlier tiles are discarded after contributing to the running maximum, denominator and weighted numerator. The partial output uses only the keys seen so far; it becomes the full attention output after the third tile. The diagram is a schedule, not a stored six-by-six score tensor.",
                "schedule": {
                    "rows": 6,
                    "columns": 6,
                    "rowStart": 0,
                    "rowCount": 2,
                    "columnStart": offset,
                    "columnCount": 2,
                    "completedColumns": offset,
                },
                "vectors": [
                    {
                        "label": "Current score tile · row-major",
                        "values": scores.flatten().tolist(),
                    },
                    {
                        "label": "Running row maximum",
                        "values": maximum.flatten().tolist(),
                    },
                    {
                        "label": "Running denominator",
                        "values": denominator.flatten().tolist(),
                    },
                    {
                        "label": "Partial output · query 0",
                        "values": partial[0, 0].tolist(),
                    },
                    {
                        "label": "Final dense output · query 0",
                        "values": reference[0, 0].tolist(),
                    },
                ],
                "metrics": [
                    {"label": "Score elements in current tile", "value": 4},
                    {"label": "Elements in a full score matrix", "value": 36},
                    {
                        "label": "Max difference from final output · rows 0–1",
                        "value": (partial - reference).abs().max().item(),
                    },
                ],
            }
        )
    return {
        "kind": "schedule",
        "title": "Finish attention without retaining every score.",
        "description": "Replay one query block through three key/value tiles. A fused GPU implementation keeps tiles in SRAM and avoids full score/probability writes to HBM. This unfused CPU reference verifies the online-softmax arithmetic; it makes no GPU memory or speed measurement.",
        "controlLabel": "Tile update",
        "cases": cases,
    }
