Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions backends/cuda/cuda_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand All @@ -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()
Expand Down
203 changes: 173 additions & 30 deletions backends/cuda/passes/lower_offgraph_kv.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -335,14 +414,17 @@ 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"]

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
Expand All @@ -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.
Expand All @@ -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``.
Expand Down Expand Up @@ -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):
Expand All @@ -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
Expand Down
3 changes: 3 additions & 0 deletions backends/cuda/tests/targets.bzl
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
Loading
Loading