@@ -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
15901626def _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+
22612371def _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
22832393def _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+
23542507def _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+
28503014def _nodes_in_dependency_order (
28513015 nodes : Sequence [FusionNode ],
28523016 edges : Sequence [FusionEdge ],
0 commit comments