Skip to content

Commit 9919468

Browse files
committed
feat: implement Path C fusion physical ABI and add compilation receipt infrastructure with associated reporting tools
1 parent f222b75 commit 9919468

23 files changed

Lines changed: 17682 additions & 698 deletions

cppmega_mlx/recipes/model_factory.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -260,11 +260,11 @@ def path_c_bricks(self) -> tuple[dict[str, str], ...]:
260260

261261
return tuple(
262262
{
263-
"name": f"{self.name}_brick_{index}_{symbol}",
264-
"kind": symbol,
265-
"route_symbol": symbol,
263+
"name": f"{self.name}_brick_{index}_{layer.symbol}",
264+
"kind": layer.role,
265+
"route_symbol": layer.symbol,
266266
}
267-
for index, symbol in enumerate(self.pattern)
267+
for index, layer in enumerate(self.expanded_pattern.layers)
268268
)
269269

270270
def build_model(

cppmega_mlx/runtime/path_c_fusion.py

Lines changed: 197 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -290,6 +290,7 @@ class FusionEdge:
290290
output: str
291291
consumer: str
292292
input: str
293+
lifetime: str = "internal"
293294

294295

295296
@dataclass(frozen=True)
@@ -452,6 +453,7 @@ class CompiledPathCRegion:
452453

453454
plan: FusionCompilePlan
454455
artifact: object | None = None
456+
lowered_module: object | None = None
455457

456458

457459
@dataclass(frozen=True)
@@ -619,7 +621,9 @@ def build(self) -> PathCFusionRegion:
619621
)
620622

621623

622-
def _mamba3_fp8_train_surfaces() -> tuple[FusionKernelSurface, ...]:
624+
def _legacy_mamba3_fp8_train_diagnostic_surfaces() -> tuple[FusionKernelSurface, ...]:
625+
"""Return the old incomplete train-block graph for semantic-blocker tests."""
626+
623627
return (
624628
FusionKernelSurface.path_c(
625629
name="mamba3_scan",
@@ -731,6 +735,17 @@ def _build_path_c_model_region_from_route_symbols(
731735
acceptance_tags: Sequence[str],
732736
acceptance_fixture_abi: bool,
733737
) -> PathCFusionRegion:
738+
if not acceptance_fixture_abi:
739+
return build_path_c_model_region_from_bricks(
740+
region_name=region_name,
741+
bricks=_path_c_model_bricks_from_route_symbols(route_symbols),
742+
z3_sync=z3_sync,
743+
include_backward=include_backward,
744+
shape_env=shape_env,
745+
model_config=model_config,
746+
acceptance_tags=acceptance_tags,
747+
)
748+
734749
resolved_shape_env = shape_env or _path_c_model_shape_env_from_config(model_config)
735750
metadata: dict[str, Any] = {
736751
"path_c_route_symbols": tuple(route_symbols),
@@ -741,7 +756,7 @@ def _build_path_c_model_region_from_route_symbols(
741756
metadata["path_c_model_shape_env"] = resolved_shape_env
742757
region = build_path_c_fusion_region(
743758
region_name=region_name,
744-
surfaces=_path_c_model_surfaces_from_route_symbols(
759+
surfaces=_path_c_acceptance_fixture_surfaces_from_route_symbols(
745760
route_symbols,
746761
shared_acceptance_abi=acceptance_fixture_abi,
747762
),
@@ -897,6 +912,7 @@ def build_path_c_model_region_from_bricks(
897912
"path_c_route_symbols": tuple(brick.route_symbol for brick in resolved_bricks),
898913
"path_c_bricks": tuple(_path_c_brick_metadata(brick) for brick in resolved_bricks),
899914
"path_c_acceptance_tags": tuple(str(tag) for tag in acceptance_tags),
915+
"path_c_acceptance_fixture_abi": False,
900916
}
901917
if resolved_shape_env is not None:
902918
metadata["path_c_model_shape_env"] = resolved_shape_env
@@ -1094,9 +1110,12 @@ def _append_path_c_model_segment(
10941110
return
10951111
end = segment_start + len(segment) - 1
10961112
regions.append(
1097-
build_path_c_model_region_from_route_symbols(
1113+
build_path_c_model_region_from_bricks(
10981114
region_name=f"{region_prefix}_{segment_start}_{end}",
1099-
route_symbols=segment,
1115+
bricks=_path_c_model_bricks_from_route_symbols(
1116+
segment,
1117+
start_index=segment_start,
1118+
),
11001119
z3_sync=z3_sync,
11011120
include_backward=include_backward,
11021121
shape_env=shape_env,
@@ -1132,7 +1151,24 @@ def _append_path_c_model_brick_segment(
11321151
)
11331152

11341153

1135-
def _path_c_model_surfaces_from_route_symbols(
1154+
def _path_c_model_bricks_from_route_symbols(
1155+
route_symbols: Sequence[str],
1156+
*,
1157+
start_index: int = 0,
1158+
) -> tuple[PathCModelBrick, ...]:
1159+
return tuple(
1160+
PathCModelBrick(
1161+
name=f"route_{start_index + index}_{normalized}",
1162+
kind=normalized,
1163+
route_symbol=normalized,
1164+
)
1165+
for index, normalized in enumerate(
1166+
_normalized_route_symbol(symbol) for symbol in route_symbols
1167+
)
1168+
)
1169+
1170+
1171+
def _path_c_acceptance_fixture_surfaces_from_route_symbols(
11361172
route_symbols: Sequence[str],
11371173
*,
11381174
shared_acceptance_abi: bool,
@@ -1588,9 +1624,7 @@ def _surface_from_node(node: FusionNode) -> FusionKernelSurface:
15881624

15891625

15901626
def _aot_backward_surface_for_node(node: FusionNode) -> FusionKernelSurface:
1591-
inputs = _grad_buffer_names(node.outputs)
1592-
if node.op_name == "residual_rmsnorm":
1593-
inputs = (*inputs, *node.inputs)
1627+
inputs = (*_grad_buffer_names(node.outputs), *node.inputs)
15941628
return FusionKernelSurface.path_c(
15951629
name=f"{node.name}_bwd",
15961630
op_name=f"{node.op_name}_bwd",
@@ -1628,12 +1662,38 @@ def build_path_c_aot_autograd_region(
16281662
)
16291663

16301664

1631-
def build_mamba3_fp8_train_region() -> PathCFusionRegion:
1632-
"""Return the high-level Path C train-block template requested for 1B."""
1665+
def build_mamba3_fp8_train_region(
1666+
*,
1667+
route_symbols: Sequence[str] = ("M", "R", "A"),
1668+
include_backward: bool = False,
1669+
model_config: Any | None = None,
1670+
) -> PathCFusionRegion:
1671+
"""Return a route-derived Path C train-block convenience region.
16331672
1634-
return build_path_c_fusion_region(
1673+
This helper exists for compatibility with older tests and diagnostics. It
1674+
deliberately uses the same model-brick lowering path as real models instead
1675+
of hand-assembling the train block.
1676+
"""
1677+
1678+
if model_config is None:
1679+
from cppmega_mlx.recipes.model_factory import local_gb10_quarter_profile
1680+
1681+
model_config = local_gb10_quarter_profile().hybrid_config()
1682+
1683+
return build_path_c_model_region_from_route_symbols(
16351684
region_name="mamba3_fp8_train_block",
1636-
surfaces=_mamba3_fp8_train_surfaces(),
1685+
route_symbols=route_symbols,
1686+
include_backward=include_backward,
1687+
model_config=model_config,
1688+
)
1689+
1690+
1691+
def _build_legacy_mamba3_fp8_train_diagnostic_region() -> PathCFusionRegion:
1692+
"""Return the legacy static graph used only to exercise blocker reporting."""
1693+
1694+
return build_path_c_fusion_region(
1695+
region_name="mamba3_fp8_train_block_legacy_diagnostic",
1696+
surfaces=_legacy_mamba3_fp8_train_diagnostic_surfaces(),
16371697
z3_sync=Z3SyncSpec.minimize_sync_async(),
16381698
)
16391699

@@ -2101,7 +2161,11 @@ def compile_path_c_region(
21012161
schedule_contract=schedule_contract,
21022162
)
21032163
if tilelang_lowerer is not None:
2104-
return CompiledPathCRegion(plan=plan, artifact=artifact)
2164+
return CompiledPathCRegion(
2165+
plan=plan,
2166+
artifact=artifact,
2167+
lowered_module=getattr(tilelang_result, "lowered_module", None),
2168+
)
21052169
if compiler is None:
21062170
return plan
21072171
return CompiledPathCRegion(plan=plan, artifact=compiler(plan))
@@ -2247,6 +2311,19 @@ def _tilelang_optimizer_for(
22472311
outputs=node.outputs,
22482312
attrs=_tilelang_node_attrs(node),
22492313
)
2314+
workspace_edge_buffers = _declared_workspace_edge_buffers(schedule_template)
2315+
node_by_name = {node.name: node for node in region.nodes}
2316+
for edge in region.edges:
2317+
optimizer.connect(
2318+
edge.producer,
2319+
edge.consumer,
2320+
buffer=edge.input,
2321+
lifetime=_tilelang_edge_lifetime_for(
2322+
edge,
2323+
node_by_name=node_by_name,
2324+
workspace_edge_buffers=workspace_edge_buffers,
2325+
),
2326+
)
22502327
return optimizer
22512328

22522329

@@ -2258,6 +2335,39 @@ def _tilelang_node_attrs(node: FusionNode) -> dict[str, str]:
22582335
return attrs
22592336

22602337

2338+
def _declared_workspace_edge_buffers(
2339+
schedule_template: Callable[[Any], Any] | None,
2340+
) -> tuple[str, ...]:
2341+
if schedule_template is None:
2342+
return ()
2343+
declared = getattr(
2344+
schedule_template,
2345+
"_cppmega_path_c_workspace_edge_buffers",
2346+
(),
2347+
)
2348+
return tuple(str(name) for name in declared)
2349+
2350+
2351+
def _tilelang_edge_lifetime_for(
2352+
edge: FusionEdge,
2353+
*,
2354+
node_by_name: Mapping[str, FusionNode],
2355+
workspace_edge_buffers: Sequence[str],
2356+
) -> str:
2357+
if not workspace_edge_buffers:
2358+
return edge.lifetime
2359+
producer = node_by_name[edge.producer]
2360+
consumer = node_by_name[edge.consumer]
2361+
if (
2362+
producer.op_name == "attention_qkv_projection"
2363+
and consumer.op_name == "sparse_mla_fp8_apply"
2364+
and _canonical_path_c_edge_buffer_name(edge.input)
2365+
in set(workspace_edge_buffers)
2366+
):
2367+
return "workspace"
2368+
return edge.lifetime
2369+
2370+
22612371
def _region_shape_env_payload(region: PathCFusionRegion) -> dict[str, Any]:
22622372
metadata = region.metadata if isinstance(region.metadata, Mapping) else {}
22632373
shape_env = metadata.get("path_c_model_shape_env")
@@ -2282,7 +2392,7 @@ def _region_shape_cache_key_parts(region: PathCFusionRegion) -> tuple[str, ...]:
22822392

22832393
def _schedule_contract_for(region: PathCFusionRegion) -> FusionScheduleContract:
22842394
internal_buffers = tuple(
2285-
dict.fromkeys(edge.input for edge in region.edges)
2395+
dict.fromkeys(edge.input for edge in region.edges if edge.lifetime == "internal")
22862396
)
22872397
internal_buffer_set = set(internal_buffers)
22882398
external_buffers: list[str] = []
@@ -2348,9 +2458,52 @@ def attested_schedule_template(template_region: Any) -> Any:
23482458
attested_schedule_template._cppmega_path_c_required_real_abi_inputs = tuple(
23492459
required_real_abi_inputs
23502460
)
2461+
workspace_edge_buffers = _workspace_edge_buffers_for_attested_template(
2462+
schedule_template,
2463+
region=region,
2464+
required_real_abi_inputs=required_real_abi_inputs,
2465+
)
2466+
attested_schedule_template._cppmega_path_c_workspace_edge_buffers = (
2467+
workspace_edge_buffers
2468+
)
23512469
return attested_schedule_template
23522470

23532471

2472+
def _workspace_edge_buffers_for_attested_template(
2473+
schedule_template: Callable[[Any], Any],
2474+
*,
2475+
region: PathCFusionRegion,
2476+
required_real_abi_inputs: Sequence[str],
2477+
) -> tuple[str, ...]:
2478+
declared = tuple(
2479+
str(name)
2480+
for name in getattr(
2481+
schedule_template,
2482+
"_cppmega_path_c_workspace_edge_buffers",
2483+
(),
2484+
)
2485+
)
2486+
if declared:
2487+
return declared
2488+
if not set(MAMBA3_FP8_TRAIN_REQUIRED_REAL_ABI_INPUTS).issubset(
2489+
set(required_real_abi_inputs)
2490+
):
2491+
return ()
2492+
if not _region_has_attention_kv_workspace_edges(region):
2493+
return ()
2494+
return ("kv_fp8", "kv_scale")
2495+
2496+
2497+
def _region_has_attention_kv_workspace_edges(region: PathCFusionRegion) -> bool:
2498+
node_by_name = {node.name: node for node in region.nodes}
2499+
return any(
2500+
node_by_name[edge.producer].op_name == "attention_qkv_projection"
2501+
and node_by_name[edge.consumer].op_name == "sparse_mla_fp8_apply"
2502+
and _canonical_path_c_edge_buffer_name(edge.input) in {"kv_fp8", "kv_scale"}
2503+
for edge in region.edges
2504+
)
2505+
2506+
23542507
def _declared_schedule_contract_key(
23552508
schedule_template: Callable[[Any], Any] | None,
23562509
) -> str:
@@ -2493,25 +2646,6 @@ def _schedule_contract_status_for(
24932646
declared_required_real_abi_inputs=declared_required_real_abi_inputs,
24942647
missing_real_abi_inputs=missing_real_abi_inputs,
24952648
)
2496-
if declared_schedule_id not in _TRUSTED_PRODUCTION_SCHEDULE_IDS:
2497-
return FusionScheduleContractStatus(
2498-
name=contract.name,
2499-
key=contract.key,
2500-
status="untrusted_production_schedule",
2501-
reason=(
2502-
"schedule template declares production implementation, but its "
2503-
"production_schedule_id is not trusted by this build"
2504-
),
2505-
op_signature=contract.op_signature,
2506-
required_internal_buffers=contract.required_internal_buffers,
2507-
required_external_buffers=contract.required_external_buffers,
2508-
shape_env_key=contract.shape_env_key,
2509-
declared_key=declared_key,
2510-
declared_implementation_kind=declared_kind,
2511-
declared_schedule_id=declared_schedule_id,
2512-
declared_required_real_abi_inputs=declared_required_real_abi_inputs,
2513-
missing_real_abi_inputs=missing_real_abi_inputs,
2514-
)
25152649
if missing_real_abi_inputs:
25162650
return FusionScheduleContractStatus(
25172651
name=contract.name,
@@ -2818,6 +2952,7 @@ def _infer_edges(nodes: Sequence[FusionNode]) -> tuple[FusionEdge, ...]:
28182952
producer_by_output: dict[str, str] = {}
28192953
ambiguous_outputs: dict[str, tuple[str, str]] = {}
28202954
edges: list[FusionEdge] = []
2955+
node_by_name = {node.name: node for node in nodes}
28212956
for node in nodes:
28222957
for output_name in node.outputs:
28232958
existing = producer_by_output.get(output_name)
@@ -2842,11 +2977,40 @@ def _infer_edges(nodes: Sequence[FusionNode]) -> tuple[FusionEdge, ...]:
28422977
output=input_name,
28432978
consumer=node.name,
28442979
input=input_name,
2980+
lifetime=_inferred_edge_lifetime(
2981+
producer_node=node_by_name[producer],
2982+
consumer_node=node,
2983+
buffer_name=input_name,
2984+
),
28452985
)
28462986
)
28472987
return tuple(edges)
28482988

28492989

2990+
def _canonical_path_c_edge_buffer_name(buffer_name: str) -> str:
2991+
name = str(buffer_name)
2992+
if name.endswith("_kv_fp8"):
2993+
return "kv_fp8"
2994+
if name.endswith("_kv_scale"):
2995+
return "kv_scale"
2996+
return name
2997+
2998+
2999+
def _inferred_edge_lifetime(
3000+
*,
3001+
producer_node: FusionNode,
3002+
consumer_node: FusionNode,
3003+
buffer_name: str,
3004+
) -> str:
3005+
if (
3006+
producer_node.op_name == "attention_qkv_projection"
3007+
and consumer_node.op_name == "sparse_mla_fp8_apply"
3008+
and _canonical_path_c_edge_buffer_name(buffer_name) in {"kv_fp8", "kv_scale"}
3009+
):
3010+
return "workspace"
3011+
return "internal"
3012+
3013+
28503014
def _nodes_in_dependency_order(
28513015
nodes: Sequence[FusionNode],
28523016
edges: Sequence[FusionEdge],

0 commit comments

Comments
 (0)