"""Small standalone literal-ISA emitter for architecture-conditioned compilers.

This module has no dependency on the baseline compiler or its schedules.
"""

import json
import re
from contextlib import contextmanager
from dataclasses import dataclass


def add(*x):
    return {"add": list(x)}


def mul(*x):
    return {"mul": list(x)}


def imm(x):
    return {"imm": x}


def rf(lane, count, offset=0):
    return {"space": "RF", "offset": offset, "count": count, "wg": "g", "lane": lane}


def hbm(offset, count, shape=None, strides=None):
    v = {"space": "HBM", "offset": offset, "count": count, "wg": None, "lane": 0}
    if shape is not None:
        v.update(shape=shape, strides=strides)
    return v


@dataclass(frozen=True)
class Tensor:
    offset: int
    shape: tuple[int, ...]
    packed: bool = False
    row_offset: int = 0
    block_cols: int = 32


def packed_offset(row, col, width, block_cols=32):
    block_row = add({"ceildiv": [add(row, 1), 32]}, -1)
    block_col = add({"ceildiv": [add(col, 1), block_cols]}, -1)
    return add(
        mul(block_row, 32 * width),
        mul(block_col, 32 * block_cols),
        mul({"mod": [row, 32]}, block_cols),
        {"mod": [col, block_cols]},
    )


def tile_view(t, row, col, nr, nc, width):
    row = add(row, t.row_offset)
    if t.packed:
        return hbm(
            add(t.offset, packed_offset(row, col, width, t.block_cols)),
            nr * nc,
            [nr, nc],
            [t.block_cols, 1],
        )
    return hbm(add(t.offset, mul(row, width), col), nr * nc, [nr, nc], [width, 1])


def row_view(t, row, width):
    row = add(row, t.row_offset)
    if t.packed:
        return hbm(
            add(t.offset, packed_offset(row, 0, width, t.block_cols)),
            width,
            [width // t.block_cols, t.block_cols],
            [32 * t.block_cols, 1],
        )
    return hbm(add(t.offset, mul(row, width)), width)


def fold(value):
    """Canonicalize constant address arithmetic without inspecting tensor values."""
    if isinstance(value, list):
        return [fold(v) for v in value]
    if not isinstance(value, dict):
        return value
    value = {k: fold(v) for k, v in value.items()}
    if len(value) != 1:
        return value
    op, args = next(iter(value.items()))
    if op in ("add", "mul"):
        items = []
        for arg in args:
            items.extend(
                arg[op] if isinstance(arg, dict) and set(arg) == {op} else [arg]
            )
        const = 0 if op == "add" else 1
        rest = []
        for arg in items:
            if type(arg) == int:
                const = const + arg if op == "add" else const * arg
            else:
                rest.append(arg)
        if op == "mul" and const == 0:
            return 0
        if const != (0 if op == "add" else 1) or not rest:
            rest.insert(0, const)
        return rest[0] if len(rest) == 1 else {op: rest}
    if op in ("ceildiv", "mod") and all(type(v) == int for v in args):
        return (
            (args[0] + args[1] - 1) // args[1] if op == "ceildiv" else args[0] % args[1]
        )
    return value


class Builder:
    def __init__(self, layout, **_):
        self.layout = layout
        self.lines = []
        self.loops = []
        self.serial = 0
        self.emit("WG.BEGIN", wg="g", sm=0, shared_bytes=0)

    def emit(self, op, **args):
        if "event" in args and args["event"] is None:
            self.serial += 1
            args["event"] = f"e{self.serial}" + "".join(
                f"_{{{name}}}" for name in self.loops
            )
        self.lines.append(op + " " + json.dumps(fold(args), separators=(",", ":")))

    def wait_last(self):
        event = f"e{self.serial}" + "".join(f"_{{{name}}}" for name in self.loops)
        self.emit("WAIT", wg="g", events=[event])

    @contextmanager
    def loop(self, name, start, stop, step=1):
        self.emit("FOR", var=name, start=start, stop=stop, step=step)
        self.loops.append(name)
        yield {"var": name}
        self.loops.pop()
        self.emit("END.FOR")

    def ld(self, src, dst):
        self.emit("LD", src=src, dst=dst, event=None)

    def st(self, src, dst):
        self.emit("ST", src=src, dst=dst, event=None)

    def vec(self, kind, src, dst):
        self.emit("VEC", kind=kind, src=src, dst=dst, event=None)

    def reduce(self, kind, src, dst):
        self.emit("REDUCE", kind=kind, src=src, dst=dst, event=None)

    def sfu(self, kind, src, dst):
        self.emit("SFU", kind=kind, src=src, dst=dst, event=None)

    def layernorm(
        self, x, gamma, beta, y, rows, width, name, addend=None, residual_out=None
    ):
        # Centered variance, stable for arbitrary finite public/hidden fixtures.
        with self.loop(name, 0, rows) as row:
            self.ld(row_view(x, row, width), rf(0, width))
            if addend is not None:
                self.ld(row_view(addend, row, width), rf(3, width))
                self.vec("add", [rf(0, width), rf(3, width)], rf(0, width))
                self.st(rf(0, width), row_view(residual_out, row, width))
            self.reduce("sum", rf(0, width), rf(1, 1))
            self.vec("mul", [rf(1, 1), imm(1 / width)], rf(1, 1))
            self.vec("sub", [rf(0, width), rf(1, 1)], rf(2, width))
            self.vec("mul", [rf(2, width), rf(2, width)], rf(3, width))
            self.reduce("sum", rf(3, width), rf(1, 1))
            self.vec("mul", [rf(1, 1), imm(1 / width)], rf(1, 1))
            self.vec("add", [rf(1, 1), imm(1e-5)], rf(1, 1))
            self.sfu("rsqrt", rf(1, 1), rf(1, 1))
            self.vec("mul", [rf(2, width), rf(1, 1)], rf(2, width))
            self.ld(hbm(gamma.offset, width), rf(3, width))
            self.vec("mul", [rf(2, width), rf(3, width)], rf(2, width))
            self.ld(hbm(beta.offset, width), rf(3, width))
            self.vec("add", [rf(2, width), rf(3, width)], rf(2, width))
            self.st(rf(2, width), row_view(y, row, width))


def rename(lines, group, sm):
    result = []
    for line in lines:
        line = line.replace('"wg":"g"', f'"wg":"g{group}"')
        line = re.sub(r'(?<=")e(?=\d)', f"g{group}_e", line)
        if line.startswith("WG.BEGIN "):
            args = json.loads(line[9:])
            args["sm"] = sm
            line = "WG.BEGIN " + json.dumps(fold(args), separators=(",", ":"))
        result.append(line)
    return result
