Skip to content

Commit 7baa322

Browse files
committed
feat(graph): introduce strip_pipeline_stage_metadata function and enhance graph signature handling
- Add `strip_pipeline_stage_metadata` to remove per-variable pipeline-stage hints from GraphDef. - Update `_scan_graph_signature` to utilize the new function for cleaner graph signature processing. - Refactor `_stack_module_states` to handle multiple graph definitions and ensure consistent signatures across modules. - Enhance error handling for heterogeneous layers in `_stack_module_states`. - Improve overall graph definition management in the core module.
1 parent 6a627a6 commit 7baa322

13 files changed

Lines changed: 4812 additions & 647 deletions

File tree

spectrax/core/containers.py

Lines changed: 24 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
import jax.numpy as jnp
2323

2424
from ..sharding.mesh import current_mesh
25-
from .graph import GraphDef, ModuleNode, VarNode, iter_variables
25+
from .graph import GraphDef, ModuleNode, VarNode, iter_variables, strip_pipeline_stage_metadata
2626
from .module import Module, Opaque, _bump_graph_epoch, _graph_epoch
2727
from .paths import str_to_path
2828
from .registry import resolve_class
@@ -60,15 +60,16 @@ def _stack_module_states(items: list[Module], *, context: str) -> tuple[GraphDef
6060
if not items:
6161
raise ValueError(f"{context} requires at least one module")
6262
exports = [export(m) for m in items]
63-
gdef = exports[0][0]
64-
signature = _scan_graph_signature(gdef)
63+
graph_defs = tuple(g for g, _state in exports)
64+
signature = _scan_graph_signature(graph_defs[0])
6565
for index, (other_gdef, _state) in enumerate(exports[1:], start=1):
6666
if _scan_graph_signature(other_gdef) != signature:
6767
raise ValueError(
6868
f"{context} requires every item to have the same graph structure; "
6969
f"item 0 and item {index} differ. Use a Python loop for heterogeneous layers."
7070
)
7171
states = [s for _, s in exports]
72+
gdef = _template_graphdef_without_mixed_stage_metadata(graph_defs)
7273
return gdef, jax.tree.map(lambda *vs: jnp.stack(vs, axis=0), *states)
7374

7475

@@ -125,26 +126,27 @@ def _scan_graph_signature(gdef: GraphDef) -> GraphDef:
125126
Returns:
126127
Return a graph signature suitable for repeated-layer scans.
127128
"""
128-
nodes = []
129-
changed = False
130-
for node in gdef.nodes:
131-
if isinstance(node, VarNode):
132-
metadata = tuple((k, v) for k, v in node.metadata if k != PIPELINE_STAGE_METADATA_KEY)
133-
if metadata != node.metadata:
134-
changed = True
135-
node = VarNode(class_name=node.class_name, collection=node.collection, metadata=metadata)
136-
nodes.append(node)
137-
if not changed:
138-
return gdef
139-
return GraphDef(
140-
nodes=tuple(nodes),
141-
root=gdef.root,
142-
var_refs=gdef.var_refs,
143-
var_canonical=gdef.var_canonical,
144-
shared_paths=gdef.shared_paths,
129+
return strip_pipeline_stage_metadata(gdef)
130+
131+
132+
def _pipeline_stage_metadata_signature(gdef: GraphDef) -> tuple[tuple[tuple[str, object], ...], ...]:
133+
"""Return only the per-variable pipeline-stage metadata from ``gdef``."""
134+
return tuple(
135+
tuple((k, v) for k, v in node.metadata if k == PIPELINE_STAGE_METADATA_KEY)
136+
for node in gdef.nodes
137+
if isinstance(node, VarNode)
145138
)
146139

147140

141+
def _template_graphdef_without_mixed_stage_metadata(graph_defs: tuple[GraphDef, ...]) -> GraphDef:
142+
"""Use the first graph template, stripping stage metadata when it varies."""
143+
template = graph_defs[0]
144+
stage_signature = _pipeline_stage_metadata_signature(template)
145+
if any(_pipeline_stage_metadata_signature(gdef) != stage_signature for gdef in graph_defs[1:]):
146+
return strip_pipeline_stage_metadata(template)
147+
return template
148+
149+
148150
def _scan_graph_topology_signature(gdef: GraphDef) -> tuple[object, ...]:
149151
"""Return the state/child topology used to decide scan compatibility.
150152
@@ -619,11 +621,11 @@ def _scan_static_template_signature(
619621
if len(graph_defs) == 1:
620622
return graph_defs[0]
621623
if family_keys is not None and all(key == family_keys[0] for key in family_keys[1:]):
622-
return graph_defs[0]
624+
return _template_graphdef_without_mixed_stage_metadata(graph_defs)
623625
key = _scan_graph_family_key(graph_defs[0])
624626
if any(_scan_graph_family_key(g) != key for g in graph_defs[1:]):
625627
return None
626-
return graph_defs[0]
628+
return _template_graphdef_without_mixed_stage_metadata(graph_defs)
627629

628630

629631
def _build_scan_plan_from_exports(exports: list[tuple[GraphDef, State]], *, context: str) -> _ScanPlan:

spectrax/core/graph.py

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,7 @@
4747
"find",
4848
"iter_variables",
4949
"live_variables",
50+
"strip_pipeline_stage_metadata",
5051
"tree_state",
5152
"update",
5253
]
@@ -228,6 +229,35 @@ def canonical_path(self, ref_id: int) -> str:
228229
raise KeyError(f"No canonical path for ref_id {ref_id}")
229230

230231

232+
def strip_pipeline_stage_metadata(graphdef: GraphDef) -> GraphDef:
233+
"""Return ``graphdef`` with per-variable pipeline-stage hints removed.
234+
235+
A graph template can represent a stack or scan segment containing several
236+
logical layers/stages. In that case one concrete layer's
237+
``pipeline_stage`` metadata is not a valid owner for the whole template.
238+
"""
239+
from .stage_assignment import PIPELINE_STAGE_METADATA_KEY
240+
241+
nodes: list[Node] = []
242+
changed = False
243+
for node in graphdef.nodes:
244+
if isinstance(node, VarNode):
245+
metadata = tuple((k, v) for k, v in node.metadata if k != PIPELINE_STAGE_METADATA_KEY)
246+
if metadata != node.metadata:
247+
changed = True
248+
node = VarNode(class_name=node.class_name, collection=node.collection, metadata=metadata)
249+
nodes.append(node)
250+
if not changed:
251+
return graphdef
252+
return GraphDef(
253+
nodes=tuple(nodes),
254+
root=graphdef.root,
255+
var_refs=graphdef.var_refs,
256+
var_canonical=graphdef.var_canonical,
257+
shared_paths=graphdef.shared_paths,
258+
)
259+
260+
231261
def _container_kind(m: Module) -> str:
232262
"""Return the :attr:`ModuleNode.container_kind` string for ``m``.
233263

spectrax/core/selector.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -505,8 +505,10 @@ def partition_state(self, module: Module, state: State) -> tuple[State, State]:
505505
writer = state._writers.get((c, path))
506506
if writer is not None:
507507
if is_match:
508+
assert matched_writers is not None
508509
matched_writers[(c, path)] = writer
509510
else:
511+
assert rest_writers is not None
510512
rest_writers[(c, path)] = writer
511513
return State._from_raw(matched_nested, writers=matched_writers), State._from_raw(
512514
rest_nested, writers=rest_writers

spectrax/runtime/mpmd/markers.py

Lines changed: 61 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,8 @@ def model(x):
4141
from jax.interpreters import ad, batching, mlir
4242
from jax.sharding import NamedSharding, PartitionSpec
4343

44+
from spectrax._internal.logging import get_logger
45+
4446
__all__ = [
4547
"cluster_jaxpr_by_markers",
4648
"has_stage_regions",
@@ -55,6 +57,9 @@ def model(x):
5557
"sxstage_region",
5658
]
5759

60+
logger = get_logger(__name__)
61+
_CLUSTER_PRUNE_DIAGNOSTICS = {"logged": 0}
62+
5863

5964
sxstage_iter_p = Primitive("sxstage_iter")
6065
sxstage_iter_p.multiple_results = True
@@ -900,6 +905,42 @@ def _collect_defined_vars_ordered(eqns: list[JaxprEqn]) -> list[Var]:
900905
return ordered
901906

902907

908+
def _prune_stage_jaxpr(sub: Jaxpr) -> Jaxpr:
909+
"""Drop stage equations that do not feed the stage outputs.
910+
911+
``cluster_jaxpr_by_markers`` already computes a minimal-ish output tuple
912+
for each stage boundary, but the stage body is later evaluated through
913+
``jax.core.eval_jaxpr``. Keeping dead equations in that private jaxpr can
914+
force expensive auxiliary computations to survive tracing even when their
915+
values are not part of the stage ABI. This pass performs plain reverse
916+
liveness over the chosen outvars before the stage reaches ``jax.jit``.
917+
"""
918+
needed: set[int] = {id(v) for v in sub.outvars if isinstance(v, Var)}
919+
kept_rev: list[JaxprEqn] = []
920+
for eqn in reversed(sub.eqns):
921+
eqn_outvars = [v for v in eqn.outvars if isinstance(v, Var)]
922+
effects = getattr(eqn, "effects", core.no_effects)
923+
keep = bool(effects) or any(id(v) in needed for v in eqn_outvars)
924+
if not keep:
925+
continue
926+
kept_rev.append(eqn)
927+
for invar in eqn.invars:
928+
if isinstance(invar, Var):
929+
needed.add(id(invar))
930+
931+
pruned_eqns = list(reversed(kept_rev))
932+
pruned_invars = [v for v in sub.invars if not isinstance(v, Var) or id(v) in needed]
933+
if len(pruned_eqns) == len(sub.eqns) and len(pruned_invars) == len(sub.invars):
934+
return sub
935+
return Jaxpr(
936+
constvars=list(sub.constvars),
937+
invars=pruned_invars,
938+
outvars=list(sub.outvars),
939+
eqns=pruned_eqns,
940+
effects=sub.effects,
941+
)
942+
943+
903944
def _stage_region_spans(jaxpr: Jaxpr) -> tuple[tuple[int, int], ...]:
904945
"""Return top-level equation spans covered by stage-region markers.
905946
@@ -1309,6 +1350,8 @@ def collect_remat_eqns(
13091350
defined_order_up_to_at[idx] = list(pre_order)
13101351

13111352
clusters: list[Jaxpr] = []
1353+
dropped_eqns_total = 0
1354+
dropped_invars_total = 0
13121355
for idx, (start, end) in enumerate(itertools.pairwise(boundaries)):
13131356
base_eqns = [
13141357
e for eqn_idx, e in enumerate(jaxpr.eqns[start:end], start=start) if eqn_idx not in boundary_marker_positions
@@ -1369,8 +1412,25 @@ def collect_remat_eqns(
13691412
eqns=eqns,
13701413
effects=core.no_effects,
13711414
)
1372-
clusters.append(sub)
1415+
pruned = _prune_stage_jaxpr(sub)
1416+
dropped_eqns_total += len(sub.eqns) - len(pruned.eqns)
1417+
dropped_invars_total += len(sub.invars) - len(pruned.invars)
1418+
clusters.append(pruned)
13731419
del idx
1420+
if dropped_eqns_total and _CLUSTER_PRUNE_DIAGNOSTICS.get("logged", 0) < 5:
1421+
try:
1422+
process_index = jax.process_index()
1423+
except Exception:
1424+
process_index = -1
1425+
if process_index == 0:
1426+
logger.warning(
1427+
"SpectraX MPMD marker clustering pruned %d dead stage equation(s) "
1428+
"and %d unused stage input(s) across %d stage(s).",
1429+
dropped_eqns_total,
1430+
dropped_invars_total,
1431+
len(clusters),
1432+
)
1433+
_CLUSTER_PRUNE_DIAGNOSTICS["logged"] = _CLUSTER_PRUNE_DIAGNOSTICS.get("logged", 0) + 1
13741434
return clusters
13751435

13761436

0 commit comments

Comments
 (0)