From 353de3e394f4186b41d2e72d4d0b2ea9e2573cd3 Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Sun, 20 Sep 2026 19:42:12 +0800 Subject: [PATCH 1/9] fix(analysis): model loop-aware value lifetimes --- docs/spec/analysis.md | 15 ++- src/tilefoundry/analysis/liveness.py | 156 +++++++++++++++++++++++ src/tilefoundry/analysis/memory.py | 110 ++++++++++------ src/tilefoundry/analysis/metadata.py | 2 +- tests/analysis/test_analyze_at_a_size.py | 28 +++- 5 files changed, 265 insertions(+), 46 deletions(-) create mode 100644 src/tilefoundry/analysis/liveness.py diff --git a/docs/spec/analysis.md b/docs/spec/analysis.md index 580eda05..5a7c486d 100644 --- a/docs/spec/analysis.md +++ b/docs/spec/analysis.md @@ -438,11 +438,11 @@ of this analysis. | `ValueLifetime.binding` | Use the parameter or binding name, suffixed with `:` and the line of the value's source span when it has one. Repeated names already differ by the printer's numeric suffix in definition order; the line locates the row in authored source, which a suffix cannot. A value with neither name nor span is `` in definition order. | No | | `ValueLifetime.memory_level` | Emit one lifetime per storage level occupied by the value's Type. | No | | `ValueLifetime.bytes` | Project the Type through every authored split at or coarser than the explicit level's `owner`, then take its logical bytes; a target-owned or undeclared level remains global. | `MemoryHierarchyFacts.explicit_levels[].owner` | -| `ValueLifetime.defined_at` | Position in the order of parameters followed by body Calls and Constants in SSA postorder. | No | -| `ValueLifetime.last_used_at` | Greatest recorded consumer position; the last position for a parameter, and also for the Function body when that body is itself a recorded value. | No | +| `ValueLifetime.defined_at` | Definition event on the function-wide structured SSA timeline. | No | +| `ValueLifetime.last_used_at` | Greatest ordinary-consumer, region-entry, loop-backedge, or region-exit use event; the final timeline event for a parameter. | No | | `ValueLifetime.persistent` | True for parameters and false for body allocations. | No | | `MemoryLevelFootprint.memory_level` | Each storage level with at least one lifetime, sorted by name. | No | -| `MemoryLevelFootprint.peak_bytes` | Largest sum of simultaneously live bytes at that level over the value order. | No | +| `MemoryLevelFootprint.peak_bytes` | Largest sum of simultaneously live bytes at that level over the structured SSA event timeline. | No | | `MemoryLevelFootprint.persistent_bytes` | Sum of persistent lifetimes at that level. | No | | `MemoryLevelFootprint.capacity_bytes` | Capacity of the matching explicit level, or `None` when it is unknown or undeclared. | `MemoryHierarchyFacts.explicit_levels[].capacity_bytes` | | `MemoryMetadata.footprint` | One `MemoryLevelFootprint` per occupied storage level. | As above | @@ -452,6 +452,15 @@ of this analysis. | `TrafficMetadata.communication` | What a Reshard sends off the unit it was on, when the shards on its two sides differ across a mesh axis that level owns. Zero where they agree. | No; the share each unit keeps follows from the mesh extents the shards name. | | `TrafficMetadata.operands` | One occurrence's per-boundary movement in order `(*call.args, call)`, the same relation-derived amounts `storage` groups. Empty on a Function and on a Function Call, neither of which has a split. | No | +One ordinary expression event uses its operands and defines its result. A region +adds separate binding and exit events: a mesh argument is used before its +parameter is defined, and a loop initial value is used before its induction and +carried parameters are defined. One representative loop-body iteration is +recorded without expanding the trip count; a carried parameter spans entry to +exit, and a yielded value remains live through the backedge event. Event +positions are monotonic across the whole Function, including nested and sibling +regions. + The target-aware loop projection is report data rather than another metadata record. `LoopFootprintMetadata` remains target-independent: diff --git a/src/tilefoundry/analysis/liveness.py b/src/tilefoundry/analysis/liveness.py new file mode 100644 index 00000000..5e977c8b --- /dev/null +++ b/src/tilefoundry/analysis/liveness.py @@ -0,0 +1,156 @@ +"""Target-independent definition/use intervals for structured HIR SSA.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from tilefoundry.ir.core import Expr, Var +from tilefoundry.ir.hir.function import Function +from tilefoundry.ir.hir.loop_region import LoopRegion +from tilefoundry.ir.hir.mesh_region import MeshRegion +from tilefoundry.ir.visitor import ExprVisitor, collect_exprs, expr_children + + +@dataclass(frozen=True) +class LiveInterval: + """One SSA value's definition and greatest use event.""" + + value: Expr + defined_at: int + last_used_at: int + + +@dataclass(frozen=True) +class Liveness: + """Definition-ordered intervals on one function-wide event timeline.""" + + intervals: tuple[LiveInterval, ...] + timeline_end: int + + +@dataclass +class _IntervalState: + value: Expr + defined_at: int + last_used_at: int + + +def _free_vars(function: Function) -> tuple[Var, ...]: + """Boundary Vars reached as uses but not owned by a structural binding site.""" + if function.body is None: + return () + values = collect_exprs(function.body) + bound_ids = {id(parameter) for parameter in function.params} + for value in values: + if isinstance(value, MeshRegion): + bound_ids.update(id(parameter) for parameter in value.params) + elif isinstance(value, LoopRegion): + bound_ids.add(id(value.induction_var)) + bound_ids.update(id(phi) for phi in value.carried_args) + return tuple(value for value in values if isinstance(value, Var) and id(value) not in bound_ids) + + +class LivenessVisitor(ExprVisitor[None]): + """Build definition/use intervals while preserving structured SSA edges.""" + + def __init__(self, function: Function) -> None: + super().__init__(root_function=function) + self._point = -1 + self._states: dict[int, _IntervalState] = {} + self._definition_order: list[int] = [] + for parameter in function.params: + self.define(parameter, self.next_event()) + for free in _free_vars(function): + self.define(free, self.next_event()) + + def next_event(self) -> int: + """Advance and return the function-wide event position.""" + self._point += 1 + return self._point + + def define(self, value: Expr, point: int) -> None: + """Record the one definition of *value*.""" + key = id(value) + if key in self._states: + raise ValueError(f"liveness: {type(value).__name__} is defined more than once") + self._states[key] = _IntervalState(value, point, point) + self._definition_order.append(key) + + def use(self, value: Expr, point: int) -> None: + """Extend *value* through one consumer event.""" + state = self._states.get(id(value)) + if state is None: + raise ValueError(f"liveness: {type(value).__name__} is used before its definition") + state.last_used_at = max(state.last_used_at, point) + + def finish(self) -> Liveness: + """Freeze the definition-ordered result.""" + states = (self._states[key] for key in self._definition_order) + return Liveness( + intervals=tuple( + LiveInterval(state.value, state.defined_at, state.last_used_at) for state in states + ), + timeline_end=self._point, + ) + + def visit_Var(self, value: Var, _ctx=None) -> None: + """A binding site defines a Var; encountering a use does not.""" + + def default_visit_leaf(self, value: Expr, _operands: tuple[None, ...], _ctx=None) -> None: + point = self.next_event() + for operand in expr_children(value): + self.use(operand, point) + self.define(value, point) + + def visit_MeshRegion(self, region: MeshRegion, ctx=None) -> None: + """Make the argument/parameter and body/result binding edges explicit.""" + for argument in region.args: + self.visit(argument, ctx) + argument_use = self.next_event() + for argument in region.args: + self.use(argument, argument_use) + parameter_definition = self.next_event() + for parameter in region.params: + self.define(parameter, parameter_definition) + + self.visit(region.body, ctx) + body_use = self.next_event() + self.use(region.body, body_use) + self.define(region, self.next_event()) + + def visit_LoopRegion(self, region: LoopRegion, ctx=None) -> None: + """Represent entry, one body iteration, backedge, and exit events.""" + for initial in region.init_args: + self.visit(initial, ctx) + entry_use = self.next_event() + for initial in region.init_args: + self.use(initial, entry_use) + + phi_definition = self.next_event() + self.define(region.induction_var, phi_definition) + for phi in region.carried_args: + self.define(phi, phi_definition) + + self.visit(region.body, ctx) + for yielded in region.yield_values: + self.visit(yielded, ctx) + backedge = self.next_event() + for yielded in region.yield_values: + self.use(yielded, backedge) + + exit_use = self.next_event() + for source in region.carried_args or (region.body,): + self.use(source, exit_use) + self.define(region, self.next_event()) + + +def analyze_liveness(function: Function) -> Liveness: + """Build target-independent intervals for a checked HIR function.""" + if function.body is None: + raise ValueError(f"liveness: function {function.name!r} has no body") + visitor = LivenessVisitor(function) + visitor.visit_function_body(function) + return visitor.finish() + + +__all__ = ["LiveInterval", "Liveness", "analyze_liveness"] diff --git a/src/tilefoundry/analysis/memory.py b/src/tilefoundry/analysis/memory.py index 4660a9af..33f90e41 100644 --- a/src/tilefoundry/analysis/memory.py +++ b/src/tilefoundry/analysis/memory.py @@ -8,17 +8,17 @@ Call, Constant, Expr, - Var, VerifyError, describe_expr, value_labels, ) from tilefoundry.ir.core import attach_metadata as attach +from tilefoundry.ir.core.module import Module from tilefoundry.ir.hir.function import Function from tilefoundry.ir.hir.loop_region import LoopRegion from tilefoundry.ir.types import TensorType, TupleType, Type, bytes_by_storage from tilefoundry.ir.types.storage import StorageKind -from tilefoundry.ir.visitor import ExprVisitor, expr_children +from tilefoundry.ir.visitor import ExprVisitor from tilefoundry.visitor_registry.access_relation import ( AccessRelations, access_relation_registry, @@ -34,6 +34,7 @@ from .errors import AnalysisError from .facts import MemoryHierarchyFacts +from .liveness import Liveness, analyze_liveness from .metadata import ( AllocationMetadata, Breakdown, @@ -307,56 +308,85 @@ def add_traffic( ) -def _lifetimes( - values: list[Expr], facts: MemoryHierarchyFacts, local: CostContext -) -> tuple[ValueLifetime, ...]: - """Each value's residency, its bytes taken in the analysed level's window. +def _resident_value_ids(function: Function, liveness: Liveness) -> frozenset[int]: + """Values whose SSA interval represents independently resident bytes.""" + result = {id(parameter) for parameter in function.params} + for interval in liveness.intervals: + value = interval.value + if isinstance(value, (Call, Constant, LoopRegion)): + result.add(id(value)) + if isinstance(value, LoopRegion): + result.update(id(phi) for phi in value.carried_args) + return frozenset(result) - The projection is the one the traffic family reads, so both halves of a report - answer for the same unit. Dividing the whole tensor by a per-level constant - outside was a second account of it, and the two disagreed. - """ - index_by_id = {id(expr): index for index, expr in enumerate(values)} - last_by_id = dict(index_by_id) - for index, consumer in enumerate(values): - for operand in expr_children(consumer): - if id(operand) in last_by_id and index > last_by_id[id(operand)]: - last_by_id[id(operand)] = index + +def _project_value_lifetimes( + liveness: Liveness, + resident_ids: frozenset[int], + parameter_ids: frozenset[int], + facts: MemoryHierarchyFacts, + local: CostContext, +) -> tuple[ValueLifetime, ...]: + """Project structural intervals into the analysed topology window.""" + intervals = tuple( + interval for interval in liveness.intervals if id(interval.value) in resident_ids + ) result: list[ValueLifetime] = [] - labels = value_labels(values) - for index, expr in enumerate(values): + labels = value_labels(interval.value for interval in intervals) + for label, interval in zip(labels, intervals, strict=True): + expr = interval.value + persistent = id(expr) in parameter_ids for memory_level, amount in bytes_by_storage(local.local_type_of(expr)).items(): if facts.explicit(memory_level) is None: continue result.append( ValueLifetime( - binding=labels[index], + binding=label, memory_level=memory_level, bytes=amount, - defined_at=index, - last_used_at=( - len(values) - 1 if isinstance(expr, Var) else last_by_id[id(expr)] - ), - persistent=isinstance(expr, Var), + defined_at=interval.defined_at, + last_used_at=(liveness.timeline_end if persistent else interval.last_used_at), + persistent=persistent, ) ) return tuple(result) +def analyze_value_lifetimes( + module: Module, + function: Function, + *, + topology_level: str | None = None, +) -> tuple[ValueLifetime, ...]: + """Project checked structural SSA liveness into memory residency.""" + liveness = analyze_liveness(function) + facts = module.resolve_target().get_facts(MemoryHierarchyFacts) + local = CostContext( + scope=FunctionScope(module, function), + topology_level=topology_level, + topologies=module.effective_topologies(), + ) + return _project_value_lifetimes( + liveness, + _resident_value_ids(function, liveness), + frozenset(id(parameter) for parameter in function.params), + facts, + local, + ) + + @dataclass class MemoryContext(AnalyzeContext): """State carried through the memory-family expression walk.""" whole: CostContext | None = None - local: CostContext | None = None locals_by_unit: dict[str, CostContext] = field(default_factory=dict) totals: dict[str, dict[str, TrafficBytes]] = field(default_factory=dict) shares: dict[str, dict[str, TrafficBytes]] = field(default_factory=dict) - values: list[Expr] = field(default_factory=list) class MemoryVisitor(ExprVisitor[None]): - """Attach per-Call traffic while collecting lifetime order and loop footprints.""" + """Attach per-Call traffic and loop footprints.""" def visit_LoopRegion(self, expr: LoopRegion, ctx: MemoryContext) -> None: child = next(item for item in ctx.current.children if item.owner is expr) @@ -371,8 +401,6 @@ def visit_LoopRegion(self, expr: LoopRegion, ctx: MemoryContext) -> None: def default_visit_leaf( self, expr: Expr, _operands: tuple[None, ...], ctx: MemoryContext ) -> None: - if isinstance(expr, (Call, Constant)): - ctx.values.append(expr) if not isinstance(expr, Call): return recorded = id(expr) in ctx.current.accesses["narrow"] @@ -408,11 +436,6 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: facts = context.target.get_facts(MemoryHierarchyFacts) topologies = module.effective_topologies() whole = CostContext(scope=FunctionScope(module, function)) - local = CostContext( - scope=FunctionScope(module, function), - topology_level=topology_level, - topologies=topologies, - ) units = tuple(topology.name for topology in topologies) or ( (topology_level,) if topology_level else () ) @@ -432,9 +455,7 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: root=context.root, current=context.current, whole=whole, - local=local, locals_by_unit=locals_by_unit, - values=list(function.params), ) MemoryVisitor().visit(function.body, memory_context) attach( @@ -445,13 +466,18 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: communication=_shares(memory_context.shares, tuple(locals_by_unit)), ), ) - lifetimes = _lifetimes(memory_context.values, facts, local) + lifetimes = analyze_value_lifetimes( + module, + function, + topology_level=topology_level, + ) levels_list: list[MemoryLevelFootprint] = [] for name in sorted({item.memory_level for item in lifetimes} | set(memory_context.totals)): declared = facts.explicit(name) rows = [item for item in lifetimes if item.memory_level == name] peak = 0 - for point in range(len(lifetimes) + 1): + end = max((item.last_used_at for item in rows), default=-1) + for point in range(end + 1): peak = max( peak, sum(item.bytes for item in rows if item.defined_at <= point <= item.last_used_at), @@ -527,4 +553,10 @@ def cache_pressure( return tuple(rows) -__all__ = ["MemoryOptions", "SELECTOR", "analyze_memory", "cache_pressure"] +__all__ = [ + "MemoryOptions", + "SELECTOR", + "analyze_memory", + "analyze_value_lifetimes", + "cache_pressure", +] diff --git a/src/tilefoundry/analysis/metadata.py b/src/tilefoundry/analysis/metadata.py index 9e015eec..309298a2 100644 --- a/src/tilefoundry/analysis/metadata.py +++ b/src/tilefoundry/analysis/metadata.py @@ -174,7 +174,7 @@ class LoopFootprintMetadata(IRMetadata): @dataclass(frozen=True) class ValueLifetime: - """One value's residency, as positions in the function's definition order. + """One value's residency on the function's structured SSA event timeline. ``persistent`` marks a value that is resident for the whole function rather than until its last use. Every parameter is persistent because a function diff --git a/tests/analysis/test_analyze_at_a_size.py b/tests/analysis/test_analyze_at_a_size.py index e1a3a106..fd248ba0 100644 --- a/tests/analysis/test_analyze_at_a_size.py +++ b/tests/analysis/test_analyze_at_a_size.py @@ -52,6 +52,11 @@ DIMS = {"ctx_len": CONTEXT} FAMILIES = ("compute-cost", "memory", "roofline", "performance") INVENTORY = [pytest.param(case, id=case.id) for case in placed_cases()] +EXPECTED_MEMORY_PEAKS = { + "qwen3_1_7b_pd.PrefillLayer.layer_prefill[ctx_len=128,seq=128]": { + "rmem": 395_264, + }, +} def _aimed(): @@ -206,9 +211,7 @@ def _every_number_counts_something(result: AnalysisResult) -> None: breakdown = getattr(held, field) for name, spread in breakdown.kinds: for value in (spread.logical, spread.total, *spread.per_unit): - assert value >= 0, ( - f"{describe_expr(expr)}: {field}[{name}] = {value}" - ) + assert value >= 0, f"{describe_expr(expr)}: {field}[{name}] = {value}" if record is TrafficMetadata: for field in ("whole", "per_unit"): for level, moved in getattr(held, field): @@ -222,6 +225,20 @@ def _every_number_counts_something(result: AnalysisResult) -> None: if record is MemoryMetadata: for level in held.footprint: assert level.peak_bytes >= 0 and level.persistent_bytes >= 0 + rows = [item for item in held.lifetimes if item.level == level.level] + end = max((item.last_used_at for item in rows), default=-1) + expected_peak = max( + ( + sum( + item.bytes + for item in rows + if item.defined_at <= point <= item.last_used_at + ) + for point in range(end + 1) + ), + default=0, + ) + assert level.peak_bytes == expected_peak for item in held.lifetimes: assert item.bytes >= 0 and 0 <= item.defined_at <= item.last_used_at assert " None: assert result.module is owner assert set(result.executed) == set(FAMILIES) assert_performance_contract(result) + expected = EXPECTED_MEMORY_PEAKS.get(case.id, {}) + placement = get_metadata(result.function, MemoryMetadata) + assert placement is not None + observed = {item.level: item.peak_bytes for item in placement.footprint} + assert {level: observed[level] for level in expected} == expected @pytest.mark.parametrize("family", FAMILIES) From b3d93fcd7f8544c031351db0af193206ba2a9eef Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Sun, 20 Sep 2026 20:12:37 +0800 Subject: [PATCH 2/9] docs(tutorial): refresh loop-aware memory peaks --- docs/tutorial/showcase.ipynb | 4 ++-- docs/tutorial/showcase.md | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/docs/tutorial/showcase.ipynb b/docs/tutorial/showcase.ipynb index ab866d67..8fa29fa7 100644 --- a/docs/tutorial/showcase.ipynb +++ b/docs/tutorial/showcase.ipynb @@ -360,7 +360,7 @@ { "name": "stdout", "output_type": "stream", - "text": "# analysis target=nvidia.h200_sxm module=Stage3_Fused function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,4392896@total,328672@cta;f32:3239808@logical,6418944@total,200592@cta other-ops=integer:9@logical,288@total,9@cta;special:33056@logical,33152@total,1036@cta\n# traffic traffic=gmem:r3476884/w4196992@total,r2558932/w4196544@cta;rmem:r656/w72@total,r656/w72@cta;smem:r5839296/w5662784@total,r682464/w672676@cta\n# peak-footprint=gmem:5047436;rmem:16;smem:33408\n# roofline ideal-ns=1599 bound-by=memory\n\n v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n" + "text": "# analysis target=nvidia.h200_sxm module=Stage3_Fused function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,4392896@total,328672@cta;f32:3239808@logical,6418944@total,200592@cta other-ops=integer:9@logical,288@total,9@cta;special:33056@logical,33152@total,1036@cta\n# traffic traffic=gmem:r3476884/w4196992@total,r2558932/w4196544@cta;rmem:r656/w72@total,r656/w72@cta;smem:r5839296/w5662784@total,r682464/w672676@cta\n# peak-footprint=gmem:7144460;rmem:16;smem:33552\n# roofline ideal-ns=1599 bound-by=memory\n\n v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n" } ], "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage3-4096.txt\").read_text(encoding=\"utf-8\")\nheader, separator, annotated = report.partition(\"\\n\\n\")\nprint(header.rstrip())\nprint()\nprint(next(line.rstrip() for line in annotated.splitlines() if \"cache_update(k_cache\" in line))\n" @@ -462,7 +462,7 @@ { "name": "stdout", "output_type": "stream", - "text": "# analysis target=nvidia.h200_sxm module=Stage5_CachePrepared function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,1246400@total,328672@cta;f32:6410280@logical,6410280@total,801285@cta other-ops=integer:32@logical,256@total,32@cta;special:33040@logical,33040@total,4130@cta\n# traffic traffic=gmem:r7672476/w4198272@total,r4001116/w4197824@cta;rmem:r2560/w0@total,r2560/w0@cta;smem:r21788704/w21486432@total,r2723588/w2685804@cta\n# peak-footprint=gmem:2950412;rmem:0;smem:33536\n# roofline ideal-ns=2474 bound-by=memory\n\n v21 = slice(k_cache, (0, v20, 0, 0), sizes=(1, 128, 2, 32), strides=(1, 1, 1, 1)) # Tensor[(1, 128, 2, 32), \"bf16\"]; compute-cost; traffic traffic=rmem:r32/w0@total,r32/w0@cta operands=0:r0/w0;1:r32/w0;result:r0/w0; roofline\n v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n" + "text": "# analysis target=nvidia.h200_sxm module=Stage5_CachePrepared function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,1246400@total,328672@cta;f32:6410280@logical,6410280@total,801285@cta other-ops=integer:32@logical,256@total,32@cta;special:33040@logical,33040@total,4130@cta\n# traffic traffic=gmem:r7672476/w4198272@total,r4001116/w4197824@cta;rmem:r2560/w0@total,r2560/w0@cta;smem:r21788704/w21486432@total,r2723588/w2685804@cta\n# peak-footprint=gmem:3474572;rmem:0;smem:33680\n# roofline ideal-ns=2474 bound-by=memory\n\n v21 = slice(k_cache, (0, v20, 0, 0), sizes=(1, 128, 2, 32), strides=(1, 1, 1, 1)) # Tensor[(1, 128, 2, 32), \"bf16\"]; compute-cost; traffic traffic=rmem:r32/w0@total,r32/w0@cta operands=0:r0/w0;1:r32/w0;result:r0/w0; roofline\n v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n" } ], "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage5-4096.txt\").read_text(encoding=\"utf-8\")\nheader, separator, annotated = report.partition(\"\\n\\n\")\nprint(header.rstrip())\nprint()\nfor needle in (\"slice(k_cache\", \"cache_update(k_cache\"):\n print(next(line.rstrip() for line in annotated.splitlines() if needle in line))\n" diff --git a/docs/tutorial/showcase.md b/docs/tutorial/showcase.md index 2305af92..2bc44914 100644 --- a/docs/tutorial/showcase.md +++ b/docs/tutorial/showcase.md @@ -730,7 +730,7 @@ print(next(line.rstrip() for line in annotated.splitlines() if "cache_update(k_c # selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline # compute-cost flops=bf16:328896@logical,4392896@total,328672@cta;f32:3239808@logical,6418944@total,200592@cta other-ops=integer:9@logical,288@total,9@cta;special:33056@logical,33152@total,1036@cta # traffic traffic=gmem:r3476884/w4196992@total,r2558932/w4196544@cta;rmem:r656/w72@total,r656/w72@cta;smem:r5839296/w5662784@total,r682464/w672676@cta -# peak-footprint=gmem:5047436;rmem:16;smem:33408 +# peak-footprint=gmem:7144460;rmem:16;smem:33552 # roofline ideal-ns=1599 bound-by=memory v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), "bf16"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory @@ -1019,7 +1019,7 @@ for needle in ("slice(k_cache", "cache_update(k_cache"): # selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline # compute-cost flops=bf16:328896@logical,1246400@total,328672@cta;f32:6410280@logical,6410280@total,801285@cta other-ops=integer:32@logical,256@total,32@cta;special:33040@logical,33040@total,4130@cta # traffic traffic=gmem:r7672476/w4198272@total,r4001116/w4197824@cta;rmem:r2560/w0@total,r2560/w0@cta;smem:r21788704/w21486432@total,r2723588/w2685804@cta -# peak-footprint=gmem:2950412;rmem:0;smem:33536 +# peak-footprint=gmem:3474572;rmem:0;smem:33680 # roofline ideal-ns=2474 bound-by=memory v21 = slice(k_cache, (0, v20, 0, 0), sizes=(1, 128, 2, 32), strides=(1, 1, 1, 1)) # Tensor[(1, 128, 2, 32), "bf16"]; compute-cost; traffic traffic=rmem:r32/w0@total,r32/w0@cta operands=0:r0/w0;1:r32/w0;result:r0/w0; roofline From 63f3317358b761b00390a17a2a66fd8f7f5a03c0 Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Sun, 20 Sep 2026 22:36:28 +0800 Subject: [PATCH 3/9] refactor(analysis): reuse live interval state --- src/tilefoundry/analysis/liveness.py | 19 +++++-------------- tests/analysis/test_analyze_at_a_size.py | 2 ++ 2 files changed, 7 insertions(+), 14 deletions(-) diff --git a/src/tilefoundry/analysis/liveness.py b/src/tilefoundry/analysis/liveness.py index 5e977c8b..0327951b 100644 --- a/src/tilefoundry/analysis/liveness.py +++ b/src/tilefoundry/analysis/liveness.py @@ -2,7 +2,7 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, replace from tilefoundry.ir.core import Expr, Var from tilefoundry.ir.hir.function import Function @@ -28,13 +28,6 @@ class Liveness: timeline_end: int -@dataclass -class _IntervalState: - value: Expr - defined_at: int - last_used_at: int - - def _free_vars(function: Function) -> tuple[Var, ...]: """Boundary Vars reached as uses but not owned by a structural binding site.""" if function.body is None: @@ -56,7 +49,7 @@ class LivenessVisitor(ExprVisitor[None]): def __init__(self, function: Function) -> None: super().__init__(root_function=function) self._point = -1 - self._states: dict[int, _IntervalState] = {} + self._states: dict[int, LiveInterval] = {} self._definition_order: list[int] = [] for parameter in function.params: self.define(parameter, self.next_event()) @@ -73,7 +66,7 @@ def define(self, value: Expr, point: int) -> None: key = id(value) if key in self._states: raise ValueError(f"liveness: {type(value).__name__} is defined more than once") - self._states[key] = _IntervalState(value, point, point) + self._states[key] = LiveInterval(value, point, point) self._definition_order.append(key) def use(self, value: Expr, point: int) -> None: @@ -81,15 +74,13 @@ def use(self, value: Expr, point: int) -> None: state = self._states.get(id(value)) if state is None: raise ValueError(f"liveness: {type(value).__name__} is used before its definition") - state.last_used_at = max(state.last_used_at, point) + self._states[id(value)] = replace(state, last_used_at=max(state.last_used_at, point)) def finish(self) -> Liveness: """Freeze the definition-ordered result.""" states = (self._states[key] for key in self._definition_order) return Liveness( - intervals=tuple( - LiveInterval(state.value, state.defined_at, state.last_used_at) for state in states - ), + intervals=tuple(states), timeline_end=self._point, ) diff --git a/tests/analysis/test_analyze_at_a_size.py b/tests/analysis/test_analyze_at_a_size.py index fd248ba0..90e7831f 100644 --- a/tests/analysis/test_analyze_at_a_size.py +++ b/tests/analysis/test_analyze_at_a_size.py @@ -54,7 +54,9 @@ INVENTORY = [pytest.param(case, id=case.id) for case in placed_cases()] EXPECTED_MEMORY_PEAKS = { "qwen3_1_7b_pd.PrefillLayer.layer_prefill[ctx_len=128,seq=128]": { + "gmem": 175_514_632, "rmem": 395_264, + "smem": 98_304, }, } From a7313be92ead45b5ab12d1a2016f9daf318dfcfe Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Mon, 21 Sep 2026 01:05:27 +0800 Subject: [PATCH 4/9] refactor(analysis): prepare allocation solver inputs [M1] --- src/tilefoundry/analysis/allocation.py | 20 ++++++++++++ src/tilefoundry/analysis/memory.py | 44 ++++++++++++++++++-------- src/tilefoundry/analysis/scope.py | 7 ++++ 3 files changed, 57 insertions(+), 14 deletions(-) create mode 100644 src/tilefoundry/analysis/allocation.py diff --git a/src/tilefoundry/analysis/allocation.py b/src/tilefoundry/analysis/allocation.py new file mode 100644 index 00000000..34bd6997 --- /dev/null +++ b/src/tilefoundry/analysis/allocation.py @@ -0,0 +1,20 @@ +"""Internal values and solver results for physical memory placement.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from tilefoundry.ir.core import Expr + +from .metadata import ValueLifetime + + +@dataclass(frozen=True) +class AllocationValue: + """Connect one logical expression to its target-aware lifetime.""" + + value: Expr + lifetime: ValueLifetime + + +__all__ = ["AllocationValue"] diff --git a/src/tilefoundry/analysis/memory.py b/src/tilefoundry/analysis/memory.py index 33f90e41..763c4d62 100644 --- a/src/tilefoundry/analysis/memory.py +++ b/src/tilefoundry/analysis/memory.py @@ -32,6 +32,7 @@ from tilefoundry.visitor_registry.contexts import Cost, CostContext, FunctionScope from tilefoundry.visitor_registry.visitors import CostEvaluator +from .allocation import AllocationValue from .errors import AnalysisError from .facts import MemoryHierarchyFacts from .liveness import Liveness, analyze_liveness @@ -320,18 +321,18 @@ def _resident_value_ids(function: Function, liveness: Liveness) -> frozenset[int return frozenset(result) -def _project_value_lifetimes( +def _project_allocation_values( liveness: Liveness, resident_ids: frozenset[int], parameter_ids: frozenset[int], facts: MemoryHierarchyFacts, local: CostContext, -) -> tuple[ValueLifetime, ...]: +) -> tuple[AllocationValue, ...]: """Project structural intervals into the analysed topology window.""" intervals = tuple( interval for interval in liveness.intervals if id(interval.value) in resident_ids ) - result: list[ValueLifetime] = [] + result: list[AllocationValue] = [] labels = value_labels(interval.value for interval in intervals) for label, interval in zip(labels, intervals, strict=True): expr = interval.value @@ -340,13 +341,18 @@ def _project_value_lifetimes( if facts.explicit(memory_level) is None: continue result.append( - ValueLifetime( - binding=label, - memory_level=memory_level, - bytes=amount, - defined_at=interval.defined_at, - last_used_at=(liveness.timeline_end if persistent else interval.last_used_at), - persistent=persistent, + AllocationValue( + expr, + ValueLifetime( + binding=label, + memory_level=memory_level, + bytes=amount, + defined_at=interval.defined_at, + last_used_at=( + liveness.timeline_end if persistent else interval.last_used_at + ), + persistent=persistent, + ), ) ) return tuple(result) @@ -366,13 +372,14 @@ def analyze_value_lifetimes( topology_level=topology_level, topologies=module.effective_topologies(), ) - return _project_value_lifetimes( + projected = _project_allocation_values( liveness, _resident_value_ids(function, liveness), frozenset(id(parameter) for parameter in function.params), facts, local, ) + return tuple(item.lifetime for item in projected) @dataclass @@ -466,11 +473,20 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: communication=_shares(memory_context.shares, tuple(locals_by_unit)), ), ) - lifetimes = analyze_value_lifetimes( - module, - function, + liveness = analyze_liveness(function) + placement = CostContext( + scope=FunctionScope(module, function), topology_level=topology_level, + topologies=topologies, + ) + allocation_values = _project_allocation_values( + liveness, + _resident_value_ids(function, liveness), + frozenset(id(parameter) for parameter in function.params), + facts, + placement, ) + lifetimes = tuple(item.lifetime for item in allocation_values) levels_list: list[MemoryLevelFootprint] = [] for name in sorted({item.memory_level for item in lifetimes} | set(memory_context.totals)): declared = facts.explicit(name) diff --git a/src/tilefoundry/analysis/scope.py b/src/tilefoundry/analysis/scope.py index 5f269462..2264469a 100644 --- a/src/tilefoundry/analysis/scope.py +++ b/src/tilefoundry/analysis/scope.py @@ -55,6 +55,7 @@ class Scope: depth: int domain: isl.set accesses: dict[str, dict[int, tuple[Call, tuple[Access, ...]]]] = field(default_factory=dict) + outputs: dict[str, dict[int, tuple[Call, tuple[Access, ...]]]] = field(default_factory=dict) relations: dict[int, tuple[Call, AccessRelations]] = field(default_factory=dict) refused: dict[str, frozenset[Call]] = field(default_factory=dict) _variance: dict[int, frozenset[int]] = field(default_factory=dict, repr=False) @@ -381,6 +382,12 @@ def record_accesses(expr: Call, scope: Scope) -> None: if access is not None: built.append(access) scope.accesses.setdefault(view, {})[id(expr)] = (expr, tuple(built)) + written: list[Access] = [] + for boundary in local_relations.outputs: + access = _bind_access(expr, expr, boundary, scope, type_ctx, narrow=narrow) + if access is not None: + written.append(access) + scope.outputs.setdefault(view, {})[id(expr)] = (expr, tuple(written)) def record_variance(expr: Expr, operands: tuple[Expr, ...]) -> None: changing: set[int] = set() From 77568b0297b8ffc63609cd69e9ad6f7ec5b34cf1 Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Mon, 21 Sep 2026 01:30:35 +0800 Subject: [PATCH 5/9] fix(analysis): solve loop-aware buffer placement [M2] --- docs/spec/analysis.md | 28 +- src/tilefoundry/analysis/allocation.py | 374 ++++++++++++++++++++++- src/tilefoundry/analysis/memory.py | 60 ++-- src/tilefoundry/analysis/metadata.py | 7 +- src/tilefoundry/analysis/scope.py | 7 +- tests/analysis/test_analysis_families.py | 28 +- 6 files changed, 455 insertions(+), 49 deletions(-) diff --git a/docs/spec/analysis.md b/docs/spec/analysis.md index 5a7c486d..f07092ea 100644 --- a/docs/spec/analysis.md +++ b/docs/spec/analysis.md @@ -272,7 +272,8 @@ class MemoryLevelFootprint: Attributes: memory_level: attribute; The memory level name. - peak_bytes: attribute; The largest simultaneous claim on the level. + peak_bytes: attribute; The solved address high-water mark for gmem/smem, + or the largest single logical value for rmem. persistent_bytes: attribute; The part that cannot be reclaimed. capacity_bytes: attribute; The stated capacity, or None when unknown. """ @@ -333,7 +334,7 @@ class AllocationMetadata: """What showing this function's buffers fit came to. Attributes: - solver_status: attribute; `"optimal"` or `"feasible"`. + solver_status: attribute; `"feasible"` for the first validated placement. """ solver_status: str @@ -376,10 +377,13 @@ anything; it does not say how much, and an Op with no relation fails closed. it empty. Capacity is settled against the authored definition order, which fixes every -buffer's lifetime before any of them is measured, so the only open question is -whether the ones live at once fit together. An arrangement answers that question -without being reported: no address or per-value buffer identity is a conclusion -of this analysis. +buffer's lifetime before any of them is measured. For `gmem` and `smem`, exact +polyhedral access relations may let the solver overlap a dead pointwise operand +with its result or embed an `insert_slice` update in its result. Every logical +SSA box remains in the model, including outside the container's lifetime. The +concrete arrangement is not reported: no address or per-value buffer identity +is a conclusion of this analysis. `rmem` is not address-solved and reports only +the largest single projected logical value. - constraints: - Capacity MUST be settled for the addressable levels `gmem` and `smem` only, @@ -392,12 +396,12 @@ of this analysis. - `allocation` MUST be absent only when no level could be projected against. A function with nothing addressable MUST record a settled `allocation`: the question was asked and there was nothing to decide. An attached - `solver_status` MUST be `"optimal"` or `"feasible"`; a domain that cannot + `solver_status` MUST be `"feasible"`; a domain that cannot fit, cannot be expressed, or does not settle in time MUST raise - `AnalysisError` saying which of the three happened and leave no record. A - domain that fits at once MUST be settled without searching, and one whose - simultaneously live bytes exceed the capacity MUST be refused without - searching. + `AnalysisError` saying which of the three happened and leave no record. The + solver MUST stop at its first feasible assignment rather than spend the + remaining timeout proving a minimum. Its reported peak is that assignment's + actual address high-water mark, not a mathematical optimum. - Every `Spread` MUST state a share for each declared topology level, not only for the level the call selected, and the record MUST name those levels once in `topologies` rather than beside each share. @@ -442,7 +446,7 @@ of this analysis. | `ValueLifetime.last_used_at` | Greatest ordinary-consumer, region-entry, loop-backedge, or region-exit use event; the final timeline event for a parameter. | No | | `ValueLifetime.persistent` | True for parameters and false for body allocations. | No | | `MemoryLevelFootprint.memory_level` | Each storage level with at least one lifetime, sorted by name. | No | -| `MemoryLevelFootprint.peak_bytes` | Largest sum of simultaneously live bytes at that level over the structured SSA event timeline. | No | +| `MemoryLevelFootprint.peak_bytes` | For `gmem` and `smem`, the address high-water mark of the first feasible whole-Function placement. Exact pointwise relations and exact `insert_slice` partitions may permit overlap; widened or unknown relations do not. For `rmem`, the largest single projected logical value, without address placement or cross-value summation. | No | | `MemoryLevelFootprint.persistent_bytes` | Sum of persistent lifetimes at that level. | No | | `MemoryLevelFootprint.capacity_bytes` | Capacity of the matching explicit level, or `None` when it is unknown or undeclared. | `MemoryHierarchyFacts.explicit_levels[].capacity_bytes` | | `MemoryMetadata.footprint` | One `MemoryLevelFootprint` per occupied storage level. | As above | diff --git a/src/tilefoundry/analysis/allocation.py b/src/tilefoundry/analysis/allocation.py index 34bd6997..4ed3f335 100644 --- a/src/tilefoundry/analysis/allocation.py +++ b/src/tilefoundry/analysis/allocation.py @@ -2,11 +2,29 @@ from __future__ import annotations +from collections import defaultdict from dataclasses import dataclass +from itertools import combinations +from typing import Protocol -from tilefoundry.ir.core import Expr +import isl +from ortools.sat.python import cp_model +from tilefoundry.ir.core import Call, Expr +from tilefoundry.ir.hir.loop_region import LoopRegion +from tilefoundry.ir.hir.tensor.insert_slice import InsertSlice +from tilefoundry.ir.hir.tensor.reshape import Reshape +from tilefoundry.ir.hir.tensor.slice import Slice + +from .errors import AnalysisError from .metadata import ValueLifetime +from .scope import Access, Scope, walk_scopes + + +class _MemoryOptions(Protocol): + timeout_seconds: float + workers: int + random_seed: int @dataclass(frozen=True) @@ -17,4 +35,356 @@ class AllocationValue: lifetime: ValueLifetime -__all__ = ["AllocationValue"] +@dataclass(frozen=True) +class AllocationResult: + """The first physical placement found for one memory level.""" + + peak_bytes: int + solver_status: str + + +def _base_value(value: Expr) -> Expr: + """Return the material allocation below non-material tensor views.""" + while isinstance(value, Call) and isinstance(value.target, (Slice, Reshape)): + value = value.args[0] + return value + + +def _coverage(accesses: tuple[Access, ...]) -> isl.set | None: + """Union the call coordinates on which exact accesses reach one buffer.""" + if not accesses or any(not access.exact for access in accesses): + return None + result = accesses[0].relation.domain() + for access in accesses[1:]: + result = result.union(access.relation.domain()) + return result.coalesce() + + +def _access_relation(accesses: tuple[Access, ...]) -> isl.map | None: + """Union exact accesses to one buffer without discarding their maps.""" + if not accesses or any(not access.exact for access in accesses): + return None + result = accesses[0].relation + for access in accesses[1:]: + result = result.union(access.relation) + return result.coalesce() + + +def _proved_overlap_groups( + values: tuple[AllocationValue, ...], root: Scope +) -> tuple[tuple[int, ...], ...]: + """Prove exact pointwise ties and dynamic-update embedded views. + + This is scheme B: the exact per-iteration offset remains in ISL. The CP + model only receives the weaker fact that a group can occupy one + result-sized allocation. + """ + by_expr = {id(item.value): index for index, item in enumerate(values)} + carried_ids = { + id(carried) + for scope in walk_scopes(root) + if isinstance(scope.owner, LoopRegion) + for carried in scope.owner.carried_args + } + groups: list[tuple[int, ...]] = [] + seen: set[tuple[int, ...]] = set() + for scope in walk_scopes(root): + for call_id, (call, inputs) in scope.accesses.get("narrow", {}).items(): + if call_id not in by_expr: + continue + recorded_output = scope.outputs.get("narrow", {}).get(call_id) + if recorded_output is None or recorded_output[0] is not call: + continue + outputs = recorded_output[1] + if not outputs or any(not access.exact for access in outputs): + continue + + result_index = by_expr[call_id] + result = values[result_index] + output_coverage = _coverage(outputs) + output_relation = _access_relation(outputs) + by_buffer: dict[int, list[Access]] = defaultdict(list) + for access in inputs: + by_buffer[id(access.buffer)].append(access) + + if output_coverage is not None and output_relation is not None: + for buffer_id, accesses in by_buffer.items(): + input_index = by_expr.get(buffer_id) + if input_index is None or input_index == result_index: + continue + source = values[input_index] + if source.lifetime.persistent: + continue + if source.lifetime.bytes != result.lifetime.bytes: + continue + if source.lifetime.last_used_at > result.lifetime.defined_at: + continue + input_relation = _access_relation(tuple(accesses)) + try: + pointwise = input_relation is not None and input_relation.is_equal( + output_relation + ) + except isl.Error: + pointwise = False + group = (result_index, input_index) + if pointwise and group not in seen: + seen.add(group) + groups.append(group) + + if not isinstance(call.target, InsertSlice): + continue + + dst = _base_value(call.args[0]) + update = _base_value(call.args[1]) + member_ids = (id(call), id(dst), id(update)) + if any(member_id not in by_expr for member_id in member_ids): + continue + indices = tuple(dict.fromkeys(by_expr[member_id] for member_id in member_ids)) + if len(indices) != 3: + continue + result_index, dst_index, update_index = indices + result = values[result_index] + destination = values[dst_index] + patch = values[update_index] + if destination.lifetime.persistent or patch.lifetime.persistent: + continue + if ( + destination.lifetime.last_used_at > result.lifetime.defined_at + and id(destination.value) not in carried_ids + ): + continue + if destination.lifetime.bytes != result.lifetime.bytes: + continue + if patch.lifetime.bytes > result.lifetime.bytes: + continue + if patch.lifetime.last_used_at > result.lifetime.defined_at: + continue + + dst_coverage = _coverage(tuple(by_buffer.get(id(dst), ()))) + update_coverage = _coverage(tuple(by_buffer.get(id(update), ()))) + written_coverage = _coverage(outputs) + if any( + coverage is None for coverage in (dst_coverage, update_coverage, written_coverage) + ): + continue + try: + partitioned = update_coverage.is_equal( + written_coverage + ) and dst_coverage.is_disjoint(written_coverage) + except isl.Error: + partitioned = False + if partitioned and indices not in seen: + seen.add(indices) + groups.append(indices) + return tuple(groups) + + +def _lifetimes_overlap(left: AllocationValue, right: AllocationValue) -> bool: + """Whether two closed structured-SSA intervals share an event.""" + return max(left.lifetime.defined_at, right.lifetime.defined_at) <= min( + left.lifetime.last_used_at, right.lifetime.last_used_at + ) + + +def _placement_hint( + values: tuple[AllocationValue, ...], + groups: tuple[tuple[int, ...], ...], + limit: int, +) -> tuple[tuple[bool, ...], tuple[int, ...], int] | None: + """Construct one complete feasible suggestion without deciding the model. + + Components exist only while building the hint. The CP model still contains + every logical box and independently validates or rejects every suggested + overlap. + """ + parent = list(range(len(values))) + + def find(index: int) -> int: + while parent[index] != index: + parent[index] = parent[parent[index]] + index = parent[index] + return index + + def members(root: int) -> set[int]: + return {index for index in range(len(values)) if find(index) == root} + + allowed_pairs: set[tuple[int, int]] = set() + selected: list[bool] = [] + for group in groups: + roots = {find(index) for index in group} + merged = set().union(*(members(root) for root in roots)) + group_pairs = { + tuple(sorted((left, right))) + for left, right in combinations(group, 2) + if _lifetimes_overlap(values[left], values[right]) + } + permitted = allowed_pairs | group_pairs + compatible = all( + not _lifetimes_overlap(values[left], values[right]) or (left, right) in permitted + for left, right in combinations(sorted(merged), 2) + ) + selected.append(compatible) + if not compatible: + continue + root = min(roots) + for other in roots: + parent[find(other)] = root + allowed_pairs.update(group_pairs) + + components: dict[int, set[int]] = defaultdict(set) + for index in range(len(values)): + components[find(index)].add(index) + ordered = sorted( + components.values(), + key=lambda component: ( + min(values[index].lifetime.defined_at for index in component), + -max(values[index].lifetime.bytes for index in component), + ), + ) + placed: list[tuple[set[int], int, int]] = [] + addresses = [0] * len(values) + peak = 0 + for component in ordered: + size = max(values[index].lifetime.bytes for index in component) + + def conflicts(other: set[int]) -> bool: + return any( + _lifetimes_overlap(values[left], values[right]) + for left in component + for right in other + ) + + blocked = tuple(item for item in placed if conflicts(item[0])) + candidates = sorted({0, *(address + held for _other, address, held in blocked)}) + address = next( + candidate + for candidate in candidates + if all( + candidate + size <= other_address or other_address + other_size <= candidate + for _other, other_address, other_size in blocked + ) + ) + if address + size > limit: + return None + for index in component: + addresses[index] = address + placed.append((component, address, size)) + peak = max(peak, address + size) + return tuple(selected), tuple(addresses), peak + + +def solve_allocation( + memory_level: str, + values: tuple[AllocationValue, ...], + root: Scope, + *, + capacity_bytes: int | None, + options: _MemoryOptions, +) -> AllocationResult: + """Return the first feasible whole-function placement for one level.""" + if any(item.lifetime.memory_level != memory_level for item in values): + raise ValueError("allocation values must all belong to the requested memory level") + if not values: + return AllocationResult(0, "optimal") + + largest = max(item.lifetime.bytes for item in values) + total = sum(item.lifetime.bytes for item in values) + limit = total if capacity_bytes is None else capacity_bytes + if largest > limit: + raise AnalysisError( + f"allocation: a value needs {largest} B in {memory_level}, " + f"which exceeds its {limit} B placement limit" + ) + + model = cp_model.CpModel() + peak = model.new_int_var(largest, limit, f"{memory_level}_peak") + addresses = [ + model.new_int_var(0, limit - item.lifetime.bytes, f"address_{index}") + for index, item in enumerate(values) + ] + for address, item in zip(addresses, values, strict=True): + model.add(address + item.lifetime.bytes <= peak) + + persistent_end = 0 + for address, item in zip(addresses, values, strict=True): + if not item.lifetime.persistent: + continue + model.add(address == persistent_end) + persistent_end += item.lifetime.bytes + for address, item in zip(addresses, values, strict=True): + if not item.lifetime.persistent: + model.add(address >= persistent_end) + + groups = _proved_overlap_groups(values, root) + choices_by_pair: dict[tuple[int, int], list[cp_model.IntVar]] = defaultdict(list) + choices: list[tuple[cp_model.IntVar, tuple[int, ...]]] = [] + for group_index, members in enumerate(groups): + selected = model.new_bool_var(f"embedded_{group_index}") + choices.append((selected, members)) + container = members[0] + for member in members[1:]: + model.add(addresses[member] >= addresses[container]).only_enforce_if(selected) + model.add( + addresses[member] + values[member].lifetime.bytes + <= addresses[container] + values[container].lifetime.bytes + ).only_enforce_if(selected) + for left, right in combinations(members, 2): + if _lifetimes_overlap(values[left], values[right]): + choices_by_pair[tuple(sorted((left, right)))].append(selected) + + order_choices: dict[tuple[int, int], tuple[cp_model.IntVar, cp_model.IntVar]] = {} + for left, right in combinations(range(len(values)), 2): + if not _lifetimes_overlap(values[left], values[right]): + continue + if values[left].lifetime.persistent or values[right].lifetime.persistent: + continue + left_before = model.new_bool_var(f"before_{left}_{right}") + right_before = model.new_bool_var(f"before_{right}_{left}") + order_choices[(left, right)] = (left_before, right_before) + model.add( + addresses[left] + values[left].lifetime.bytes <= addresses[right] + ).only_enforce_if(left_before) + model.add( + addresses[right] + values[right].lifetime.bytes <= addresses[left] + ).only_enforce_if(right_before) + model.add_bool_or( + left_before, + right_before, + *choices_by_pair.get((left, right), ()), + ) + + hint = _placement_hint(values, groups, limit) + if hint is not None: + selected_hints, address_hints, peak_hint = hint + model.add(peak <= peak_hint) + for (selected, _members), suggested in zip(choices, selected_hints, strict=True): + model.add_hint(selected, int(suggested)) + for address, suggested in zip(addresses, address_hints, strict=True): + model.add_hint(address, suggested) + for (left, right), (left_before, right_before) in order_choices.items(): + left_end = address_hints[left] + values[left].lifetime.bytes + right_end = address_hints[right] + values[right].lifetime.bytes + model.add_hint(left_before, int(left_end <= address_hints[right])) + model.add_hint(right_before, int(right_end <= address_hints[left])) + model.add_hint(peak, peak_hint) + + solver = cp_model.CpSolver() + solver.parameters.max_time_in_seconds = options.timeout_seconds + solver.parameters.num_search_workers = options.workers + solver.parameters.random_seed = options.random_seed + solver.parameters.stop_after_first_solution = True + solver.parameters.search_branching = cp_model.HINT_SEARCH + status = solver.solve(model) + if status not in (cp_model.OPTIMAL, cp_model.FEASIBLE): + status_name = solver.status_name(status).lower() + raise AnalysisError( + f"allocation: no feasible {memory_level} placement was found ({status_name})" + ) + actual_peak = max( + solver.value(address) + item.lifetime.bytes + for address, item in zip(addresses, values, strict=True) + ) + return AllocationResult(actual_peak, "feasible") + + +__all__ = ["AllocationResult", "AllocationValue", "solve_allocation"] diff --git a/src/tilefoundry/analysis/memory.py b/src/tilefoundry/analysis/memory.py index 763c4d62..2d5d68c4 100644 --- a/src/tilefoundry/analysis/memory.py +++ b/src/tilefoundry/analysis/memory.py @@ -32,7 +32,7 @@ from tilefoundry.visitor_registry.contexts import Cost, CostContext, FunctionScope from tilefoundry.visitor_registry.visitors import CostEvaluator -from .allocation import AllocationValue +from .allocation import AllocationValue, solve_allocation from .errors import AnalysisError from .facts import MemoryHierarchyFacts from .liveness import Liveness, analyze_liveness @@ -487,41 +487,63 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: placement, ) lifetimes = tuple(item.lifetime for item in allocation_values) + solver_options = ( + context.options if isinstance(context.options, MemoryOptions) else MemoryOptions() + ) + solver_statuses: list[str] = [] levels_list: list[MemoryLevelFootprint] = [] for name in sorted({item.memory_level for item in lifetimes} | set(memory_context.totals)): declared = facts.explicit(name) - rows = [item for item in lifetimes if item.memory_level == name] - peak = 0 - end = max((item.last_used_at for item in rows), default=-1) - for point in range(end + 1): - peak = max( - peak, - sum(item.bytes for item in rows if item.defined_at <= point <= item.last_used_at), + values = tuple(item for item in allocation_values if item.lifetime.memory_level == name) + rows = [item.lifetime for item in values] + capacity = declared.capacity_bytes if declared is not None else None + for item in rows: + if capacity is not None and item.bytes > capacity: + raise AnalysisError( + f"function {function.name!r}: value {item.binding!r} needs " + f"{item.bytes} B in {item.memory_level}, which exceeds the " + f"{capacity} B the target states for that level" + ) + if name in (str(StorageKind.GMEM), str(StorageKind.SMEM)) and values: + solved = solve_allocation( + name, + values, + context.root, + capacity_bytes=capacity, + options=solver_options, ) + peak = solved.peak_bytes + solver_statuses.append(solved.solver_status) + elif name == str(StorageKind.RMEM): + peak = max((item.bytes for item in rows), default=0) + else: + peak = 0 + end = max((item.last_used_at for item in rows), default=-1) + for point in range(end + 1): + peak = max( + peak, + sum( + item.bytes for item in rows if item.defined_at <= point <= item.last_used_at + ), + ) levels_list.append( MemoryLevelFootprint( memory_level=name, peak_bytes=peak, persistent_bytes=sum(item.bytes for item in rows if item.persistent), - capacity_bytes=declared.capacity_bytes if declared is not None else None, + capacity_bytes=capacity, ) ) levels = tuple(levels_list) - for item in lifetimes: - declared = facts.explicit(item.memory_level) - capacity = None if declared is None else declared.capacity_bytes - if capacity is not None and item.bytes > capacity: - raise AnalysisError( - f"function {function.name!r}: value {item.binding!r} needs " - f"{item.bytes} B in {item.memory_level}, which exceeds the " - f"{capacity} B the target states for that level" - ) + allocation = None + if solver_statuses: + allocation = AllocationMetadata(solver_status="feasible") attach( function, MemoryMetadata( footprint=levels, lifetimes=lifetimes, - allocation=AllocationMetadata(solver_status="optimal"), + allocation=allocation, ), ) diff --git a/src/tilefoundry/analysis/metadata.py b/src/tilefoundry/analysis/metadata.py index 309298a2..d4f55522 100644 --- a/src/tilefoundry/analysis/metadata.py +++ b/src/tilefoundry/analysis/metadata.py @@ -123,7 +123,7 @@ def per_unit(self) -> tuple[tuple[str, TrafficBytes], ...]: @dataclass(frozen=True) class MemoryLevelFootprint: - """How much of one memory level a function needs at its peak. + """One level's solved high-water mark or largest logical value. ``persistent_bytes`` is the part that cannot be reclaimed within the function, so it is the floor the peak can never fall below. @@ -199,10 +199,11 @@ def level(self) -> str: @dataclass(frozen=True) class AllocationMetadata: - """What showing this function's buffers fit took. + """What showing this function's addressable buffers fit took. Where any of them would sit is the solver's business and appears nowhere - here. What a reader can act on is whether the question was settled. + here. ``feasible`` means the first validated placement was returned without + claiming that its high-water mark is minimal. """ solver_status: str diff --git a/src/tilefoundry/analysis/scope.py b/src/tilefoundry/analysis/scope.py index 2264469a..8c9607ed 100644 --- a/src/tilefoundry/analysis/scope.py +++ b/src/tilefoundry/analysis/scope.py @@ -43,6 +43,7 @@ class Access: relation: isl.map buffer: Expr + exact: bool = True @dataclass(eq=False) @@ -278,6 +279,7 @@ def _bind_access( narrow: bool, ) -> Access | None: relation = relation_of(boundary.pattern) + exact = True loops = [] cursor = scope while cursor is not None: @@ -302,6 +304,7 @@ def _bind_access( term = None if term is None: term = _widest_allowed(relation, name, operand.type) + exact = False if term is None: relation = relation.project_out(isl.dim_type.PARAM, 0, 1) continue @@ -337,7 +340,9 @@ def placed(kind: str, sign: int, constant: int) -> isl.constraint: folded = renaming_relation(operand, ctx, stated=scope.stated_relations(operand, ctx)) relation = relation.apply_range(relation_of(folded)) operand = operand.args[0] - return Access(relation, operand) + if relation.dim(isl.dim_type.PARAM): + exact = False + return Access(relation, operand, exact) def build_scopes( diff --git a/tests/analysis/test_analysis_families.py b/tests/analysis/test_analysis_families.py index 2109a3e2..854293bd 100644 --- a/tests/analysis/test_analysis_families.py +++ b/tests/analysis/test_analysis_families.py @@ -38,6 +38,7 @@ _local_duration_ns, ) from tilefoundry.analysis.errors import AnalysisError +from tilefoundry.analysis.memory import MemoryOptions from tilefoundry.dsl import ConstTensor, DimVar, Mesh, Tensor, Topology, tf from tilefoundry.ir.core import ( Call, @@ -339,13 +340,11 @@ def test_a_matmul_counts_its_rows_once_whichever_axis_the_mesh_split() -> None: def test_a_program_whose_buffers_have_nowhere_to_sit_is_refused() -> None: """Placing the buffers is what makes the rest of the answer worth having. - One shared tile of this program is twice what the machine states for that - level, and no ordering makes room for it: a value that cannot be placed at - all is refused, because memory decides this and everything downstream reads - what it decided. Restating the capacity changes only the answer: the - lifetimes are the same either way. What is refused is one value against the - capacity and not the working set, so the same program on the unrestated - machine keeps two tiles that each fit live at once and is answered. + One shared tile of this program is twice what the tight machine states for + that level, so the value cannot be placed at all. On the real machine the + pointwise add result reuses its dead input: two logical lifetimes overlap at + the call event, but their exact access maps prove one physical placement. + Restating capacity changes only the answer, not those logical lifetimes. """ tight = replace(_SharedTile, target=_TightShared("nvidia.h200_sxm")) split = next(function for function in tight.functions if function.name == "split") @@ -359,9 +358,15 @@ def test_a_program_whose_buffers_have_nowhere_to_sit_is_refused() -> None: unrestated = next(item for item in _SharedTile.functions if item.name == "split") held = get_metadata( - analyze(_SharedTile, unrestated, analysis="memory").function, MemoryMetadata + analyze( + _SharedTile, + unrestated, + analysis="memory", + options=MemoryOptions(timeout_seconds=1.0), + ).function, + MemoryMetadata, ).footprint - assert next(item.peak_bytes for item in held if item.level == "smem") == 422_400 + assert next(item.peak_bytes for item in held if item.level == "smem") == 211_200 fits = analyze( roomy, @@ -370,7 +375,7 @@ def test_a_program_whose_buffers_have_nowhere_to_sit_is_refused() -> None: ) summary = get_metadata(fits.function, PerformanceSummaryMetadata) assert summary is not None - assert get_metadata(fits.function, MemoryMetadata).allocation.solver_status == "optimal" + assert get_metadata(fits.function, MemoryMetadata).allocation.solver_status == "feasible" assert summary.timeline.end_ns > 0 wider = replace(_SharedTile, target=_RoomierShared("nvidia.h200_sxm")) @@ -407,8 +412,7 @@ def test_a_price_is_refused_where_the_machine_states_no_rate_to_pay_it_at() -> N work = next( record for record in records - if record is not None - and any(spread.per_unit[0] for _name, spread in record.flops.kinds) + if record is not None and any(spread.per_unit[0] for _name, spread in record.flops.kinds) ) with pytest.raises( From 39ff5c907e8e6aa9b2f95c74dcbf63f3d832dc9d Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Mon, 21 Sep 2026 10:13:00 +0800 Subject: [PATCH 6/9] fix(analysis): report placement capacity errors [M2] --- docs/spec/analysis.md | 32 ++++++++++-------- src/tilefoundry/analysis/allocation.py | 8 +---- src/tilefoundry/analysis/memory.py | 15 ++++----- src/tilefoundry/analysis/metadata.py | 9 ++--- src/tilefoundry/inspection/analysis_report.py | 10 +++--- src/tilefoundry/inspection/values.py | 13 ++++++-- tests/analysis/test_analysis_families.py | 33 ++++++++++++------- 7 files changed, 68 insertions(+), 52 deletions(-) diff --git a/docs/spec/analysis.md b/docs/spec/analysis.md index f07092ea..83396a85 100644 --- a/docs/spec/analysis.md +++ b/docs/spec/analysis.md @@ -345,12 +345,14 @@ class MemoryMetadata(IRMetadata): Attributes: footprint: attribute; One row per level the function places values in. lifetimes: attribute; One entry per value residency. + errors: attribute; Solved placement peaks that exceed stated capacity. advisories: attribute; Capacity findings that do not invalidate the program. allocation: attribute; What showing the addressable buffers fit came to. """ footprint: tuple[MemoryLevelFootprint, ...] = () lifetimes: tuple[ValueLifetime, ...] = () + errors: tuple[str, ...] = () advisories: tuple[str, ...] = () allocation: AllocationMetadata | None = None ``` @@ -386,7 +388,7 @@ is a conclusion of this analysis. `rmem` is not address-solved and reports only the largest single projected logical value. - constraints: - - Capacity MUST be settled for the addressable levels `gmem` and `smem` only, + - Placement MUST be settled for the addressable levels `gmem` and `smem` only, once per capacity domain that holds a buffer -- the whole target for a level owned target-wide, one per owning position otherwise -- with two buffers in one domain never live in the same bytes at once. Residency at another level @@ -396,12 +398,13 @@ the largest single projected logical value. - `allocation` MUST be absent only when no level could be projected against. A function with nothing addressable MUST record a settled `allocation`: the question was asked and there was nothing to decide. An attached - `solver_status` MUST be `"feasible"`; a domain that cannot - fit, cannot be expressed, or does not settle in time MUST raise - `AnalysisError` saying which of the three happened and leave no record. The + `solver_status` MUST be `"feasible"`; a domain that cannot be expressed or + does not settle in time MUST raise `AnalysisError` and leave no record. The solver MUST stop at its first feasible assignment rather than spend the remaining timeout proving a minimum. Its reported peak is that assignment's - actual address high-water mark, not a mathematical optimum. + actual address high-water mark, not a mathematical optimum. Capacity MUST + NOT restrict that address space: after solving, a high-water mark above + capacity MUST add a non-fatal `errors` entry and preserve the complete result. - Every `Spread` MUST state a share for each declared topology level, not only for the level the call selected, and the record MUST name those levels once in `topologies` rather than beside each share. @@ -451,7 +454,8 @@ the largest single projected logical value. | `MemoryLevelFootprint.capacity_bytes` | Capacity of the matching explicit level, or `None` when it is unknown or undeclared. | `MemoryHierarchyFacts.explicit_levels[].capacity_bytes` | | `MemoryMetadata.footprint` | One `MemoryLevelFootprint` per occupied storage level. | As above | | `MemoryMetadata.lifetimes` | Every value residency except a `Reshape` or a `Transpose`, each of which describes bytes its operand already holds. | As above | -| `MemoryMetadata.advisories` | Explicit peak overflow, cache/shared-capacity division, and same-scope authored-loop access-footprint findings. | `MemoryHierarchyFacts` | +| `MemoryMetadata.errors` | One non-fatal error for each solved `gmem` or `smem` placement whose `peak_bytes` exceeds stated `capacity_bytes`; this includes a single value larger than capacity. | `MemoryHierarchyFacts.explicit_levels[].capacity_bytes` | +| `MemoryMetadata.advisories` | Cache/shared-capacity division and same-scope authored-loop access-footprint findings. | `MemoryHierarchyFacts` | | `TrafficMetadata.storage` | One occurrence's per-boundary movement asked of the Op's access relations, charged to the storage levels its operand Types name. The total is asked in the whole program's window and each level's share in that level's, over Types projected through the authored `Split`s at or coarser than it. On a Function, summed over every reachable occurrence, each counted as often as its authored loops repeat it. A Type with leaves at several levels keeps those leaf bytes separate. A `UMAT` leaf has no residency of its own: when it appears in `Call.args`, charge its own bytes at the target's established `rmem` materialization level; when it appears only in an Op attribute, charge nothing. A Function Call takes the callee's grouped total. | No; projection reads resolved Mesh and effective Module topology extents. | | `TrafficMetadata.communication` | What a Reshard sends off the unit it was on, when the shards on its two sides differ across a mesh axis that level owns. Zero where they agree. | No; the share each unit keeps follows from the mesh extents the shards name. | | `TrafficMetadata.operands` | One occurrence's per-boundary movement in order `(*call.args, call)`, the same relation-derived amounts `storage` groups. Empty on a Function and on a Function Call, neither of which has a split. | No | @@ -555,16 +559,17 @@ line per advisory: ```text traffic traffic=:r/w@total,r/w@[,...][;:...] peak-footprint=:[,:...] +error="" advisory="" ``` -An empty footprint states the family name alone; each advisory is its own line -and is quoted and escaped +An empty footprint states the family name alone; each error and advisory is its +own line and is quoted and escaped ([inspection §2.8](./inspection.md#28-record-comment-forms)). The record's own comment form projects the footprint it holds, and `lifetimes` is read from JSON: ```text -memory peak=:[,...] persistent= advisories= +memory peak=:[,...] persistent= errors= advisories= ``` Every measured Call also receives a `traffic` annotation, whose `operands` split @@ -588,6 +593,7 @@ attached only to the Function. Its full JSON projection is under "lifetimes": [{"binding": , "memory_level": , "bytes": , "defined_at": , "last_used_at": , "persistent": }, ...], + "errors": [, ...], "advisories": [, ...]} ``` @@ -625,10 +631,10 @@ attached only to the Function. Its full JSON projection is under owner. - Analysis MUST NOT infer memory ownership from a storage level's name or capacity scope. - - One value exceeding an explicit level's capacity MUST raise `AnalysisError`. - An aggregate explicit peak or an authored-loop access footprint exceeding an - implicit cache capacity MUST instead produce an advisory and MUST NOT fail - the call. + - A solved explicit-level peak exceeding capacity, whether from one value or + the aggregate placement, MUST produce a report `error` and MUST NOT fail the + call. An authored-loop access footprint exceeding an implicit cache capacity + MUST instead produce an advisory. #### 1.2.3 `roofline` diff --git a/src/tilefoundry/analysis/allocation.py b/src/tilefoundry/analysis/allocation.py index 4ed3f335..26423583 100644 --- a/src/tilefoundry/analysis/allocation.py +++ b/src/tilefoundry/analysis/allocation.py @@ -278,7 +278,6 @@ def solve_allocation( values: tuple[AllocationValue, ...], root: Scope, *, - capacity_bytes: int | None, options: _MemoryOptions, ) -> AllocationResult: """Return the first feasible whole-function placement for one level.""" @@ -289,12 +288,7 @@ def solve_allocation( largest = max(item.lifetime.bytes for item in values) total = sum(item.lifetime.bytes for item in values) - limit = total if capacity_bytes is None else capacity_bytes - if largest > limit: - raise AnalysisError( - f"allocation: a value needs {largest} B in {memory_level}, " - f"which exceeds its {limit} B placement limit" - ) + limit = total model = cp_model.CpModel() peak = model.new_int_var(largest, limit, f"{memory_level}_peak") diff --git a/src/tilefoundry/analysis/memory.py b/src/tilefoundry/analysis/memory.py index 2d5d68c4..f525458b 100644 --- a/src/tilefoundry/analysis/memory.py +++ b/src/tilefoundry/analysis/memory.py @@ -497,19 +497,11 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: values = tuple(item for item in allocation_values if item.lifetime.memory_level == name) rows = [item.lifetime for item in values] capacity = declared.capacity_bytes if declared is not None else None - for item in rows: - if capacity is not None and item.bytes > capacity: - raise AnalysisError( - f"function {function.name!r}: value {item.binding!r} needs " - f"{item.bytes} B in {item.memory_level}, which exceeds the " - f"{capacity} B the target states for that level" - ) if name in (str(StorageKind.GMEM), str(StorageKind.SMEM)) and values: solved = solve_allocation( name, values, context.root, - capacity_bytes=capacity, options=solver_options, ) peak = solved.peak_bytes @@ -535,6 +527,12 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: ) ) levels = tuple(levels_list) + errors = tuple( + f"{item.memory_level} placement peak {item.peak_bytes} B exceeds " + f"capacity {item.capacity_bytes} B" + for item in levels + if item.exceeds_capacity + ) allocation = None if solver_statuses: allocation = AllocationMetadata(solver_status="feasible") @@ -543,6 +541,7 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: MemoryMetadata( footprint=levels, lifetimes=lifetimes, + errors=errors, allocation=allocation, ), ) diff --git a/src/tilefoundry/analysis/metadata.py b/src/tilefoundry/analysis/metadata.py index d4f55522..f63f369b 100644 --- a/src/tilefoundry/analysis/metadata.py +++ b/src/tilefoundry/analysis/metadata.py @@ -213,10 +213,10 @@ class AllocationMetadata: class MemoryMetadata(IRMetadata): """Record one function's memory behavior against a target hierarchy. - Function attachment reflects that peaks span all live ranges. Advisories - report cache working-set and order-dependent peak overflow; only a single - value exceeding an addressable level is an error because no schedule can - place it. + Function attachment reflects that peaks span all live ranges. ``errors`` + reports a solved placement whose high-water exceeds stated capacity without + suppressing the rest of the analysis result. Advisories carry lower-severity + capacity findings. ``allocation`` is absent when the function has no addressable buffer to place at the level being analysed, which is a different answer from having @@ -225,6 +225,7 @@ class MemoryMetadata(IRMetadata): footprint: tuple[MemoryLevelFootprint, ...] = () lifetimes: tuple[ValueLifetime, ...] = () + errors: tuple[str, ...] = () advisories: tuple[str, ...] = () allocation: "AllocationMetadata | None" = None diff --git a/src/tilefoundry/inspection/analysis_report.py b/src/tilefoundry/inspection/analysis_report.py index 7dfc8a00..8510972d 100644 --- a/src/tilefoundry/inspection/analysis_report.py +++ b/src/tilefoundry/inspection/analysis_report.py @@ -23,6 +23,7 @@ from tilefoundry.inspection.python_printer import HirPrinter, PythonPrintOptions from tilefoundry.inspection.values import ( AdvisorySummary, + ErrorSummary, MemorySummary, PerformanceSummaryView, Prose, @@ -49,9 +50,7 @@ def selected_types(result: AnalysisResult) -> tuple[type[IRMetadata], ...]: return _selected_types(result.module, result.analyses, result.metadata_types) -def render_analysis( - result: AnalysisResult, *, operands: bool = False -) -> AnalysisRendering: +def render_analysis(result: AnalysisResult, *, operands: bool = False) -> AnalysisRendering: """Render one result once for both annotated source and report data.""" selected_types_ = selected_types(result) rendered = HirPrinter().render( @@ -98,9 +97,7 @@ def _summary( function=data["function"], topology=data["topology"] or "none", ), - ReportSelection( - requested=tuple(data["requested"]), executed=tuple(data["executed"]) - ), + ReportSelection(requested=tuple(data["requested"]), executed=tuple(data["executed"])), ] if "totals" in data and "compute-cost" in data["executed"]: views.append(get_metadata(function, ComputeCostMetadata) or ComputeCostMetadata()) @@ -110,6 +107,7 @@ def _summary( memory = get_metadata(function, MemoryMetadata) views.append(MemorySummary(peak_footprint(memory))) if MemoryMetadata in selected: + views.extend(ErrorSummary(Prose(note)) for note in memory.errors) views.extend(AdvisorySummary(Prose(note)) for note in memory.advisories) if "roofline" in function_records: views.append(get_metadata(function, RooflineMetadata)) diff --git a/src/tilefoundry/inspection/values.py b/src/tilefoundry/inspection/values.py index 79d8cc8b..09b72b44 100644 --- a/src/tilefoundry/inspection/values.py +++ b/src/tilefoundry/inspection/values.py @@ -51,6 +51,11 @@ class AdvisorySummary(IRMetadata): text: Prose +@dataclass(frozen=True) +class ErrorSummary(IRMetadata): + text: Prose + + @dataclass(frozen=True) class PerformanceSummaryView(IRMetadata): root: str = "" @@ -127,8 +132,7 @@ def _spread(self, spread, topologies, *, logical=True): def _breakdown(self, held, topologies, *, logical=True): return { - kind: self._spread(spread, topologies, logical=logical) - for kind, spread in held.kinds + kind: self._spread(spread, topologies, logical=logical) for kind, spread in held.kinds } def print_ComputeCostMetadata(self, record, **_): @@ -158,6 +162,7 @@ def print_MemoryMetadata(self, record, **_): ( ("peak", {item.memory_level: item.peak_bytes for item in record.footprint}), ("persistent", sum(item.persistent_bytes for item in record.footprint), 0), + ("errors", len(record.errors), 0), ("advisories", len(record.advisories), 0), ), ) @@ -217,6 +222,9 @@ def print_MemorySummary(self, record, **_): def print_AdvisorySummary(self, record, **_): return self._single("advisory", record.text) + def print_ErrorSummary(self, record, **_): + return self._single("error", record.text) + def print_PerformanceSummaryView(self, record, **_): return self._record( "performance", @@ -256,6 +264,7 @@ def peak_footprint(record): "ReportSelection", "MemorySummary", "AdvisorySummary", + "ErrorSummary", "PerformanceSummaryView", "peak_footprint", ] diff --git a/tests/analysis/test_analysis_families.py b/tests/analysis/test_analysis_families.py index 854293bd..700be7fa 100644 --- a/tests/analysis/test_analysis_families.py +++ b/tests/analysis/test_analysis_families.py @@ -40,6 +40,7 @@ from tilefoundry.analysis.errors import AnalysisError from tilefoundry.analysis.memory import MemoryOptions from tilefoundry.dsl import ConstTensor, DimVar, Mesh, Tensor, Topology, tf +from tilefoundry.inspection.analysis_report import render_analysis, render_text from tilefoundry.ir.core import ( Call, get_metadata, @@ -337,24 +338,30 @@ def test_a_matmul_counts_its_rows_once_whichever_axis_the_mesh_split() -> None: assert per_layout["last_axis"] == per_layout["strip_major"] -def test_a_program_whose_buffers_have_nowhere_to_sit_is_refused() -> None: - """Placing the buffers is what makes the rest of the answer worth having. +def test_a_program_whose_peak_exceeds_capacity_reports_an_error() -> None: + """A placement error is report data rather than an aborted analysis. One shared tile of this program is twice what the tight machine states for - that level, so the value cannot be placed at all. On the real machine the - pointwise add result reuses its dead input: two logical lifetimes overlap at - the call event, but their exact access maps prove one physical placement. - Restating capacity changes only the answer, not those logical lifetimes. + that level. The solver still places it and its pointwise add result in one + buffer, then reports that solved high-water against capacity. Restating + capacity changes only the error, not the logical lifetimes or placement. """ tight = replace(_SharedTile, target=_TightShared("nvidia.h200_sxm")) split = next(function for function in tight.functions if function.name == "split") roomy = replace(_SharedTile, target=_RoomyShared("nvidia.h200_sxm")) - refusal = r"value 'v\d+:\d+' needs 211200 B in smem, which exceeds the 105600 B" - with pytest.raises(AnalysisError, match=refusal): - analyze(tight, split, analysis="memory") - with pytest.raises(AnalysisError, match=refusal): - analyze(tight, split, analysis="performance") + tight_memory = analyze(tight, split, analysis="memory") + tight_record = get_metadata(tight_memory.function, MemoryMetadata) + assert tight_record.errors == ("smem placement peak 211200 B exceeds capacity 105600 B",) + assert '# error="smem placement peak 211200 B exceeds capacity 105600 B"' in render_text( + render_analysis(tight_memory) + ) + assert render_analysis(tight_memory).data["function_records"]["memory"]["errors"] == [ + "smem placement peak 211200 B exceeds capacity 105600 B" + ] + + tight_performance = analyze(tight, split, analysis="performance") + assert get_metadata(tight_performance.function, MemoryMetadata).errors == tight_record.errors unrestated = next(item for item in _SharedTile.functions if item.name == "split") held = get_metadata( @@ -375,7 +382,9 @@ def test_a_program_whose_buffers_have_nowhere_to_sit_is_refused() -> None: ) summary = get_metadata(fits.function, PerformanceSummaryMetadata) assert summary is not None - assert get_metadata(fits.function, MemoryMetadata).allocation.solver_status == "feasible" + fits_memory = get_metadata(fits.function, MemoryMetadata) + assert fits_memory.allocation.solver_status == "feasible" + assert fits_memory.errors == () assert summary.timeline.end_ns > 0 wider = replace(_SharedTile, target=_RoomierShared("nvidia.h200_sxm")) From fb677dbeda9079314db2dcf7365a2defcb445b01 Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Mon, 21 Sep 2026 10:19:09 +0800 Subject: [PATCH 7/9] fix(analysis): settle empty address placements [M2] --- src/tilefoundry/analysis/memory.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/src/tilefoundry/analysis/memory.py b/src/tilefoundry/analysis/memory.py index f525458b..0a58db26 100644 --- a/src/tilefoundry/analysis/memory.py +++ b/src/tilefoundry/analysis/memory.py @@ -490,7 +490,6 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: solver_options = ( context.options if isinstance(context.options, MemoryOptions) else MemoryOptions() ) - solver_statuses: list[str] = [] levels_list: list[MemoryLevelFootprint] = [] for name in sorted({item.memory_level for item in lifetimes} | set(memory_context.totals)): declared = facts.explicit(name) @@ -505,7 +504,6 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: options=solver_options, ) peak = solved.peak_bytes - solver_statuses.append(solved.solver_status) elif name == str(StorageKind.RMEM): peak = max((item.bytes for item in rows), default=0) else: @@ -533,9 +531,7 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: for item in levels if item.exceeds_capacity ) - allocation = None - if solver_statuses: - allocation = AllocationMetadata(solver_status="feasible") + allocation = AllocationMetadata(solver_status="feasible") attach( function, MemoryMetadata( From d4d72f73f696f70c08687170c56cc67e5ff3b71a Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Mon, 21 Sep 2026 10:21:21 +0800 Subject: [PATCH 8/9] test(analysis): lock solved memory peaks [M3] --- docs/tutorial/showcase.ipynb | 34 ++--- docs/tutorial/showcase.md | 63 ++++----- tests/analysis/test_analyze_at_a_size.py | 167 ++++++++++++++++++++--- 3 files changed, 194 insertions(+), 70 deletions(-) diff --git a/docs/tutorial/showcase.ipynb b/docs/tutorial/showcase.ipynb index 8fa29fa7..0647d479 100644 --- a/docs/tutorial/showcase.ipynb +++ b/docs/tutorial/showcase.ipynb @@ -90,7 +90,7 @@ { "name": "stdout", "output_type": "stream", - "text": "# analysis target=nvidia.h200_sxm module=Stage0_Naive function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,328896@total,328896@cta;f32:200448@logical,200448@total,200448@cta other-ops=special:1024@logical,1024@total,1024@cta\n# traffic traffic=gmem:r2225620/w806592@total,r2225620/w806592@cta\n# peak-footprint=gmem:1675788\n# roofline ideal-ns=632 bound-by=memory\n\n v0 = matmul(hidden, w_q, a_layout=\"MK\", b_layout=\"KN\") # Tensor[(1, 1, 256), \"bf16\"]; compute-cost flops=bf16:131072@logical,131072@total,131072@cta; traffic traffic=gmem:r131584/w512@total,r131584/w512@cta operands=0:r512/w0;1:r131072/w0;result:r0/w512; roofline ideal-ns=28 bound-by=memory\n v11 = cache_update(k_cache, cur_pos, write_len, v10) # Tensor[(1, 128, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n v34 = matmul(v33, w_o, a_layout=\"MK\", b_layout=\"KN\") # Tensor[(1, 1, 256), \"bf16\"]; compute-cost flops=bf16:131072@logical,131072@total,131072@cta; traffic traffic=gmem:r131584/w512@total,r131584/w512@cta operands=0:r512/w0;1:r131072/w0;result:r0/w512; roofline ideal-ns=28 bound-by=memory\n" + "text": "# analysis target=nvidia.h200_sxm module=Stage0_Naive function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,328896@total,328896@cta;f32:200448@logical,200448@total,200448@cta other-ops=special:1024@logical,1024@total,1024@cta\n# traffic traffic=gmem:r2225620/w806592@total,r2225620/w806592@cta\n# peak-footprint=gmem:1690380\n# roofline ideal-ns=632 bound-by=memory\n\n v0 = matmul(hidden, w_q, a_layout=\"MK\", b_layout=\"KN\") # Tensor[(1, 1, 256), \"bf16\"]; compute-cost flops=bf16:131072@logical,131072@total,131072@cta; traffic traffic=gmem:r131584/w512@total,r131584/w512@cta operands=0:r512/w0;1:r131072/w0;result:r0/w512; roofline ideal-ns=28 bound-by=memory\n v11 = cache_update(k_cache, cur_pos, write_len, v10) # Tensor[(1, 128, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n v34 = matmul(v33, w_o, a_layout=\"MK\", b_layout=\"KN\") # Tensor[(1, 1, 256), \"bf16\"]; compute-cost flops=bf16:131072@logical,131072@total,131072@cta; traffic traffic=gmem:r131584/w512@total,r131584/w512@cta operands=0:r512/w0;1:r131072/w0;result:r0/w512; roofline ideal-ns=28 bound-by=memory\n" } ], "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage0-128.txt\").read_text(encoding=\"utf-8\")\nheader, separator, annotated = report.partition(\"\\n\\n\")\nprint(header.rstrip())\nprint()\nfor needle in (\"matmul(hidden, w_q\", \"cache_update(k_cache\", \"matmul(v33, w_o\"):\n line = next(line for line in annotated.splitlines() if needle in line)\n print(line.rstrip())\n" @@ -137,7 +137,7 @@ { "name": "stdout", "output_type": "stream", - "text": "| `ctx_len` | f32 flops `global@CTA` | traffic `global@CTA` | peak gmem bytes | ideal ns | bound |\n|---:|---:|---|---:|---:|---|\n| 128 | `200448@logical,200448@total,200448@cta` | `gmem:r2225620/w806592@total,r2225620/w806592@cta` | 1675788 | 632 | memory |\n| 512 | `799488@logical,799488@total,799488@cta` | `gmem:r4744660/w3202752@total,r4744660/w3202752@cta` | 2572812 | 1656 | memory |\n| 1024 | `1598208@logical,1598208@total,1598208@cta` | `gmem:r8103380/w6397632@total,r8103380/w6397632@cta` | 3768844 | 3022 | memory |\n| 2048 | `3195648@logical,3195648@total,3195648@cta` | `gmem:r14820820/w12787392@total,r14820820/w12787392@cta` | 6160908 | 5752 | memory |\n| 4096 | `6390528@logical,6390528@total,6390528@cta` | `gmem:r28255700/w25566912@total,r28255700/w25566912@cta` | 10945036 | 11214 | memory |\n| 8192 | `12780288@logical,12780288@total,12780288@cta` | `gmem:r55125460/w51125952@total,r55125460/w51125952@cta` | 20513292 | 22136 | memory |\n" + "text": "| `ctx_len` | f32 flops `global@CTA` | traffic `global@CTA` | peak gmem bytes | ideal ns | bound |\n|---:|---:|---|---:|---:|---|\n| 128 | `200448@logical,200448@total,200448@cta` | `gmem:r2225620/w806592@total,r2225620/w806592@cta` | 1690380 | 632 | memory |\n| 512 | `799488@logical,799488@total,799488@cta` | `gmem:r4744660/w3202752@total,r4744660/w3202752@cta` | 2624268 | 1656 | memory |\n| 1024 | `1598208@logical,1598208@total,1598208@cta` | `gmem:r8103380/w6397632@total,r8103380/w6397632@cta` | 3869452 | 3022 | memory |\n| 2048 | `3195648@logical,3195648@total,3195648@cta` | `gmem:r14820820/w12787392@total,r14820820/w12787392@cta` | 6359820 | 5752 | memory |\n| 4096 | `6390528@logical,6390528@total,6390528@cta` | `gmem:r28255700/w25566912@total,r28255700/w25566912@cta` | 11340556 | 11214 | memory |\n| 8192 | `12780288@logical,12780288@total,12780288@cta` | `gmem:r55125460/w51125952@total,r55125460/w51125952@cta` | 21302028 | 22136 | memory |\n" } ], "source": "import re\nfrom pathlib import Path\n\n\ndef metrics(ctx_len):\n report = Path(f\"tutorial-reports/stage0-{ctx_len}.txt\").read_text(encoding=\"utf-8\")\n lines = report.splitlines()\n compute = next(line for line in lines if line.startswith(\"# compute-cost \"))\n traffic = next(line for line in lines if line.startswith(\"# traffic \"))\n peak = next(line for line in lines if line.startswith(\"# peak-footprint=\"))\n roofline = next(line for line in lines if line.startswith(\"# roofline \"))\n f32 = re.search(r\"f32:([^ ]+)\", compute).group(1)\n traffic_value = traffic.removeprefix(\"# traffic traffic=\")\n gmem_peak = re.search(r\"gmem:([^,]+)\", peak).group(1)\n ideal, bound = re.search(r\"ideal-ns=([^ ]+) bound-by=([^ ]+)\", roofline).groups()\n return f32, traffic_value, gmem_peak, ideal, bound\n\n\nprint(\"| `ctx_len` | f32 flops `global@CTA` | traffic `global@CTA` | peak gmem bytes | ideal ns | bound |\")\nprint(\"|---:|---:|---|---:|---:|---|\")\nfor ctx_len in (128, 512, 1024, 2048, 4096, 8192):\n f32, traffic, peak, ideal, bound = metrics(ctx_len)\n print(f\"| {ctx_len} | `{f32}` | `{traffic}` | {peak} | {ideal} | {bound} |\")\n" @@ -145,7 +145,7 @@ { "cell_type": "markdown", "metadata": {}, - "source": "The table says:\n\n```text\nper-CTA work = global work -> no authored split yet\nweight bytes = fixed -> projection weights are a staging target\ncache scan = grows with ctx_len -> a full-cache residency decision will fail first\n```\n\n## 2. Specialize at the capacity boundary\n\nThe full-cache sharded program in `Stage2_Sharded` places one local query head's K and V cache in smem. The boundary is derived from the target capacity:\n\n```text\nbytes per ctx per CTA = K + V\n = 2 * HEAD_DIM * sizeof(bf16)\n = 2 * 32 * 2\n = 128 B\n\nT = floor(232448 B / 128 B)\n = 1816\n```\n\n`Stage1_Specialized` expresses the dispatch as two closed `DimVarRangePat` variants: `[1, 1816]` and `[1817, 8192]`. The Stage1 body is deliberately the unsplit baseline, so the dispatch contract can be read independently from the later implementations.\n" + "source": "The table says:\n\n```text\nper-CTA work = global work -> no authored split yet\nweight bytes = fixed -> projection weights are a staging target\ncache scan = grows with ctx_len -> a full-cache residency decision will fail first\n```\n\n## 2. Specialize at the cache-only capacity estimate\n\nThe full-cache sharded program in `Stage2_Sharded` places one local query head's K and V cache in smem. A cache-only estimate follows from the target capacity:\n\n```text\nbytes per ctx per CTA = K + V\n = 2 * HEAD_DIM * sizeof(bf16)\n = 2 * 32 * 2\n = 128 B\n\nT = floor(232448 B / 128 B)\n = 1816\n```\n\nThis is not the complete placement peak: other simultaneously resident values also occupy smem. `Stage1_Specialized` still uses the estimate to express two closed `DimVarRangePat` variants, `[1, 1816]` and `[1817, 8192]`; the CP placement report below decides whether either variant actually fits. The Stage1 body is deliberately the unsplit baseline, so the dispatch contract can be read independently from the later implementations.\n" }, { "cell_type": "code", @@ -161,7 +161,7 @@ { "cell_type": "markdown", "metadata": {}, - "source": "The Bash cell writes the valid boundary report. The following Python cell loads the file\nand prints its report header:\n" + "source": "The Bash cell writes the complete report at the cache-only estimate. Analysis still returns when the solved placement exceeds capacity; the report header carries a non-fatal `# error=` line:\n" }, { "cell_type": "code", @@ -188,7 +188,7 @@ { "name": "stdout", "output_type": "stream", - "text": "# analysis target=nvidia.h200_sxm module=Stage2_Sharded function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,2629376@total,328672@cta;f32:2833728@logical,2833728@total,354216@cta other-ops=special:14528@logical,14528@total,1816@cta\n# traffic traffic=gmem:r5563796/w3721856@total,r3936212/w3721408@cta;smem:r9597248/w9480000@total,r1199656/w1185000@cta\n# peak-footprint=gmem:3701260;smem:472160\n# roofline ideal-ns=1935 bound-by=memory\n" + "text": "# analysis target=nvidia.h200_sxm module=Stage2_Sharded function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,2629376@total,328672@cta;f32:2833728@logical,2833728@total,354216@cta other-ops=special:14528@logical,14528@total,1816@cta\n# traffic traffic=gmem:r5563796/w3721856@total,r3936212/w3721408@cta;smem:r9597248/w9480000@total,r1199656/w1185000@cta\n# peak-footprint=gmem:3933836;smem:581312\n# error=\"smem placement peak 581312 B exceeds capacity 232448 B\"\n# roofline ideal-ns=1935 bound-by=memory\n" } ], "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage2-1816.txt\").read_text(encoding=\"utf-8\")\nprint(report.partition(\"\\n\\n\")[0].rstrip())\n" @@ -196,7 +196,7 @@ { "cell_type": "markdown", "metadata": {}, - "source": "A larger context crosses the stated capacity. The Bash cell preserves the non-zero\nCLI refusal in a file; the following Python cell loads and prints the actual error output:\n" + "source": "A larger context increases the solved peak. The CLI still succeeds and writes the complete report, including its capacity error:\n" }, { "cell_type": "code", @@ -204,11 +204,11 @@ "metadata": { "tilefoundry": { "cell_type": "bash", - "command": "stage2-1820-refusal" + "command": "stage2-1820-error" } }, "outputs": [], - "source": "%%bash\nset -euo pipefail\nmkdir -p tutorial-reports\nset +e\ntilefoundry analyze attn_layer.py:Stage2_Sharded \\\n tutorial-reports/stage2-1820.txt \\\n --compute-cost --memory --roofline --dim ctx_len=1820 \\\n 2> tutorial-reports/stage2-1820.err\nstatus=$?\nset -e\nif [ \"$status\" -eq 0 ]; then\n echo \"expected Stage2_Sharded to refuse ctx_len=1820\" >&2\n exit 1\nfi\n" + "source": "%%bash\nset -euo pipefail\nmkdir -p tutorial-reports\ntilefoundry analyze attn_layer.py:Stage2_Sharded \\\n tutorial-reports/stage2-1820.txt \\\n --compute-cost --memory --roofline --dim ctx_len=1820\n" }, { "cell_type": "code", @@ -216,22 +216,22 @@ "metadata": { "tilefoundry": { "output_format": "text", - "analysis": "REFUSAL_STAGE2_1820" + "analysis": "STAGE2_1820_ERROR" } }, "outputs": [ { "name": "stdout", "output_type": "stream", - "text": "tilefoundry: error: function 'gqa_decode': value 'v16:240' needs 232960 B in smem, which exceeds the 232448 B the target states for that level\n" + "text": "# analysis target=nvidia.h200_sxm module=Stage2_Sharded function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,2629376@total,328672@cta;f32:2839968@logical,2839968@total,354996@cta other-ops=special:14560@logical,14560@total,1820@cta\n# traffic traffic=gmem:r5573012/w3730048@total,r3941844/w3729600@cta;smem:r9618368/w9500864@total,r1202296/w1187608@cta\n# peak-footprint=gmem:3939468;smem:582592\n# error=\"smem placement peak 582592 B exceeds capacity 232448 B\"\n# roofline ideal-ns=1939 bound-by=memory\n" } ], - "source": "from pathlib import Path\n\nerror = Path(\"tutorial-reports/stage2-1820.err\").read_text(encoding=\"utf-8\")\nprint(error.rstrip())\n" + "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage2-1820.txt\").read_text(encoding=\"utf-8\")\nprint(report.partition(\"\\n\\n\")[0].rstrip())\n" }, { "cell_type": "markdown", "metadata": {}, - "source": "The formula chooses the dispatch boundary. It is not a benchmark-tuned magic number.\n\n## 3. Split the query heads\n\nThe next change is `Stage2_Sharded`: one `cta.head` owns one query head. The same-size comparison isolates placement from context growth.\n" + "source": "The formula supplies a cache-only dispatch estimate, while the solved placement and its report error expose the complete capacity requirement.\n\n## 3. Split the query heads\n\nThe next change is `Stage2_Sharded`: one `cta.head` owns one query head. The same-size comparison isolates placement from context growth.\n" }, { "cell_type": "code", @@ -274,7 +274,7 @@ { "name": "stdout", "output_type": "stream", - "text": "# compute-cost flops=bf16:328896@logical,328896@total,328896@cta;f32:200448@logical,200448@total,200448@cta other-ops=special:1024@logical,1024@total,1024@cta\n# traffic traffic=gmem:r2225620/w806592@total,r2225620/w806592@cta\n# peak-footprint=gmem:1675788\n# roofline ideal-ns=632 bound-by=memory\n" + "text": "# compute-cost flops=bf16:328896@logical,328896@total,328896@cta;f32:200448@logical,200448@total,200448@cta other-ops=special:1024@logical,1024@total,1024@cta\n# traffic traffic=gmem:r2225620/w806592@total,r2225620/w806592@cta\n# peak-footprint=gmem:1690380\n# roofline ideal-ns=632 bound-by=memory\n" } ], "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage0-128-summary.txt\").read_text(encoding=\"utf-8\")\nfor line in report.splitlines():\n if line.startswith((\"# compute-cost \", \"# traffic \", \"# peak-footprint=\", \"# roofline \")):\n print(line)\n" @@ -309,7 +309,7 @@ { "name": "stdout", "output_type": "stream", - "text": "# compute-cost flops=bf16:328896@logical,2629376@total,328672@cta;f32:200448@logical,200448@total,25056@cta other-ops=special:1024@logical,1024@total,128@cta\n# traffic traffic=gmem:r1674644/w264832@total,r1559508/w264384@cta;smem:r684608/w675392@total,r85576/w84424@cta\n# peak-footprint=gmem:1540620;smem:33280\n# roofline ideal-ns=405 bound-by=memory\n" + "text": "# compute-cost flops=bf16:328896@logical,2629376@total,328672@cta;f32:200448@logical,200448@total,25056@cta other-ops=special:1024@logical,1024@total,128@cta\n# traffic traffic=gmem:r1674644/w264832@total,r1559508/w264384@cta;smem:r684608/w675392@total,r85576/w84424@cta\n# peak-footprint=gmem:1557132;smem:41152\n# roofline ideal-ns=405 bound-by=memory\n" } ], "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage2-128-summary.txt\").read_text(encoding=\"utf-8\")\nfor line in report.splitlines():\n if line.startswith((\"# compute-cost \", \"# traffic \", \"# peak-footprint=\", \"# roofline \")):\n print(line)\n" @@ -360,7 +360,7 @@ { "name": "stdout", "output_type": "stream", - "text": "# analysis target=nvidia.h200_sxm module=Stage3_Fused function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,4392896@total,328672@cta;f32:3239808@logical,6418944@total,200592@cta other-ops=integer:9@logical,288@total,9@cta;special:33056@logical,33152@total,1036@cta\n# traffic traffic=gmem:r3476884/w4196992@total,r2558932/w4196544@cta;rmem:r656/w72@total,r656/w72@cta;smem:r5839296/w5662784@total,r682464/w672676@cta\n# peak-footprint=gmem:7144460;rmem:16;smem:33552\n# roofline ideal-ns=1599 bound-by=memory\n\n v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n" + "text": "# analysis target=nvidia.h200_sxm module=Stage3_Fused function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,4392896@total,328672@cta;f32:3239808@logical,6418944@total,200592@cta other-ops=integer:9@logical,288@total,9@cta;special:33056@logical,33152@total,1036@cta\n# traffic traffic=gmem:r3476884/w4196992@total,r2558932/w4196544@cta;rmem:r656/w72@total,r656/w72@cta;smem:r5839296/w5662784@total,r682464/w672676@cta\n# peak-footprint=gmem:7145228;rmem:8;smem:42312\n# roofline ideal-ns=1599 bound-by=memory\n\n v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n" } ], "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage3-4096.txt\").read_text(encoding=\"utf-8\")\nheader, separator, annotated = report.partition(\"\\n\\n\")\nprint(header.rstrip())\nprint()\nprint(next(line.rstrip() for line in annotated.splitlines() if \"cache_update(k_cache\" in line))\n" @@ -411,7 +411,7 @@ { "name": "stdout", "output_type": "stream", - "text": "# analysis target=nvidia.h200_sxm module=Stage4_WeightPrepared function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,337408@total,42176@cta;f32:6390528@logical,51124224@total,6390528@cta other-ops=special:32768@logical,262144@total,32768@cta\n# traffic traffic=gmem:r28254676/w25566912@total,r27967956/w25565792@cta;smem:r331008/w329984@total,r43168/w42144@cta\n# peak-footprint=gmem:10945036;smem:16960\n# roofline ideal-ns=11213 bound-by=memory\n\n v1 = reshard(w_q, layout=(1, 256, 8 @ mesh.head, 32), storage=smem) # Tensor[(1, 256, 256), \"bf16\", ((1, 256, 8 @ mesh.head, 32), (0, 32, 0, 1)), \"smem\"]; compute-cost; traffic traffic=gmem:r131072/w0@total,r16384/w0@cta;smem:r0/w131072@total,r0/w16384@cta operands=0:r131072/w0;result:r0/w131072; roofline ideal-ns=28 bound-by=memory\n v2 = matmul(v0, v1, a_layout=\"MK\", b_layout=\"KN\") # Tensor[(1, 1, 256), \"bf16\", ((1, 1, 8 @ mesh.head, 32), (256, 256, 32, 1)), \"smem\"]; compute-cost flops=bf16:131072@logical,131072@total,16384@cta; traffic traffic=smem:r131584/w512@total,r16896/w64@cta operands=0:r512/w0;1:r131072/w0;result:r0/w512; roofline ideal-ns=1 bound-by=compute\n\n v42 = reshard(w_o, layout=(1, 256, 8 @ mesh.head, 32), storage=smem) # Tensor[(1, 256, 256), \"bf16\", ((1, 256, 8 @ mesh.head, 32), (0, 32, 0, 1)), \"smem\"]; compute-cost; traffic traffic=gmem:r131072/w0@total,r16384/w0@cta;smem:r0/w131072@total,r0/w16384@cta operands=0:r131072/w0;result:r0/w131072; roofline ideal-ns=28 bound-by=memory\n v43 = matmul(v41, v42, a_layout=\"MK\", b_layout=\"KN\") # Tensor[(1, 1, 256), \"bf16\", ((1, 1, 8 @ mesh.head, 32), (256, 256, 32, 1)), \"smem\"]; compute-cost flops=bf16:131072@logical,131072@total,16384@cta; traffic traffic=smem:r131584/w512@total,r16896/w64@cta operands=0:r512/w0;1:r131072/w0;result:r0/w512; roofline ideal-ns=1 bound-by=compute\n" + "text": "# analysis target=nvidia.h200_sxm module=Stage4_WeightPrepared function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,337408@total,42176@cta;f32:6390528@logical,51124224@total,6390528@cta other-ops=special:32768@logical,262144@total,32768@cta\n# traffic traffic=gmem:r28254676/w25566912@total,r27967956/w25565792@cta;smem:r331008/w329984@total,r43168/w42144@cta\n# peak-footprint=gmem:11340556;smem:16960\n# roofline ideal-ns=11213 bound-by=memory\n\n v1 = reshard(w_q, layout=(1, 256, 8 @ mesh.head, 32), storage=smem) # Tensor[(1, 256, 256), \"bf16\", ((1, 256, 8 @ mesh.head, 32), (0, 32, 0, 1)), \"smem\"]; compute-cost; traffic traffic=gmem:r131072/w0@total,r16384/w0@cta;smem:r0/w131072@total,r0/w16384@cta operands=0:r131072/w0;result:r0/w131072; roofline ideal-ns=28 bound-by=memory\n v2 = matmul(v0, v1, a_layout=\"MK\", b_layout=\"KN\") # Tensor[(1, 1, 256), \"bf16\", ((1, 1, 8 @ mesh.head, 32), (256, 256, 32, 1)), \"smem\"]; compute-cost flops=bf16:131072@logical,131072@total,16384@cta; traffic traffic=smem:r131584/w512@total,r16896/w64@cta operands=0:r512/w0;1:r131072/w0;result:r0/w512; roofline ideal-ns=1 bound-by=compute\n\n v42 = reshard(w_o, layout=(1, 256, 8 @ mesh.head, 32), storage=smem) # Tensor[(1, 256, 256), \"bf16\", ((1, 256, 8 @ mesh.head, 32), (0, 32, 0, 1)), \"smem\"]; compute-cost; traffic traffic=gmem:r131072/w0@total,r16384/w0@cta;smem:r0/w131072@total,r0/w16384@cta operands=0:r131072/w0;result:r0/w131072; roofline ideal-ns=28 bound-by=memory\n v43 = matmul(v41, v42, a_layout=\"MK\", b_layout=\"KN\") # Tensor[(1, 1, 256), \"bf16\", ((1, 1, 8 @ mesh.head, 32), (256, 256, 32, 1)), \"smem\"]; compute-cost flops=bf16:131072@logical,131072@total,16384@cta; traffic traffic=smem:r131584/w512@total,r16896/w64@cta operands=0:r512/w0;1:r131072/w0;result:r0/w512; roofline ideal-ns=1 bound-by=compute\n" } ], "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage4-4096.txt\").read_text(encoding=\"utf-8\")\nheader, separator, annotated = report.partition(\"\\n\\n\")\nprint(header.rstrip())\nlines = annotated.splitlines()\nfor needle in (\"reshard(w_q\", \"reshard(w_o\"):\n start = next(index for index, line in enumerate(lines) if needle in line)\n end = start\n while end + 1 < len(lines):\n end += 1\n if end > start and \" # \" in lines[end]:\n break\n print()\n print(\"\\n\".join(line.rstrip() for line in lines[start : end + 1]))\n" @@ -462,7 +462,7 @@ { "name": "stdout", "output_type": "stream", - "text": "# analysis target=nvidia.h200_sxm module=Stage5_CachePrepared function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,1246400@total,328672@cta;f32:6410280@logical,6410280@total,801285@cta other-ops=integer:32@logical,256@total,32@cta;special:33040@logical,33040@total,4130@cta\n# traffic traffic=gmem:r7672476/w4198272@total,r4001116/w4197824@cta;rmem:r2560/w0@total,r2560/w0@cta;smem:r21788704/w21486432@total,r2723588/w2685804@cta\n# peak-footprint=gmem:3474572;rmem:0;smem:33680\n# roofline ideal-ns=2474 bound-by=memory\n\n v21 = slice(k_cache, (0, v20, 0, 0), sizes=(1, 128, 2, 32), strides=(1, 1, 1, 1)) # Tensor[(1, 128, 2, 32), \"bf16\"]; compute-cost; traffic traffic=rmem:r32/w0@total,r32/w0@cta operands=0:r0/w0;1:r32/w0;result:r0/w0; roofline\n v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n" + "text": "# analysis target=nvidia.h200_sxm module=Stage5_CachePrepared function=gqa_decode topology=cta\n# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline\n# compute-cost flops=bf16:328896@logical,1246400@total,328672@cta;f32:6410280@logical,6410280@total,801285@cta other-ops=integer:32@logical,256@total,32@cta;special:33040@logical,33040@total,4130@cta\n# traffic traffic=gmem:r7672476/w4198272@total,r4001116/w4197824@cta;rmem:r2560/w0@total,r2560/w0@cta;smem:r21788704/w21486432@total,r2723588/w2685804@cta\n# peak-footprint=gmem:3475212;rmem:0;smem:41920\n# roofline ideal-ns=2474 bound-by=memory\n\n v21 = slice(k_cache, (0, v20, 0, 0), sizes=(1, 128, 2, 32), strides=(1, 1, 1, 1)) # Tensor[(1, 128, 2, 32), \"bf16\"]; compute-cost; traffic traffic=rmem:r32/w0@total,r32/w0@cta operands=0:r0/w0;1:r32/w0;result:r0/w0; roofline\n v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), \"bf16\"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory\n" } ], "source": "from pathlib import Path\n\nreport = Path(\"tutorial-reports/stage5-4096.txt\").read_text(encoding=\"utf-8\")\nheader, separator, annotated = report.partition(\"\\n\\n\")\nprint(header.rstrip())\nprint()\nfor needle in (\"slice(k_cache\", \"cache_update(k_cache\"):\n print(next(line.rstrip() for line in annotated.splitlines() if needle in line))\n" diff --git a/docs/tutorial/showcase.md b/docs/tutorial/showcase.md index 2bc44914..ab318d21 100644 --- a/docs/tutorial/showcase.md +++ b/docs/tutorial/showcase.md @@ -199,7 +199,7 @@ for needle in ("matmul(hidden, w_q", "cache_update(k_cache", "matmul(v33, w_o"): # selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline # compute-cost flops=bf16:328896@logical,328896@total,328896@cta;f32:200448@logical,200448@total,200448@cta other-ops=special:1024@logical,1024@total,1024@cta # traffic traffic=gmem:r2225620/w806592@total,r2225620/w806592@cta -# peak-footprint=gmem:1675788 +# peak-footprint=gmem:1690380 # roofline ideal-ns=632 bound-by=memory v0 = matmul(hidden, w_q, a_layout="MK", b_layout="KN") # Tensor[(1, 1, 256), "bf16"]; compute-cost flops=bf16:131072@logical,131072@total,131072@cta; traffic traffic=gmem:r131584/w512@total,r131584/w512@cta operands=0:r512/w0;1:r131072/w0;result:r0/w512; roofline ideal-ns=28 bound-by=memory @@ -259,12 +259,12 @@ for ctx_len in (128, 512, 1024, 2048, 4096, 8192): | `ctx_len` | f32 flops `global@CTA` | traffic `global@CTA` | peak gmem bytes | ideal ns | bound | |---:|---:|---|---:|---:|---| -| 128 | `200448@logical,200448@total,200448@cta` | `gmem:r2225620/w806592@total,r2225620/w806592@cta` | 1675788 | 632 | memory | -| 512 | `799488@logical,799488@total,799488@cta` | `gmem:r4744660/w3202752@total,r4744660/w3202752@cta` | 2572812 | 1656 | memory | -| 1024 | `1598208@logical,1598208@total,1598208@cta` | `gmem:r8103380/w6397632@total,r8103380/w6397632@cta` | 3768844 | 3022 | memory | -| 2048 | `3195648@logical,3195648@total,3195648@cta` | `gmem:r14820820/w12787392@total,r14820820/w12787392@cta` | 6160908 | 5752 | memory | -| 4096 | `6390528@logical,6390528@total,6390528@cta` | `gmem:r28255700/w25566912@total,r28255700/w25566912@cta` | 10945036 | 11214 | memory | -| 8192 | `12780288@logical,12780288@total,12780288@cta` | `gmem:r55125460/w51125952@total,r55125460/w51125952@cta` | 20513292 | 22136 | memory | +| 128 | `200448@logical,200448@total,200448@cta` | `gmem:r2225620/w806592@total,r2225620/w806592@cta` | 1690380 | 632 | memory | +| 512 | `799488@logical,799488@total,799488@cta` | `gmem:r4744660/w3202752@total,r4744660/w3202752@cta` | 2624268 | 1656 | memory | +| 1024 | `1598208@logical,1598208@total,1598208@cta` | `gmem:r8103380/w6397632@total,r8103380/w6397632@cta` | 3869452 | 3022 | memory | +| 2048 | `3195648@logical,3195648@total,3195648@cta` | `gmem:r14820820/w12787392@total,r14820820/w12787392@cta` | 6359820 | 5752 | memory | +| 4096 | `6390528@logical,6390528@total,6390528@cta` | `gmem:r28255700/w25566912@total,r28255700/w25566912@cta` | 11340556 | 11214 | memory | +| 8192 | `12780288@logical,12780288@total,12780288@cta` | `gmem:r55125460/w51125952@total,r55125460/w51125952@cta` | 21302028 | 22136 | memory | The table says: @@ -274,9 +274,9 @@ weight bytes = fixed -> projection weights are a staging target cache scan = grows with ctx_len -> a full-cache residency decision will fail first ``` -## 2. Specialize at the capacity boundary +## 2. Specialize at the cache-only capacity estimate -The full-cache sharded program in `Stage2_Sharded` places one local query head's K and V cache in smem. The boundary is derived from the target capacity: +The full-cache sharded program in `Stage2_Sharded` places one local query head's K and V cache in smem. A cache-only estimate follows from the target capacity: ```text bytes per ctx per CTA = K + V @@ -288,7 +288,7 @@ T = floor(232448 B / 128 B) = 1816 ``` -`Stage1_Specialized` expresses the dispatch as two closed `DimVarRangePat` variants: `[1, 1816]` and `[1817, 8192]`. The Stage1 body is deliberately the unsplit baseline, so the dispatch contract can be read independently from the later implementations. +This is not the complete placement peak: other simultaneously resident values also occupy smem. `Stage1_Specialized` still uses the estimate to express two closed `DimVarRangePat` variants, `[1, 1816]` and `[1817, 8192]`; the CP placement report below decides whether either variant actually fits. The Stage1 body is deliberately the unsplit baseline, so the dispatch contract can be read independently from the later implementations. @@ -405,8 +405,7 @@ class Stage1_Specialized: gqa_decode_specialized = Stage1_Specialized.entry_function() ``` -The Bash cell writes the valid boundary report. The following Python cell loads the file -and prints its report header: +The Bash cell writes the complete report at the cache-only estimate. Analysis still returns when the solved placement exceeds capacity; the report header carries a non-fatal `# error=` line: ```bash set -euo pipefail @@ -428,41 +427,39 @@ print(report.partition("\n\n")[0].rstrip()) # selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline # compute-cost flops=bf16:328896@logical,2629376@total,328672@cta;f32:2833728@logical,2833728@total,354216@cta other-ops=special:14528@logical,14528@total,1816@cta # traffic traffic=gmem:r5563796/w3721856@total,r3936212/w3721408@cta;smem:r9597248/w9480000@total,r1199656/w1185000@cta -# peak-footprint=gmem:3701260;smem:472160 +# peak-footprint=gmem:3933836;smem:581312 +# error="smem placement peak 581312 B exceeds capacity 232448 B" # roofline ideal-ns=1935 bound-by=memory ``` -A larger context crosses the stated capacity. The Bash cell preserves the non-zero -CLI refusal in a file; the following Python cell loads and prints the actual error output: +A larger context increases the solved peak. The CLI still succeeds and writes the complete report, including its capacity error: ```bash set -euo pipefail mkdir -p tutorial-reports -set +e tilefoundry analyze attn_layer.py:Stage2_Sharded \ tutorial-reports/stage2-1820.txt \ - --compute-cost --memory --roofline --dim ctx_len=1820 \ - 2> tutorial-reports/stage2-1820.err -status=$? -set -e -if [ "$status" -eq 0 ]; then - echo "expected Stage2_Sharded to refuse ctx_len=1820" >&2 - exit 1 -fi + --compute-cost --memory --roofline --dim ctx_len=1820 ``` ```python from pathlib import Path -error = Path("tutorial-reports/stage2-1820.err").read_text(encoding="utf-8") -print(error.rstrip()) +report = Path("tutorial-reports/stage2-1820.txt").read_text(encoding="utf-8") +print(report.partition("\n\n")[0].rstrip()) ``` ```text -tilefoundry: error: function 'gqa_decode': value 'v16:240' needs 232960 B in smem, which exceeds the 232448 B the target states for that level +# analysis target=nvidia.h200_sxm module=Stage2_Sharded function=gqa_decode topology=cta +# selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline +# compute-cost flops=bf16:328896@logical,2629376@total,328672@cta;f32:2839968@logical,2839968@total,354996@cta other-ops=special:14560@logical,14560@total,1820@cta +# traffic traffic=gmem:r5573012/w3730048@total,r3941844/w3729600@cta;smem:r9618368/w9500864@total,r1202296/w1187608@cta +# peak-footprint=gmem:3939468;smem:582592 +# error="smem placement peak 582592 B exceeds capacity 232448 B" +# roofline ideal-ns=1939 bound-by=memory ``` -The formula chooses the dispatch boundary. It is not a benchmark-tuned magic number. +The formula supplies a cache-only dispatch estimate, while the solved placement and its report error expose the complete capacity requirement. ## 3. Split the query heads @@ -556,7 +553,7 @@ for line in report.splitlines(): ```text # compute-cost flops=bf16:328896@logical,328896@total,328896@cta;f32:200448@logical,200448@total,200448@cta other-ops=special:1024@logical,1024@total,1024@cta # traffic traffic=gmem:r2225620/w806592@total,r2225620/w806592@cta -# peak-footprint=gmem:1675788 +# peak-footprint=gmem:1690380 # roofline ideal-ns=632 bound-by=memory ``` @@ -583,7 +580,7 @@ for line in report.splitlines(): ```text # compute-cost flops=bf16:328896@logical,2629376@total,328672@cta;f32:200448@logical,200448@total,25056@cta other-ops=special:1024@logical,1024@total,128@cta # traffic traffic=gmem:r1674644/w264832@total,r1559508/w264384@cta;smem:r684608/w675392@total,r85576/w84424@cta -# peak-footprint=gmem:1540620;smem:33280 +# peak-footprint=gmem:1557132;smem:41152 # roofline ideal-ns=405 bound-by=memory ``` @@ -730,7 +727,7 @@ print(next(line.rstrip() for line in annotated.splitlines() if "cache_update(k_c # selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline # compute-cost flops=bf16:328896@logical,4392896@total,328672@cta;f32:3239808@logical,6418944@total,200592@cta other-ops=integer:9@logical,288@total,9@cta;special:33056@logical,33152@total,1036@cta # traffic traffic=gmem:r3476884/w4196992@total,r2558932/w4196544@cta;rmem:r656/w72@total,r656/w72@cta;smem:r5839296/w5662784@total,r682464/w672676@cta -# peak-footprint=gmem:7144460;rmem:16;smem:33552 +# peak-footprint=gmem:7145228;rmem:8;smem:42312 # roofline ideal-ns=1599 bound-by=memory v6 = cache_update(k_cache, cur_pos, write_len, v5) # Tensor[(1, 4096, 2, 32), "bf16"]; compute-cost; traffic traffic=gmem:r136/w128@total,r136/w128@cta operands=0:r0/w0;1:r4/w0;2:r4/w0;3:r128/w0;result:r0/w128; roofline ideal-ns=1 bound-by=memory @@ -867,7 +864,7 @@ for needle in ("reshard(w_q", "reshard(w_o"): # selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline # compute-cost flops=bf16:328896@logical,337408@total,42176@cta;f32:6390528@logical,51124224@total,6390528@cta other-ops=special:32768@logical,262144@total,32768@cta # traffic traffic=gmem:r28254676/w25566912@total,r27967956/w25565792@cta;smem:r331008/w329984@total,r43168/w42144@cta -# peak-footprint=gmem:10945036;smem:16960 +# peak-footprint=gmem:11340556;smem:16960 # roofline ideal-ns=11213 bound-by=memory v1 = reshard(w_q, layout=(1, 256, 8 @ mesh.head, 32), storage=smem) # Tensor[(1, 256, 256), "bf16", ((1, 256, 8 @ mesh.head, 32), (0, 32, 0, 1)), "smem"]; compute-cost; traffic traffic=gmem:r131072/w0@total,r16384/w0@cta;smem:r0/w131072@total,r0/w16384@cta operands=0:r131072/w0;result:r0/w131072; roofline ideal-ns=28 bound-by=memory @@ -1019,7 +1016,7 @@ for needle in ("slice(k_cache", "cache_update(k_cache"): # selection requested=compute-cost,memory,roofline executed=compute-cost,memory,roofline # compute-cost flops=bf16:328896@logical,1246400@total,328672@cta;f32:6410280@logical,6410280@total,801285@cta other-ops=integer:32@logical,256@total,32@cta;special:33040@logical,33040@total,4130@cta # traffic traffic=gmem:r7672476/w4198272@total,r4001116/w4197824@cta;rmem:r2560/w0@total,r2560/w0@cta;smem:r21788704/w21486432@total,r2723588/w2685804@cta -# peak-footprint=gmem:3474572;rmem:0;smem:33680 +# peak-footprint=gmem:3475212;rmem:0;smem:41920 # roofline ideal-ns=2474 bound-by=memory v21 = slice(k_cache, (0, v20, 0, 0), sizes=(1, 128, 2, 32), strides=(1, 1, 1, 1)) # Tensor[(1, 128, 2, 32), "bf16"]; compute-cost; traffic traffic=rmem:r32/w0@total,r32/w0@cta operands=0:r0/w0;1:r32/w0;result:r0/w0; roofline diff --git a/tests/analysis/test_analyze_at_a_size.py b/tests/analysis/test_analyze_at_a_size.py index 90e7831f..bd607c96 100644 --- a/tests/analysis/test_analyze_at_a_size.py +++ b/tests/analysis/test_analyze_at_a_size.py @@ -51,15 +51,157 @@ CONTEXT = 32 DIMS = {"ctx_len": CONTEXT} FAMILIES = ("compute-cost", "memory", "roofline", "performance") -INVENTORY = [pytest.param(case, id=case.id) for case in placed_cases()] +CASES = placed_cases() +INVENTORY = [pytest.param(case, id=case.id) for case in CASES] EXPECTED_MEMORY_PEAKS = { + "derived_prefill.DerivedPrefill.prefill[prefill_n=64,topology_only=128]": { + "gmem": 288, + }, + "flash_split_k_decode.FlashSplitKDecode.flash_split_k_decode[ctx=128]": { + "gmem": 788_480, + "rmem": 8, + "smem": 83_592, + }, + "fused_boundary.FusedBoundary.inner.run[static]": {"rmem": 128}, + "fused_boundary.FusedBoundary.inner.scale[static]": {"rmem": 128}, + "fused_boundary.FusedBoundary.root[static]": { + "gmem": 512, + "rmem": 128, + "smem": 32, + }, + "fused_boundary.FusedBoundary.stage[static]": {"smem": 64}, + "gqa_decode.GqaOnline._ctx_combine[static]": {"gmem": 291_968}, + "gqa_decode.GqaOnline._ctx_partials[ctx_len=128]": {"gmem": 5_662_720}, + "gqa_decode.GqaOnline.gqa_online_attend[ctx_len=128]": { + "gmem": 283_752, + "rmem": 0, + }, + "leaf_weights.Mod.entry[static]": { + "gmem": 51_539_608_064, + "rmem": 0, + "smem": 160, + }, + "leaf_weights.Mod.leaf[static]": {"gmem": 512, "smem": 64}, + "leaf_weights.Mod.other[static]": { + "gmem": 51_539_608_064, + "rmem": 0, + "smem": 160, + }, + "mesh_slice_start.Fixed.scan[static]": { + "gmem": 5_120, + "rmem": 0, + "smem": 1_408, + }, + "mesh_slice_start.OutOfWindow.oob[static]": {"gmem": 6_144, "rmem": 0}, + "mesh_slice_start.Strided.scan[static]": { + "gmem": 5_120, + "rmem": 8, + "smem": 1_408, + }, + "mha_decode_paged.Batch2Page256.mha_decode_paged[static]": { + "gmem": 5_245_000, + "rmem": 32_768, + "smem": 16_384, + }, + "mha_decode_paged.LongerCache.mha_decode_paged[static]": { + "gmem": 4_195_364, + "rmem": 16_384, + "smem": 8_192, + }, + "mha_decode_paged.ShorterCache.mha_decode_paged[static]": { + "gmem": 2_098_196, + "rmem": 8_192, + "smem": 4_096, + }, + "mha_decode_paged.SingleTokenPage128.mha_decode_paged[static]": { + "gmem": 8_392_740, + "rmem": 32_768, + "smem": 16_384, + }, + "moe_mega_kernel.MoEMegaKernel.experts[static]": {"gmem": 61_440}, + "moe_mega_kernel.MoEMegaKernel.routed_expert[static]": {"gmem": 61_440}, + "moe_mega_kernel.MoEMegaKernel.shared_expert[static]": {"gmem": 61_440}, + "nested_twin.Weighted.scaled[static]": {"gmem": 1_348, "rmem": 4}, + "performance_findings.Compare.kernel[static]": {"gmem": 136_208}, + "performance_findings.GmemSquare.kernel[static]": {"gmem": 68_096}, + "performance_findings.Levels.kernel[static]": { + "gmem": 2_113_536, + "rmem": 16_384, + }, + "performance_findings.LevelsNested.kernel[static]": { + "gmem": 2_113_536, + "rmem": 16_384, + }, + "performance_findings.LevelsOnOneMesh.kernel[static]": { + "gmem": 2_113_536, + "rmem": 16_384, + }, + "performance_findings.LocalTier.kernel[static]": {"gmem": 68_096, "rmem": 512}, + "prefill_decode_attention.PrefillDecodeAttention.attend[ctx=128,seq=128]": { + "gmem": 1_572_864, + "rmem": 0, + "smem": 245_760, + }, + "qwen3_1_7b_pd.PrefillLayer.layer_decode[ctx_len=128,seq=128]": { + "gmem": 148_521_988, + "rmem": 520, + "smem": 65_792, + }, "qwen3_1_7b_pd.PrefillLayer.layer_prefill[ctx_len=128,seq=128]": { - "gmem": 175_514_632, - "rmem": 395_264, - "smem": 98_304, + "gmem": 178_922_500, + "rmem": 66_560, + "smem": 131_072, + }, + "qwen3_1_7b_pd.PrefillLayer.model[ctx_len=0,seq=512]": { + "gmem": 5_750_002_180, + "rmem": 66_560, + "smem": 131_072, + }, + "qwen3_1_7b_pd.PrefillLayer.model[ctx_len=4608,seq=1]": { + "gmem": 4_763_301_384, + "rmem": 520, + "smem": 65_792, + }, + "qwen3_1_7b_pd.PrefillLayer.model[ctx_len=512,seq=1]": { + "gmem": 4_763_301_384, + "rmem": 520, + "smem": 65_792, }, + "qwen3_1_7b_pd.PrefillLayer.model[ctx_len=512,seq=512]": { + "gmem": 5_750_002_180, + "rmem": 66_560, + "smem": 131_072, + }, + "region_boundaries.RegionBoundaries.helper[static]": {"gmem": 64, "rmem": 32}, + "region_boundaries.RegionBoundaries.run[static]": { + "gmem": 64, + "rmem": 32, + "smem": 32, + }, + "rmsnorm.RmsnormModule.rmsnorm[static]": {"gmem": 6_144, "rmem": 6_144}, + "rmsnorm_quant_seq2.RmsnormQuantSeq2Module.rmsnorm_quant_seq_2[static]": { + "gmem": 9_312, + "rmem": 12_288, + }, + "rmsnorm_seq2.RmsnormSeq2Module.rmsnorm_seq_2[static]": { + "gmem": 12_288, + "rmem": 12_288, + }, + "specialize_through_call.Direct.pick[n=128]": {"gmem": 1_024, "smem": 128}, + "specialize_through_call.Direct.run[n=128]": {"gmem": 1_024, "smem": 128}, + "specialize_through_call.ToCallee.pick[n=128]": {"gmem": 1_024, "smem": 128}, + "specialize_through_call.ToCallee.run[n=128]": {"gmem": 1_024, "smem": 128}, + "square_cuda.Model.main[static]": {"gmem": 676, "rmem": 4}, + "tiny_tp_decoder.DecoderLayer.decode[static]": {"gmem": 48, "rmem": 16}, + "tiny_tp_decoder.DecoderLayer.project[static]": {"gmem": 128}, + "tiny_tp_decoder.TinyTPDecoderLM.layer.decode[static]": {"gmem": 48, "rmem": 16}, + "tiny_tp_decoder.TinyTPDecoderLM.layer.project[static]": {"gmem": 128}, + "tp_all_to_all.TransposeShard.transpose_shard[static]": {"gmem": 256}, + "weighted_twin.Weighted.scaled[static]": {"gmem": 1_348, "rmem": 4}, } +assert set(EXPECTED_MEMORY_PEAKS) == {case.id for case in CASES} + def _aimed(): """The decode example, aimed at one machine.""" @@ -227,20 +369,6 @@ def _every_number_counts_something(result: AnalysisResult) -> None: if record is MemoryMetadata: for level in held.footprint: assert level.peak_bytes >= 0 and level.persistent_bytes >= 0 - rows = [item for item in held.lifetimes if item.level == level.level] - end = max((item.last_used_at for item in rows), default=-1) - expected_peak = max( - ( - sum( - item.bytes - for item in rows - if item.defined_at <= point <= item.last_used_at - ) - for point in range(end + 1) - ), - default=0, - ) - assert level.peak_bytes == expected_peak for item in held.lifetimes: assert item.bytes >= 0 and 0 <= item.defined_at <= item.last_used_at assert " None: assert result.module is owner assert set(result.executed) == set(FAMILIES) assert_performance_contract(result) - expected = EXPECTED_MEMORY_PEAKS.get(case.id, {}) placement = get_metadata(result.function, MemoryMetadata) assert placement is not None observed = {item.level: item.peak_bytes for item in placement.footprint} - assert {level: observed[level] for level in expected} == expected + assert observed == EXPECTED_MEMORY_PEAKS[case.id] @pytest.mark.parametrize("family", FAMILIES) From 840b9a26a44bd1a274391d5f720b6968f1e0e321 Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Mon, 21 Sep 2026 12:48:22 +0800 Subject: [PATCH 9/9] fix(analysis): solve node operand reuse [M4] --- src/tilefoundry/analysis/allocation.py | 531 ++++++++++++++++------- src/tilefoundry/analysis/liveness.py | 19 +- src/tilefoundry/analysis/memory.py | 1 + tests/analysis/test_analyze_at_a_size.py | 6 +- 4 files changed, 401 insertions(+), 156 deletions(-) diff --git a/src/tilefoundry/analysis/allocation.py b/src/tilefoundry/analysis/allocation.py index 26423583..f2b3083c 100644 --- a/src/tilefoundry/analysis/allocation.py +++ b/src/tilefoundry/analysis/allocation.py @@ -3,7 +3,7 @@ from __future__ import annotations from collections import defaultdict -from dataclasses import dataclass +from dataclasses import dataclass, field, replace from itertools import combinations from typing import Protocol @@ -15,10 +15,16 @@ from tilefoundry.ir.hir.tensor.insert_slice import InsertSlice from tilefoundry.ir.hir.tensor.reshape import Reshape from tilefoundry.ir.hir.tensor.slice import Slice +from tilefoundry.ir.types import TensorType +from tilefoundry.ir.types.utils import local_type_of +from tilefoundry.ir.visitor import ExprVisitor +from tilefoundry.utils.isl_utils import equates +from tilefoundry.visitor_registry.access_relation import index_set from .errors import AnalysisError +from .liveness import Liveness from .metadata import ValueLifetime -from .scope import Access, Scope, walk_scopes +from .scope import Access, Scope class _MemoryOptions(Protocol): @@ -43,6 +49,30 @@ class AllocationResult: solver_status: str +@dataclass(frozen=True) +class _OperandConstraint: + """An exact logical relation from one material operand into a result.""" + + operand: Expr + relation: isl.map + + +@dataclass +class _ConstraintContext: + """One memory-level model while the HIR visitor applies logical relations.""" + + current: Scope + liveness: Liveness + values: tuple[AllocationValue, ...] + boxes_by_expr: dict[int, int] + model: cp_model.CpModel + addresses: tuple[cp_model.IntVar, ...] + selected_by_pair: dict[tuple[int, int], list[cp_model.IntVar]] = field( + default_factory=lambda: defaultdict(list) + ) + applied: list[tuple[cp_model.IntVar, int, int]] = field(default_factory=list) + + def _base_value(value: Expr) -> Expr: """Return the material allocation below non-material tensor views.""" while isinstance(value, Call) and isinstance(value.target, (Slice, Reshape)): @@ -50,6 +80,16 @@ def _base_value(value: Expr) -> Expr: return value +def _is_view_of(value: Expr, source: Expr) -> bool: + """Whether ``value`` reaches ``source`` through only non-material views.""" + while True: + if value is source: + return True + if not isinstance(value, Call) or not isinstance(value.target, (Slice, Reshape)): + return False + value = value.args[0] + + def _coverage(accesses: tuple[Access, ...]) -> isl.set | None: """Union the call coordinates on which exact accesses reach one buffer.""" if not accesses or any(not access.exact for access in accesses): @@ -70,113 +110,312 @@ def _access_relation(accesses: tuple[Access, ...]) -> isl.map | None: return result.coalesce() -def _proved_overlap_groups( - values: tuple[AllocationValue, ...], root: Scope -) -> tuple[tuple[int, ...], ...]: - """Prove exact pointwise ties and dynamic-update embedded views. +def _value_domain(value: Expr, scope: Scope) -> isl.set | None: + """The complete loop-aware coordinate domain of one material value.""" + try: + held = local_type_of(value.type) + except (TypeError, ValueError, NotImplementedError): + return None + if not isinstance(held, TensorType): + return None + box = index_set(held.shape) + if box is None: + return None + return scope.domain.flat_product(box).coalesce() - This is scheme B: the exact per-iteration offset remains in ISL. The CP - model only receives the weaker fact that a group can occupy one - result-sized allocation. - """ - by_expr = {id(item.value): index for index, item in enumerate(values)} - carried_ids = { - id(carried) - for scope in walk_scopes(root) - if isinstance(scope.owner, LoopRegion) - for carried in scope.owner.carried_args - } - groups: list[tuple[int, ...]] = [] - seen: set[tuple[int, ...]] = set() - for scope in walk_scopes(root): - for call_id, (call, inputs) in scope.accesses.get("narrow", {}).items(): - if call_id not in by_expr: - continue - recorded_output = scope.outputs.get("narrow", {}).get(call_id) - if recorded_output is None or recorded_output[0] is not call: - continue - outputs = recorded_output[1] - if not outputs or any(not access.exact for access in outputs): - continue - result_index = by_expr[call_id] - result = values[result_index] - output_coverage = _coverage(outputs) - output_relation = _access_relation(outputs) - by_buffer: dict[int, list[Access]] = defaultdict(list) - for access in inputs: - by_buffer[id(access.buffer)].append(access) - - if output_coverage is not None and output_relation is not None: - for buffer_id, accesses in by_buffer.items(): - input_index = by_expr.get(buffer_id) - if input_index is None or input_index == result_index: - continue - source = values[input_index] - if source.lifetime.persistent: - continue - if source.lifetime.bytes != result.lifetime.bytes: - continue - if source.lifetime.last_used_at > result.lifetime.defined_at: - continue - input_relation = _access_relation(tuple(accesses)) - try: - pointwise = input_relation is not None and input_relation.is_equal( - output_relation - ) - except isl.Error: - pointwise = False - group = (result_index, input_index) - if pointwise and group not in seen: - seen.add(group) - groups.append(group) - - if not isinstance(call.target, InsertSlice): - continue +def _operand_to_result_relation( + inputs: isl.map, outputs: isl.map, loop_depth: int +) -> isl.map | None: + """Compose call accesses into ``[loops..., operand] -> [result]``.""" + try: + common = inputs.domain().intersect(outputs.domain()).coalesce() + if common.is_empty(): + return None + inputs = inputs.intersect_domain(common) + outputs = outputs.intersect_domain(common) + call_dims = inputs.dim(isl.dim_type.IN) + if call_dims != outputs.dim(isl.dim_type.IN) or loop_depth > call_dims: + return None + loop_prefix = isl.map.identity(common.get_space().map_from_set()).project_out( + isl.dim_type.OUT, loop_depth, call_dims - loop_depth + ) + operand_with_loops = loop_prefix.flat_range_product(inputs) + result_with_loops = loop_prefix.flat_range_product(outputs) + per_iteration = operand_with_loops.reverse().apply_range(result_with_loops).coalesce() + result = operand_with_loops.reverse().apply_range(outputs).coalesce() + return ( + result + if result.is_single_valued() + and per_iteration.is_single_valued() + and per_iteration.is_injective() + else None + ) + except isl.Error: + return None - dst = _base_value(call.args[0]) - update = _base_value(call.args[1]) - member_ids = (id(call), id(dst), id(update)) - if any(member_id not in by_expr for member_id in member_ids): - continue - indices = tuple(dict.fromkeys(by_expr[member_id] for member_id in member_ids)) - if len(indices) != 3: - continue - result_index, dst_index, update_index = indices - result = values[result_index] - destination = values[dst_index] - patch = values[update_index] - if destination.lifetime.persistent or patch.lifetime.persistent: - continue - if ( - destination.lifetime.last_used_at > result.lifetime.defined_at - and id(destination.value) not in carried_ids - ): - continue - if destination.lifetime.bytes != result.lifetime.bytes: - continue - if patch.lifetime.bytes > result.lifetime.bytes: - continue - if patch.lifetime.last_used_at > result.lifetime.defined_at: - continue - dst_coverage = _coverage(tuple(by_buffer.get(id(dst), ()))) - update_coverage = _coverage(tuple(by_buffer.get(id(update), ()))) - written_coverage = _coverage(outputs) - if any( - coverage is None for coverage in (dst_coverage, update_coverage, written_coverage) - ): - continue +def _complete_operand_relation( + operand: Expr, + inputs: isl.map | None, + outputs: isl.map | None, + scope: Scope, +) -> isl.map | None: + """Return a composed relation only when it covers the whole operand.""" + if inputs is None or outputs is None: + return None + relation = _operand_to_result_relation(inputs, outputs, scope.depth) + expected = _value_domain(operand, scope) + if relation is None or expected is None: + return None + try: + return relation if relation.domain().is_equal(expected) else None + except isl.Error: + return None + + +def _identity_result_relation(node: Call, scope: Scope) -> isl.map | None: + """Map a same-shaped operand to the result while retaining loop axes.""" + domain = _value_domain(node, scope) + if domain is None: + return None + try: + return ( + isl.map.identity(domain.get_space().map_from_set()) + .intersect_domain(domain) + .project_out(isl.dim_type.OUT, 0, scope.depth) + ) + except isl.Error: + return None + + +def _intervals_by_expr(liveness: Liveness) -> dict[int, tuple[int, int]]: + """Index the immutable liveness answer without traversing the HIR again.""" + return { + id(interval.value): (interval.defined_at, interval.last_used_at) + for interval in liveness.intervals + } + + +def _corresponding_carry(source: Expr, operand: Expr, scope: Scope) -> LoopRegion | None: + """Find the carry whose own yield is ``source``, without walking its body.""" + cursor: Scope | None = scope + while cursor is not None: + loop = cursor.owner + if isinstance(loop, LoopRegion): + for slot, carried in enumerate(loop.carried_args): + if carried is operand and slot < len(loop.yield_values): + return loop if _is_view_of(loop.yield_values[slot], source) else None + cursor = cursor.parent + return None + + +def _tie_is_live(source: Expr, operand: Expr, scope: Scope, liveness: Liveness) -> bool: + """Prove that reusing ``operand`` cannot clobber a later ordinary use.""" + intervals = _intervals_by_expr(liveness) + source_interval = intervals.get(id(source)) + operand_interval = intervals.get(id(operand)) + if source_interval is None or operand_interval is None: + return False + if any( + not use.synthetic + and use.at > source_interval[0] + and _base_value(use.value) is operand + and not _is_view_of(use.value, source) + for use in liveness.uses + ): + return False + if operand_interval[1] <= source_interval[0]: + return True + + return _corresponding_carry(source, operand, scope) is not None + + +def _analyze_operand_constraints( + node: Call, scope: Scope, liveness: Liveness +) -> tuple[_OperandConstraint, ...]: + """Prove exact logical relations between one result and its operands.""" + recorded_inputs = scope.accesses.get("narrow", {}).get(id(node)) + recorded_outputs = scope.outputs.get("narrow", {}).get(id(node)) + if ( + recorded_inputs is None + or recorded_inputs[0] is not node + or recorded_outputs is None + or recorded_outputs[0] is not node + ): + return () + inputs = recorded_inputs[1] + outputs = recorded_outputs[1] + output_coverage = _coverage(outputs) + output_relation = _access_relation(outputs) + full_result = _value_domain(node, scope) + if output_coverage is None or output_relation is None or full_result is None: + return () + + by_buffer: dict[int, list[Access]] = defaultdict(list) + operands: dict[int, Expr] = {} + for access in inputs: + key = id(access.buffer) + by_buffer[key].append(access) + operands[key] = access.buffer + + result: list[_OperandConstraint] = [] + seen: set[int] = set() + + try: + covers_result = output_coverage.is_equal(full_result) + except isl.Error: + covers_result = False + if covers_result: + for key, operand in operands.items(): + input_relation = _access_relation(tuple(by_buffer[key])) try: - partitioned = update_coverage.is_equal( - written_coverage - ) and dst_coverage.is_disjoint(written_coverage) + pointwise = input_relation is not None and input_relation.is_equal(output_relation) except isl.Error: - partitioned = False - if partitioned and indices not in seen: - seen.add(indices) - groups.append(indices) - return tuple(groups) + pointwise = False + relation = ( + _complete_operand_relation(operand, input_relation, output_relation, scope) + if pointwise + else None + ) + if relation is not None and _tie_is_live(node, operand, scope, liveness): + result.append(_OperandConstraint(operand, relation)) + seen.add(key) + + if not isinstance(node.target, InsertSlice): + return tuple(result) + + dst = _base_value(node.args[0]) + update = _base_value(node.args[1]) + written = output_coverage + dst_coverage = _coverage(tuple(by_buffer.get(id(dst), ()))) + update_accesses = tuple(by_buffer.get(id(update), ())) + update_coverage = _coverage(update_accesses) + + if dst_coverage is not None: + relation = _identity_result_relation(node, scope) + expected_dst = _value_domain(dst, scope) + try: + partitioned = dst_coverage.is_disjoint(written) and dst_coverage.union( + written + ).coalesce().is_equal(full_result) + complete_identity = ( + relation is not None + and expected_dst is not None + and relation.domain().is_equal(expected_dst) + ) + except isl.Error: + partitioned = complete_identity = False + if ( + partitioned + and complete_identity + and id(dst) not in seen + and _tie_is_live(node, dst, scope, liveness) + ): + result.append(_OperandConstraint(dst, relation)) + seen.add(id(dst)) + + try: + update_matches_write = update_coverage is not None and update_coverage.is_equal(written) + except isl.Error: + update_matches_write = False + update_relation = ( + _complete_operand_relation( + update, + _access_relation(update_accesses), + output_relation, + scope, + ) + if update_matches_write + else None + ) + if ( + update_relation is not None + and id(update) not in seen + and _tie_is_live(node, update, scope, liveness) + ): + result.append(_OperandConstraint(update, update_relation)) + return tuple(result) + + +def _proves_zero_offset(relation: isl.map) -> bool: + """Whether every result coordinate equals the operand coordinate.""" + output_dims = relation.dim(isl.dim_type.OUT) + loop_dims = relation.dim(isl.dim_type.IN) - output_dims + if loop_dims < 0: + return False + try: + return all( + equates(relation, output_axis, loop_dims + output_axis) + for output_axis in range(output_dims) + ) + except isl.Error: + return False + + +def _apply_constraints( + node: Call, + constraints: tuple[_OperandConstraint, ...], + ctx: _ConstraintContext, +) -> None: + """Compile logical relations into this memory level's CP model.""" + result_index = ctx.boxes_by_expr.get(id(node)) + if result_index is None: + return + result = ctx.values[result_index] + for constraint in constraints: + operand_index = ctx.boxes_by_expr.get(id(constraint.operand)) + if operand_index is None or operand_index == result_index: + continue + operand = ctx.values[operand_index] + if operand.lifetime.persistent or not _lifetimes_overlap(result, operand): + continue + if _proves_zero_offset(constraint.relation): + if operand.lifetime.bytes != result.lifetime.bytes: + continue + selected = ctx.model.new_bool_var(f"embedded_{len(ctx.applied)}") + ctx.model.add( + ctx.addresses[operand_index] == ctx.addresses[result_index] + ).only_enforce_if(selected) + else: + if operand.lifetime.bytes > result.lifetime.bytes: + continue + selected = ctx.model.new_bool_var(f"embedded_{len(ctx.applied)}") + ctx.model.add( + ctx.addresses[operand_index] >= ctx.addresses[result_index] + ).only_enforce_if(selected) + ctx.model.add( + ctx.addresses[operand_index] + operand.lifetime.bytes + <= ctx.addresses[result_index] + result.lifetime.bytes + ).only_enforce_if(selected) + pair = tuple(sorted((result_index, operand_index))) + ctx.selected_by_pair[pair].append(selected) + ctx.applied.append((selected, result_index, operand_index)) + + +class AllocationConstraintVisitor(ExprVisitor[None]): + """Visit the HIR DAG once and apply each node/operand placement relation.""" + + def visit_LoopRegion(self, node: LoopRegion, ctx: _ConstraintContext) -> None: + child = next(item for item in ctx.current.children if item.owner is node) + inner = replace(ctx, current=child) + for operand in node.init_args: + self.visit(operand, ctx) + self.visit(node.body, inner) + for operand in node.yield_values: + self.visit(operand, inner) + + def default_visit_leaf( + self, node: Expr, _operands: tuple[None, ...], ctx: _ConstraintContext + ) -> None: + if ( + not isinstance(node, Call) + or id(node) not in ctx.current.accesses.get("narrow", {}) + or id(node) not in ctx.boxes_by_expr + ): + return + constraints = _analyze_operand_constraints(node, ctx.current, ctx.liveness) + _apply_constraints(node, constraints, ctx) def _lifetimes_overlap(left: AllocationValue, right: AllocationValue) -> bool: @@ -186,12 +425,12 @@ def _lifetimes_overlap(left: AllocationValue, right: AllocationValue) -> bool: ) -def _placement_hint( +def _construct_feasible_seed( values: tuple[AllocationValue, ...], - groups: tuple[tuple[int, ...], ...], + applied: list[tuple[cp_model.IntVar, int, int]], limit: int, -) -> tuple[tuple[bool, ...], tuple[int, ...], int] | None: - """Construct one complete feasible suggestion without deciding the model. +) -> tuple[tuple[bool, ...], tuple[int, ...], int]: + """Construct one complete feasible seed from already-applied CP edges. Components exist only while building the hint. The CP model still contains every logical box and independently validates or rejects every suggested @@ -210,15 +449,11 @@ def members(root: int) -> set[int]: allowed_pairs: set[tuple[int, int]] = set() selected: list[bool] = [] - for group in groups: - roots = {find(index) for index in group} + for _choice, left, right in applied: + pair = tuple(sorted((left, right))) + roots = {find(left), find(right)} merged = set().union(*(members(root) for root in roots)) - group_pairs = { - tuple(sorted((left, right))) - for left, right in combinations(group, 2) - if _lifetimes_overlap(values[left], values[right]) - } - permitted = allowed_pairs | group_pairs + permitted = allowed_pairs | {pair} compatible = all( not _lifetimes_overlap(values[left], values[right]) or (left, right) in permitted for left, right in combinations(sorted(merged), 2) @@ -229,7 +464,7 @@ def members(root: int) -> set[int]: root = min(roots) for other in roots: parent[find(other)] = root - allowed_pairs.update(group_pairs) + allowed_pairs.add(pair) components: dict[int, set[int]] = defaultdict(set) for index in range(len(values)): @@ -265,7 +500,7 @@ def conflicts(other: set[int]) -> bool: ) ) if address + size > limit: - return None + raise AnalysisError("allocation: failed to construct a bounded feasible seed") for index in component: addresses[index] = address placed.append((component, address, size)) @@ -276,6 +511,7 @@ def conflicts(other: set[int]) -> bool: def solve_allocation( memory_level: str, values: tuple[AllocationValue, ...], + liveness: Liveness, root: Scope, *, options: _MemoryOptions, @@ -292,10 +528,10 @@ def solve_allocation( model = cp_model.CpModel() peak = model.new_int_var(largest, limit, f"{memory_level}_peak") - addresses = [ + addresses = tuple( model.new_int_var(0, limit - item.lifetime.bytes, f"address_{index}") for index, item in enumerate(values) - ] + ) for address, item in zip(addresses, values, strict=True): model.add(address + item.lifetime.bytes <= peak) @@ -309,22 +545,15 @@ def solve_allocation( if not item.lifetime.persistent: model.add(address >= persistent_end) - groups = _proved_overlap_groups(values, root) - choices_by_pair: dict[tuple[int, int], list[cp_model.IntVar]] = defaultdict(list) - choices: list[tuple[cp_model.IntVar, tuple[int, ...]]] = [] - for group_index, members in enumerate(groups): - selected = model.new_bool_var(f"embedded_{group_index}") - choices.append((selected, members)) - container = members[0] - for member in members[1:]: - model.add(addresses[member] >= addresses[container]).only_enforce_if(selected) - model.add( - addresses[member] + values[member].lifetime.bytes - <= addresses[container] + values[container].lifetime.bytes - ).only_enforce_if(selected) - for left, right in combinations(members, 2): - if _lifetimes_overlap(values[left], values[right]): - choices_by_pair[tuple(sorted((left, right)))].append(selected) + context = _ConstraintContext( + current=root, + liveness=liveness, + values=values, + boxes_by_expr={id(item.value): index for index, item in enumerate(values)}, + model=model, + addresses=addresses, + ) + AllocationConstraintVisitor(root_function=root.owner).visit_function_body(root.owner, context) order_choices: dict[tuple[int, int], tuple[cp_model.IntVar, cp_model.IntVar]] = {} for left, right in combinations(range(len(values)), 2): @@ -344,23 +573,25 @@ def solve_allocation( model.add_bool_or( left_before, right_before, - *choices_by_pair.get((left, right), ()), + *context.selected_by_pair.get((left, right), ()), ) - hint = _placement_hint(values, groups, limit) - if hint is not None: - selected_hints, address_hints, peak_hint = hint - model.add(peak <= peak_hint) - for (selected, _members), suggested in zip(choices, selected_hints, strict=True): - model.add_hint(selected, int(suggested)) - for address, suggested in zip(addresses, address_hints, strict=True): - model.add_hint(address, suggested) - for (left, right), (left_before, right_before) in order_choices.items(): - left_end = address_hints[left] + values[left].lifetime.bytes - right_end = address_hints[right] + values[right].lifetime.bytes - model.add_hint(left_before, int(left_end <= address_hints[right])) - model.add_hint(right_before, int(right_end <= address_hints[left])) - model.add_hint(peak, peak_hint) + selected_hints, address_hints, peak_hint = _construct_feasible_seed( + values, context.applied, limit + ) + model.add(peak <= peak_hint) + for (selected, _result, _operand), suggested in zip( + context.applied, selected_hints, strict=True + ): + model.add_hint(selected, int(suggested)) + for address, suggested in zip(addresses, address_hints, strict=True): + model.add_hint(address, suggested) + for (left, right), (left_before, right_before) in order_choices.items(): + left_end = address_hints[left] + values[left].lifetime.bytes + right_end = address_hints[right] + values[right].lifetime.bytes + model.add_hint(left_before, int(left_end <= address_hints[right])) + model.add_hint(right_before, int(right_end <= address_hints[left])) + model.add_hint(peak, peak_hint) solver = cp_model.CpSolver() solver.parameters.max_time_in_seconds = options.timeout_seconds diff --git a/src/tilefoundry/analysis/liveness.py b/src/tilefoundry/analysis/liveness.py index 0327951b..a9c74e08 100644 --- a/src/tilefoundry/analysis/liveness.py +++ b/src/tilefoundry/analysis/liveness.py @@ -20,11 +20,21 @@ class LiveInterval: last_used_at: int +@dataclass(frozen=True) +class UseEvent: + """One SSA use and whether it exists only to extend structured liveness.""" + + value: Expr + at: int + synthetic: bool = False + + @dataclass(frozen=True) class Liveness: """Definition-ordered intervals on one function-wide event timeline.""" intervals: tuple[LiveInterval, ...] + uses: tuple[UseEvent, ...] timeline_end: int @@ -51,6 +61,7 @@ def __init__(self, function: Function) -> None: self._point = -1 self._states: dict[int, LiveInterval] = {} self._definition_order: list[int] = [] + self._uses: list[UseEvent] = [] for parameter in function.params: self.define(parameter, self.next_event()) for free in _free_vars(function): @@ -69,18 +80,20 @@ def define(self, value: Expr, point: int) -> None: self._states[key] = LiveInterval(value, point, point) self._definition_order.append(key) - def use(self, value: Expr, point: int) -> None: + def use(self, value: Expr, point: int, *, synthetic: bool = False) -> None: """Extend *value* through one consumer event.""" state = self._states.get(id(value)) if state is None: raise ValueError(f"liveness: {type(value).__name__} is used before its definition") self._states[id(value)] = replace(state, last_used_at=max(state.last_used_at, point)) + self._uses.append(UseEvent(value, point, synthetic)) def finish(self) -> Liveness: """Freeze the definition-ordered result.""" states = (self._states[key] for key in self._definition_order) return Liveness( intervals=tuple(states), + uses=tuple(self._uses), timeline_end=self._point, ) @@ -131,7 +144,7 @@ def visit_LoopRegion(self, region: LoopRegion, ctx=None) -> None: exit_use = self.next_event() for source in region.carried_args or (region.body,): - self.use(source, exit_use) + self.use(source, exit_use, synthetic=bool(region.carried_args)) self.define(region, self.next_event()) @@ -144,4 +157,4 @@ def analyze_liveness(function: Function) -> Liveness: return visitor.finish() -__all__ = ["LiveInterval", "Liveness", "analyze_liveness"] +__all__ = ["LiveInterval", "Liveness", "UseEvent", "analyze_liveness"] diff --git a/src/tilefoundry/analysis/memory.py b/src/tilefoundry/analysis/memory.py index 0a58db26..a6f3ae3b 100644 --- a/src/tilefoundry/analysis/memory.py +++ b/src/tilefoundry/analysis/memory.py @@ -500,6 +500,7 @@ def analyze_memory(function: Function, context: AnalyzeContext) -> None: solved = solve_allocation( name, values, + liveness, context.root, options=solver_options, ) diff --git a/tests/analysis/test_analyze_at_a_size.py b/tests/analysis/test_analyze_at_a_size.py index bd607c96..de8a0c1a 100644 --- a/tests/analysis/test_analyze_at_a_size.py +++ b/tests/analysis/test_analyze_at_a_size.py @@ -138,17 +138,17 @@ }, "performance_findings.LocalTier.kernel[static]": {"gmem": 68_096, "rmem": 512}, "prefill_decode_attention.PrefillDecodeAttention.attend[ctx=128,seq=128]": { - "gmem": 1_572_864, + "gmem": 1_310_720, "rmem": 0, "smem": 245_760, }, "qwen3_1_7b_pd.PrefillLayer.layer_decode[ctx_len=128,seq=128]": { - "gmem": 148_521_988, + "gmem": 145_933_316, "rmem": 520, "smem": 65_792, }, "qwen3_1_7b_pd.PrefillLayer.layer_prefill[ctx_len=128,seq=128]": { - "gmem": 178_922_500, + "gmem": 177_087_496, "rmem": 66_560, "smem": 131_072, },