From d0b6945773fcf8b85d53afebf6a380debecdf26d Mon Sep 17 00:00:00 2001 From: gasoonjia Date: Fri, 2 Oct 2026 12:06:31 -0700 Subject: [PATCH] Update [ghstack-poisoned] --- backends/cuda/cuda_backend.py | 18 + backends/cuda/passes/lower_offgraph_kv.py | 203 ++++++++-- backends/cuda/tests/targets.bzl | 3 + backends/cuda/tests/test_offgraph_kv.py | 469 +++++++++++++++++++++- 4 files changed, 662 insertions(+), 31 deletions(-) diff --git a/backends/cuda/cuda_backend.py b/backends/cuda/cuda_backend.py index 41152b7c98e..d1db364505c 100644 --- a/backends/cuda/cuda_backend.py +++ b/backends/cuda/cuda_backend.py @@ -922,6 +922,11 @@ def get_extra_aoti_compile_context_manager( f"Invalid low_memory_mode: {mode}. Expected 'ON' or 'OFF'." ) low_memory_mode = mode + cell_layout = any( + spec.key == OFFGRAPH_KV_COMPILE_SPEC + and parse_offgraph_kv_manifest(spec.value)["layout"] == "cell" + for spec in compile_specs or [] + ) @contextlib.contextmanager def _combined(): @@ -947,6 +952,19 @@ def _combined(): _compile_time_cpu_clones(torch.device(cls.get_device_name())) ) trim_host_memory() + if cell_layout: + # The cell layout's step buffers are constants the graph + # only reads, while the runtime rewrites them before every + # forward. Folding them, or inlining a small one as a + # literal, would evaluate them at compile time -- on + # storage that does not exist yet -- and bake the result + # into the program. + stack.enter_context( + torch._inductor.config.patch( + joint_graph_constant_folding=False, + always_keep_tensor_constants=True, + ) + ) yield return _combined() diff --git a/backends/cuda/passes/lower_offgraph_kv.py b/backends/cuda/passes/lower_offgraph_kv.py index 637ed9cad00..5b58c88aeff 100644 --- a/backends/cuda/passes/lower_offgraph_kv.py +++ b/backends/cuda/passes/lower_offgraph_kv.py @@ -27,6 +27,19 @@ OFFGRAPH_KV_STEP_WIDTH_COMPILE_SPEC = "offgraph_kv_step_width" OFFGRAPH_KV_FQN_PREFIX = "__et_offgraph_kv_" +# Cell layout: every sequence of a batch shares one pool of per-token cells, +# and the runtime cache writes the step's placement and visibility into these +# buffers before each forward. Each is declared once per program at a fixed +# shape, so the runtime binds it at a fixed address that a captured CUDA graph +# may keep. +OFFGRAPH_KV_CELLS_FQN = OFFGRAPH_KV_FQN_PREFIX + "cells" +OFFGRAPH_KV_READ_LEN_FQN = OFFGRAPH_KV_FQN_PREFIX + "read_len" + + +def offgraph_kv_mask_fqn(window: int) -> str: + """The step mask shared by every layer of ``window`` (0 = full history).""" + return f"{OFFGRAPH_KV_FQN_PREFIX}mask_w{window}" + def ring_physical_capacity(window: int, max_write: int) -> int: """Slots a ring layer needs to serve one step of up to ``max_write`` tokens. @@ -56,6 +69,13 @@ def parse_offgraph_kv_manifest(value: bytes) -> dict[str, Any]: max_write = manifest.get("max_write") if not isinstance(max_write, int) or not 0 < max_write <= maximum_capacity: raise ValueError("off-graph max_write must be in [1, maximum_capacity]") + layout = manifest.setdefault("layout", "sequence") + if layout not in ("sequence", "cell"): + raise ValueError(f"invalid off-graph KV layout {layout!r}") + if layout == "cell": + max_cells = manifest.get("max_cells") + if not isinstance(max_cells, int) or max_cells < max_write: + raise ValueError("off-graph cell layout needs max_cells >= max_write") layers = manifest.get("layers") if not isinstance(layers, list) or not layers: raise ValueError("off-graph manifest must contain layers") @@ -170,6 +190,60 @@ def offgraph_step( return out +def cell_step( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + k_pool: torch.Tensor, + v_pool: torch.Tensor, + cells: torch.Tensor, + read_len: torch.Tensor, + mask: torch.Tensor, + scale: float, + out_dtype: torch.dtype | None = None, +) -> torch.Tensor: + """Write this step's K/V into the cells the cache placed it in, then attend. + + Placement and visibility are decided on the host by the runtime cache, + which writes them into ``cells``, ``mask`` and ``read_len`` before the + forward: token i lands in cell ``cells[i]`` and attends cell j iff + ``mask[0, 0, i, j]``. Cells of several sequences interleave in one pool, + so no causal alignment describes visibility and the mask is always + explicit. ``read_len`` bounds the sweep to the occupied extent. + + ``cells`` and ``mask`` are declared at the widest step; this one uses + their first T rows. ``k``/``v`` are BHSD; the pools are BSHD. sdpa returns + the query's dtype; ``out_dtype`` is the neutral op's requested output dtype. + """ + width = q.shape[2] + slots = cells[:width] + _write(k_pool, slots, k) + _write(v_pool, slots, v) + out = torch.ops.triton.sdpa( + q, + k_pool.transpose(1, 2), + v_pool.transpose(1, 2), + mask[:, :, :width, :], + 0.0, + False, + scale, + True, + read_len, + ) + if out_dtype is not None and out.dtype != out_dtype: + out = out.to(out_dtype) + return out + + +def _cell_step_fn(scale: float, out_dtype): + def fn(q, k, v, k_pool, v_pool, cells, read_len, mask): + return cell_step( + q, k, v, k_pool, v_pool, cells, read_len, mask, scale, out_dtype + ) + + return fn + + def _mask_fn(buf_size: int, window: int): """Freeze the constants in a closure rather than as default arguments. @@ -303,6 +377,11 @@ class LowerOffGraphKVPass: def __init__(self, manifest: dict[str, Any]) -> None: self._manifest = manifest self._masks: dict[tuple[str, int, int], Any] = {} + self._cell_buffers: dict[str, Any] = {} + + @property + def _cell_layout(self) -> bool: + return self._manifest.get("layout", "sequence") == "cell" @staticmethod def _compile_storage(shape, dtype: torch.dtype, device): @@ -335,6 +414,10 @@ def _check_supported(node) -> None: ) def _layer_capacity(self, layer: dict[str, Any]) -> int: + if self._cell_layout: + # A window bounds what a query sees, not where its token lives: a + # cell layout keeps every layer's history in the one pool. + return self._manifest["max_cells"] if layer["policy"] == "ring": return ring_physical_capacity(layer["window"], self._manifest["max_write"]) return self._manifest["maximum_capacity"] @@ -342,7 +425,6 @@ def _layer_capacity(self, layer: dict[str, Any]) -> int: def _storage_nodes( self, exported_program: ExportedProgram, node, layer: dict[str, Any] ): - graph = exported_program.graph_module.graph layer_id = layer["layer_id"] # Shape and dtype come from the step's own K, so the storage cannot # disagree with what gets written into it. The manifest only carries @@ -367,21 +449,70 @@ def _storage_nodes( self._compile_storage(shape, kv.dtype, kv.device), self._compile_storage(shape, kv.dtype, kv.device), ) - result = [] - first_node = next(iter(graph.nodes)) - for name, value in zip(names, values): - with graph.inserting_before(first_node): - result.append( - create_constant_placeholder( - exp_program=exported_program, - graph=graph, - name=name, - kind=InputKind.BUFFER, - data=value, - persistent_buffer=False, - ) - ) - return result + return [ + self._runtime_buffer(exported_program, name, value) + for name, value in zip(names, values) + ] + + @staticmethod + def _runtime_buffer(exported_program: ExportedProgram, name: str, value): + graph = exported_program.graph_module.graph + with graph.inserting_before(next(iter(graph.nodes))): + return create_constant_placeholder( + exp_program=exported_program, + graph=graph, + name=name, + kind=InputKind.BUFFER, + data=value, + persistent_buffer=False, + ) + + def _cell_buffer(self, exported_program, name: str, shape, dtype, device): + """A step buffer of the cell layout, declared once per program. + + Unlike the pools these are small, so they carry real (zero) bytes at + compile time: Inductor reads small constants while it compiles, and + they stay runtime-owned regardless, since the weight collector never + serializes a runtime-owned FQN. + """ + node = self._cell_buffers.get(name) + if node is None: + node = self._runtime_buffer( + exported_program, name, torch.zeros(shape, dtype=dtype, device=device) + ) + self._cell_buffers[name] = node + return node + + def _lower_cell( + self, exported_program, node, layer, k_storage, v_storage, out_dtype + ): + max_write = self._manifest["max_write"] + max_cells = self._manifest["max_cells"] + device = node.args[1].meta["val"].device + window = layer["window"] if layer["policy"] == "ring" else 0 + cells = self._cell_buffer( + exported_program, OFFGRAPH_KV_CELLS_FQN, (max_write,), torch.int64, device + ) + read_len = self._cell_buffer( + exported_program, OFFGRAPH_KV_READ_LEN_FQN, (1,), torch.int64, device + ) + mask = self._cell_buffer( + exported_program, + offgraph_kv_mask_fqn(window), + (1, 1, max_write, max_cells), + torch.bool, + device, + ) + call_args = ( + *node.args[0:3], # q, k, v + k_storage, + v_storage, + cells, + read_len, + mask, + ) + example_args = tuple(n.meta["val"] for n in call_args) + return _cell_step_fn(node.args[5], out_dtype), call_args, example_args def _ring_mask(self, graph, position, buf_size: int, window: int, before): """Emit the ring mask once and share it across layers that match. @@ -405,6 +536,22 @@ def _ring_mask(self, graph, position, buf_size: int, window: int, before): self._masks[key] = mask return mask + def _lower_sequence(self, graph, node, layer, k_storage, v_storage, out_dtype): + inputs = list(node.args[0:4]) # q, k, v, position + call_args = (*inputs, k_storage, v_storage) + example_args = tuple(n.meta["val"] for n in call_args) + buf_size = self._layer_capacity(layer) + masked = layer["policy"] == "ring" + if masked: + mask = self._ring_mask(graph, inputs[3], buf_size, layer["window"], node) + call_args = (*call_args, mask) + example_args = (*example_args, mask.meta["val"]) + return ( + _step_fn(node.args[5], buf_size, masked, out_dtype), + call_args, + example_args, + ) + @staticmethod def _inline(graph, fn, example_args, call_args, before): """Trace ``fn`` and splice its body in ahead of ``before``. @@ -433,8 +580,10 @@ def _inline(graph, fn, example_args, call_args, before): def __call__(self, exported_program: ExportedProgram) -> ExportedProgram: graph_module = exported_program.graph_module - # Mask nodes belong to one graph; never carry them into another. + # Mask and buffer nodes belong to one graph; never carry them into + # another. self._masks = {} + self._cell_buffers = {} target = exir_ops.edge.kvcache.update_and_attend.default modified = False for node in list(graph_module.graph.nodes): @@ -448,24 +597,18 @@ def __call__(self, exported_program: ExportedProgram) -> ExportedProgram: raise ValueError(f"off-graph manifest has no layer {layer_id}") self._check_supported(node) - inputs = list(node.args[0:4]) # q, k, v, position - scale = node.args[5] out_dtype = ( node.args[6] if len(node.args) > 6 else node.kwargs.get("out_dtype") ) k_storage, v_storage = self._storage_nodes(exported_program, node, layer) - call_args = (*inputs, k_storage, v_storage) - example_args = tuple(n.meta["val"] for n in call_args) - - buf_size = self._layer_capacity(layer) - masked = layer["policy"] == "ring" - if masked: - mask = self._ring_mask( - graph_module.graph, inputs[3], buf_size, layer["window"], node + if self._cell_layout: + fn, call_args, example_args = self._lower_cell( + exported_program, node, layer, k_storage, v_storage, out_dtype + ) + else: + fn, call_args, example_args = self._lower_sequence( + graph_module.graph, node, layer, k_storage, v_storage, out_dtype ) - call_args = (*call_args, mask) - example_args = (*example_args, mask.meta["val"]) - fn = _step_fn(scale, buf_size, masked, out_dtype) new_node = self._inline( graph_module.graph, fn, example_args, call_args, node diff --git a/backends/cuda/tests/targets.bzl b/backends/cuda/tests/targets.bzl index c88bd742c06..3e83fdd93a2 100644 --- a/backends/cuda/tests/targets.bzl +++ b/backends/cuda/tests/targets.bzl @@ -60,8 +60,11 @@ def define_common_targets(is_fbcode = False): ], deps = [ "//caffe2:torch", + "//executorch/backends/cuda:cuda_backend", + "//executorch/backends/cuda:cuda_partitioner", "//executorch/backends/cuda:cuda_passes", "//executorch/exir:lib", + "//executorch/exir/backend:compile_spec_schema", "//executorch/exir/dialects:lib", # The oracle is the neutral op itself, not a reference rebuilt here. "//executorch/extension/llm/cache:cache", diff --git a/backends/cuda/tests/test_offgraph_kv.py b/backends/cuda/tests/test_offgraph_kv.py index d77047c9292..f167375c795 100644 --- a/backends/cuda/tests/test_offgraph_kv.py +++ b/backends/cuda/tests/test_offgraph_kv.py @@ -12,19 +12,29 @@ # The very functions the lowering pass traces and splices into the graph, so a # passing test cannot agree with a decomposition the pass does not emit. +from executorch.backends.cuda.cuda_backend import CudaBackend +from executorch.backends.cuda.cuda_partitioner import CudaPartitioner from executorch.backends.cuda.passes.lower_offgraph_kv import ( + cell_step, CheckOffGraphKVStepWidthPass, LowerOffGraphKVPass, + OFFGRAPH_KV_CELLS_FQN, + OFFGRAPH_KV_COMPILE_SPEC, OFFGRAPH_KV_FQN_PREFIX, + OFFGRAPH_KV_READ_LEN_FQN, + OFFGRAPH_KV_STEP_WIDTH_COMPILE_SPEC, + offgraph_kv_mask_fqn, offgraph_step, parse_offgraph_kv_manifest, ring_attention_mask, ring_physical_capacity, ) -from executorch.exir import EdgeCompileConfig, to_edge +from executorch.exir import EdgeCompileConfig, to_edge, to_edge_transform_and_lower +from executorch.exir.backend.compile_spec_schema import CompileSpec from executorch.exir.dialects._ops import ops as exir_ops from executorch.extension.llm.cache.reference_cache import ( CacheConfig, + CellReferenceCache, LayerPolicy, SequenceReferenceCache, ) @@ -494,3 +504,460 @@ def test_lowered_program_matches_the_neutral_op(self) -> None: out = module(q, k, v, position.reshape(-1, 1)) expected = flat.step(q, k, v, position) + ring.step(q, k, v, position) self.assertLess(_max_abs_diff(out, expected), 2e-2) + + +# -- cell layout --------------------------------------------------------------- + + +class _CellOracle: + """The neutral op over ``CellReferenceCache``: the placement and masking + contract the cell layout implements, for several sequences at once. + + One layer per window, so a test can drive flat and windowed layers of the + same step. After a step it also exposes that step's plan -- the cells it + placed and the per-window masks -- which is exactly what the runtime cache + writes into the lowered program's step buffers. + """ + + def __init__(self, max_cells: int, windows=(0,)) -> None: + self.windows = tuple(windows) + self._cache = CellReferenceCache( + CacheConfig( + n_layers=len(self.windows), + n_kv_heads=N_KV_HEADS, + head_dim=HEAD_DIM, + capacity=max_cells, + dtype=torch.float32, + layers=tuple( + LayerPolicy.ring(w) if w else LayerPolicy.flat() + for w in self.windows + ), + ) + ) + self._key = f"cuda-offgraph-cell-{id(self)}" + REGISTRY.install(self._key, self._cache) + + def step(self, seq_ids, q, k, v, position): + """Declares the step, runs every layer; returns one output per layer.""" + self._cache.declare_step(seq_ids) + outputs = [] + with REGISTRY.active(self._key): + for layer_id in range(len(self.windows)): + outputs.append( + torch.ops.kvcache.update_and_attend( + q.float().cpu(), + k.float().cpu(), + v.float().cpu(), + position.reshape(-1, 1).cpu(), + layer_id, + SCALE, + torch.float32, + ) + ) + return outputs + + def plan(self): + # The reference keeps the step's plan private; it is the same lowest- + # free placement and ownership mask the C++ CellCache computes. + plan = self._cache._plan + return plan.cells, plan.base.shape[-1], plan.mask_for + + def close(self) -> None: + REGISTRY.uninstall(self._key) + + +class _CellBuffers: + """The runtime side of the cell layout: pools plus the step buffers.""" + + def __init__(self, max_write: int, max_cells: int, windows=(0,)) -> None: + self.windows = tuple(windows) + self.pools = [_storage(max_cells) for _ in self.windows] + self.cells = torch.zeros(max_write, dtype=torch.long, device="cuda") + self.read_len = torch.zeros(1, dtype=torch.long, device="cuda") + self.masks = { + w: torch.zeros(1, 1, max_write, max_cells, dtype=torch.bool, device="cuda") + for w in set(self.windows) + } + + def load(self, plan) -> None: + cells, read_len, mask_for = plan + width = cells.numel() + self.cells[:width] = cells.to("cuda") + self.read_len.fill_(read_len) + for window, mask in self.masks.items(): + mask.zero_() + mask[0, 0, :width, :read_len] = mask_for(window).to("cuda") + + def step(self, layer: int, q, k, v): + k_pool, v_pool = self.pools[layer] + return cell_step( + q, + k, + v, + k_pool, + v_pool, + self.cells, + self.read_len, + self.masks[self.windows[layer]], + SCALE, + ) + + +def _batch(groups): + """Packs (seq_id, start, length) groups onto one token axis.""" + parts = [_inputs(start, length) for _, start, length in groups] + seq_ids = [seq for seq, _, length in groups for _ in range(length)] + q = torch.cat([p[0] for p in parts], dim=2) + k = torch.cat([p[1] for p in parts], dim=2) + v = torch.cat([p[2] for p in parts], dim=2) + position = torch.cat([p[3] for p in parts]) + return seq_ids, q, k, v, position + + +# Two sequences prefilling in one step, decoding together, then one decoding +# beside the other's chunk: every combination a packed batch produces. +_INTERLEAVED = ( + ((0, 0, 5), (1, 0, 3)), + ((0, 5, 1), (1, 3, 1)), + ((1, 4, 4), (0, 6, 1)), + ((0, 7, 1),), +) + + +class OffGraphKVCellDecompositionTest(unittest.TestCase): + """cell_step, fed the reference cache's plan, must match the neutral op.""" + + @classmethod + def setUpClass(cls) -> None: + _skip_if_no_cuda() + + def _run(self, steps, max_cells: int, windows=(0,), max_write: int = 16): + oracle = _CellOracle(max_cells, windows) + self.addCleanup(oracle.close) + buffers = _CellBuffers(max_write, max_cells, windows) + for groups in steps: + seq_ids, q, k, v, position = _batch(groups) + expected = oracle.step(seq_ids, q, k, v, position) + buffers.load(oracle.plan()) + for layer in range(len(windows)): + out = buffers.step(layer, q, k, v) + self.assertLess(_max_abs_diff(out, expected[layer]), 1e-2) + return buffers, oracle + + def test_interleaved_sequences_match_oracle(self) -> None: + torch.manual_seed(10) + self._run(_INTERLEAVED, max_cells=64) + + def test_windowed_layer_masks_older_cells(self) -> None: + torch.manual_seed(11) + steps = (((0, 0, 6), (1, 0, 6)),) + tuple( + ((0, 6 + i, 1), (1, 6 + i, 1)) for i in range(5) + ) + self._run(steps, max_cells=64, windows=(0, 4)) + + def test_writes_land_in_placed_cells(self) -> None: + torch.manual_seed(12) + oracle = _CellOracle(64) + self.addCleanup(oracle.close) + buffers = _CellBuffers(16, 64) + seq_ids, q, k, v, position = _batch(((0, 0, 3), (1, 0, 2))) + oracle.step(seq_ids, q, k, v, position) + plan = oracle.plan() + buffers.load(plan) + buffers.step(0, q, k, v) + + cells = plan[0].to("cuda") + k_pool, v_pool = buffers.pools[0] + self.assertTrue(torch.equal(k_pool[0, cells], k[0].transpose(0, 1))) + self.assertTrue(torch.equal(v_pool[0, cells], v[0].transpose(0, 1))) + untouched = torch.ones(64, dtype=torch.bool, device="cuda") + untouched[cells] = False + self.assertFalse(k_pool[0, untouched].any()) + + def test_batched_sequence_sees_only_its_own_cells(self) -> None: + # Independent of the oracle: a sequence's output is the same whether + # or not another sequence shares the forward. + torch.manual_seed(13) + a = _inputs(0, 6) + b = _inputs(0, 4) + solo_oracle = _CellOracle(64) + self.addCleanup(solo_oracle.close) + solo = _CellBuffers(16, 64) + solo_oracle.step([0] * 6, *a) + solo.load(solo_oracle.plan()) + alone = solo.step(0, a[0], a[1], a[2]) + + pair_oracle = _CellOracle(64) + self.addCleanup(pair_oracle.close) + pair = _CellBuffers(16, 64) + q, k, v = (torch.cat([b[i], a[i]], dim=2) for i in range(3)) + position = torch.cat([b[3], a[3]]) + pair_oracle.step([1] * 4 + [0] * 6, q, k, v, position) + pair.load(pair_oracle.plan()) + together = pair.step(0, q, k, v) + + self.assertLess( + (together[:, :, 4:].float() - alone.float()).abs().max().item(), 1e-2 + ) + + def test_read_len_bounds_the_sweep(self) -> None: + # The runtime writes only the mask's [:width, :read_len] each step and + # leaves stale columns past read_len from earlier, wider steps. Poison + # everything past read_len -- mask true, pool NaN -- and require the + # output unchanged, on both the plain and the split-K kernels. + for seed, (groups, max_cells) in enumerate( + ((((0, 0, 5), (1, 0, 3)), 64), (((0, 0, 290), (1, 0, 3)), 512)) + ): + torch.manual_seed(20 + seed) + oracle = _CellOracle(max_cells) + self.addCleanup(oracle.close) + buffers = _CellBuffers(296, max_cells) + seq_ids, q, k, v, position = _batch(groups) + (expected,) = oracle.step(seq_ids, q, k, v, position) + plan = oracle.plan() + buffers.load(plan) + read_len = plan[1] + buffers.masks[0][..., read_len:] = True + buffers.pools[0][0][:, read_len:] = float("nan") + buffers.pools[0][1][:, read_len:] = float("nan") + out = buffers.step(0, q, k, v) + self.assertFalse(torch.isnan(out).any()) + self.assertLess(_max_abs_diff(out, expected), 1e-2) + + def test_split_k_decodes_over_a_large_pool(self) -> None: + # read_len past sdpa's split-K threshold, with one, then three, query + # rows: the decode and small-query split-K kernels under an explicit + # mask, which no single-sequence path exercises. + torch.manual_seed(14) + steps = ( + ((0, 0, 200), (1, 0, 80)), + ((2, 0, 16),), + ((0, 200, 1),), + ((0, 201, 1), (1, 80, 1), (2, 16, 1)), + ) + self._run(steps, max_cells=512, max_write=296) + + +class LowerOffGraphKVCellPassTest(unittest.TestCase): + """The pass in cell layout, on an exported program.""" + + MAX_CELLS = 64 + MAX_WRITE = 8 + WINDOW = 4 + + @classmethod + def setUpClass(cls) -> None: + _skip_if_no_cuda() + + @classmethod + def manifest(cls) -> bytes: + return json.dumps( + { + "version": 1, + "layout": "cell", + "maximum_capacity": cls.MAX_CELLS, + "max_cells": cls.MAX_CELLS, + "max_write": cls.MAX_WRITE, + "layers": [ + {"layer_id": 0, "policy": "flat"}, + {"layer_id": 1, "policy": "ring", "window": cls.WINDOW}, + ], + } + ).encode() + + def _lowered(self): + q, k, v, position = _inputs(0, self.MAX_WRITE) + t = torch.export.Dim("t", min=1, max=self.MAX_WRITE) + program = torch.export.export( + _FlatAndRing(), + (q, k, v, position.reshape(-1, 1)), + dynamic_shapes=({2: t}, {2: t}, {2: t}, {0: t}), + strict=False, + ) + edge = to_edge( + program, compile_config=EdgeCompileConfig(_check_ir_validity=False) + ).exported_program() + return LowerOffGraphKVPass(parse_offgraph_kv_manifest(self.manifest()))(edge) + + def test_manifest_layout_is_validated(self) -> None: + base = json.loads(self.manifest()) + sequence = dict(base, layout="sequence") + del sequence["max_cells"] + self.assertEqual( + parse_offgraph_kv_manifest(json.dumps(sequence).encode())["layout"], + "sequence", + ) + del sequence["layout"] + self.assertEqual( + parse_offgraph_kv_manifest(json.dumps(sequence).encode())["layout"], + "sequence", + ) + for bad in ( + dict(base, layout="paged"), + {k: v for k, v in base.items() if k != "max_cells"}, + dict(base, max_cells=self.MAX_WRITE - 1), + ): + with self.assertRaises(ValueError): + parse_offgraph_kv_manifest(json.dumps(bad).encode()) + + def test_declares_one_pool_per_layer_and_shared_step_buffers(self) -> None: + lowered = self._lowered() + graph = lowered.graph_module.graph + target = exir_ops.edge.kvcache.update_and_attend.default + self.assertFalse(any(n.target == target for n in graph.nodes)) + + buffers = { + n.name: tuple(n.meta["val"].shape) + for n in graph.nodes + if n.op == "placeholder" and n.name.startswith(OFFGRAPH_KV_FQN_PREFIX) + } + pool = (1, self.MAX_CELLS, N_KV_HEADS, HEAD_DIM) + mask = (1, 1, self.MAX_WRITE, self.MAX_CELLS) + # Both layers keep their history in a full pool; the window only picks + # which mask the layer reads. + self.assertEqual( + buffers, + { + f"{OFFGRAPH_KV_FQN_PREFIX}layer_0_k": pool, + f"{OFFGRAPH_KV_FQN_PREFIX}layer_0_v": pool, + f"{OFFGRAPH_KV_FQN_PREFIX}layer_1_k": pool, + f"{OFFGRAPH_KV_FQN_PREFIX}layer_1_v": pool, + OFFGRAPH_KV_CELLS_FQN: (self.MAX_WRITE,), + OFFGRAPH_KV_READ_LEN_FQN: (1,), + offgraph_kv_mask_fqn(0): mask, + offgraph_kv_mask_fqn(self.WINDOW): mask, + }, + ) + # The pools are declared without bytes; the step buffers are small and + # carry zeros for the compiler. Both are runtime-owned. + for name in buffers: + nbytes = lowered.constants[name].untyped_storage().nbytes() + if "layer_" in name: + self.assertEqual(nbytes, 0, name) + else: + self.assertFalse(lowered.constants[name].any(), name) + sdpa_calls = [ + n for n in graph.nodes if n.target == torch.ops.triton.sdpa.default + ] + # Always an explicit mask, never sdpa's own causal alignment. + self.assertEqual( + [(n.args[3] is not None, n.args[5]) for n in sdpa_calls], + [(True, False), (True, False)], + ) + + def test_lowered_program_matches_the_neutral_op(self) -> None: + torch.manual_seed(15) + lowered = self._lowered() + runtime = {} + for name, value in list(lowered.constants.items()): + if name.startswith(OFFGRAPH_KV_FQN_PREFIX): + runtime[name] = torch.zeros( + value.shape, device="cuda", dtype=value.dtype + ) + lowered.constants[name] = runtime[name] + module = lowered.module() + oracle = _CellOracle(self.MAX_CELLS, windows=(0, self.WINDOW)) + self.addCleanup(oracle.close) + + for groups in _INTERLEAVED: + seq_ids, q, k, v, position = _batch(groups) + flat, ring = oracle.step(seq_ids, q, k, v, position) + cells, read_len, mask_for = oracle.plan() + width = cells.numel() + # What the runtime cache writes before the forward. + runtime[OFFGRAPH_KV_CELLS_FQN][:width] = cells.to("cuda") + runtime[OFFGRAPH_KV_READ_LEN_FQN].fill_(read_len) + for window in (0, self.WINDOW): + mask = runtime[offgraph_kv_mask_fqn(window)] + mask.zero_() + mask[0, 0, :width, :read_len] = mask_for(window).to("cuda") + out = module(q, k, v, position.reshape(-1, 1)) + self.assertLess(_max_abs_diff(out, flat + ring), 2e-2) + + +class _CellAttention(torch.nn.Module): + """One layer's worth: enough to make AOTI compile the cell step.""" + + def forward(self, q, k, v, position): + return torch.ops.kvcache.update_and_attend( + q, k, v, position.reshape(-1, 1), 0, SCALE, torch.bfloat16 + ) + + +class OffGraphKVCellCompileTest(unittest.TestCase): + """A static one-token method and a dynamic one from two tokens both compile. + + The batching executor routes one-token steps to the first and wider ones to + the second, so the second's width starts at 2. triton::sdpa dispatches on + the query length (one token, 2..4 tokens, wider), and a pool past its + split-K threshold makes those branches live under a symbolic length. + """ + + MAX_CELLS = 512 + MAX_WRITE = 32 + + @classmethod + def setUpClass(cls) -> None: + _skip_if_no_cuda() + + def test_decode_and_prefill_from_two_tokens_lower(self) -> None: + manifest = json.dumps( + { + "version": 1, + "layout": "cell", + "maximum_capacity": self.MAX_CELLS, + "max_cells": self.MAX_CELLS, + "max_write": self.MAX_WRITE, + "layers": [{"layer_id": 0, "policy": "flat"}], + } + ).encode() + t = torch.export.Dim("t", min=2, max=self.MAX_WRITE) + programs = { + "decode": torch.export.export( + _CellAttention(), _inputs(0, 1), strict=True + ), + "prefill": torch.export.export( + _CellAttention(), + _inputs(0, 8), + dynamic_shapes=({2: t}, {2: t}, {2: t}, {0: t}), + strict=True, + ), + } + # The PCH path shells out to `openssl sha512` and is flaky; every + # off-graph export turns it off, so the test compiles the same way. + import torch._inductor.config as inductor_config + + with inductor_config.patch({"aot_inductor.precompile_headers": False}): + lowered = self._lower(programs, manifest) + for name in programs: + graph = lowered.exported_program(name).graph + self.assertTrue( + any("executorch_call_delegate" in str(n.target) for n in graph.nodes), + name, + ) + + @staticmethod + def _lower(programs, manifest): + return to_edge_transform_and_lower( + programs, + partitioner={ + name: [ + CudaPartitioner( + [ + CudaBackend.generate_method_name_compile_spec(name), + CompileSpec(OFFGRAPH_KV_COMPILE_SPEC, manifest), + # The step width names the delegate's input + # order, which partitioning decides: position + # comes first. + CompileSpec(OFFGRAPH_KV_STEP_WIDTH_COMPILE_SPEC, b"0:0"), + # Runtime-owned storage is declared without bytes; + # low-memory mode is what compiles such constants, + # as every off-graph export does. + CompileSpec("low_memory_mode", b"ON"), + ] + ) + ] + for name in programs + }, + compile_config=EdgeCompileConfig(_check_ir_validity=False), + )