Skip to content

Commit 2f1ca1d

Browse files
authored
[graph_trainer] Match Eager FSDP bucket order (pytorch#3590)
## Summary Align GraphTrainer FSDP bucket packing order with Eager FSDP2. Eager FSDP2 packs bucket payloads in managed parameter registration order, while GraphTrainer was packing FSDP collective buckets in FX graph execution order. For reduce-scatter, this changes byte offsets inside NCCL buckets. On multinode FSDP groups, NCCL may use different channel ring orders per offset, so the same parameter elements can be reduced through a different bf16 accumulation order and lose bitwise parity. This change derives the FSDP2 module order from traced state FQNs and sorts each FSDP bucket group by that order before delegating to the upstream manual overlap bucketer. The Torchtitan-specific logic is limited to choosing bucket payload order; upstream still owns merging, insertion, wait remapping, and bucketed-node tagging. Verified on 64,128,256 TP=1,8 llama3 8b
1 parent 7c5ea51 commit 2f1ca1d

2 files changed

Lines changed: 54 additions & 0 deletions

File tree

‎torchtitan/experiments/graph_trainer/fsdp_passes.py‎

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,7 @@
3333
_move_overlap_nodes = None
3434
from torch._inductor.fx_passes.overlap_manual_scheduling import (
3535
manual_overlap_bucketing,
36+
ManualOverlapPreservingBucketer,
3637
ManualOverlapScheduler,
3738
)
3839
from torch._inductor.fx_passes.overlap_scheduling import (
@@ -193,6 +194,40 @@ def transformer_block_bucketing_reordering_pass(
193194
return gm
194195

195196

197+
def get_fsdp_param_module_order(state_fqns: list[str]) -> dict[str, int]:
198+
"""Return module order matching FSDP2's first-seen parameter order."""
199+
order: dict[str, int] = {}
200+
for fqn in state_fqns:
201+
if "." not in fqn:
202+
continue
203+
module_fqn = fqn.rsplit(".", 1)[0]
204+
order.setdefault(module_fqn, len(order))
205+
return order
206+
207+
208+
class FSDPParamOrderBucketer(ManualOverlapPreservingBucketer):
209+
"""Pack FSDP buckets in Eager FSDP2 parameter order."""
210+
211+
def __init__(
212+
self,
213+
*args: Any,
214+
fsdp_param_module_order: dict[str, int] | None = None,
215+
**kwargs: Any,
216+
) -> None:
217+
super().__init__(*args, **kwargs)
218+
self.fsdp_param_module_order = fsdp_param_module_order or {}
219+
220+
def _param_order_key(self, node: fx.Node) -> tuple[int, int]:
221+
module_fqn = node.meta.get("custom", {}).get(_MODULE_FQN)
222+
param_idx = self.fsdp_param_module_order.get(module_fqn, len(self.node_idx))
223+
return (param_idx, self.node_idx[node])
224+
225+
def _bucket_group(self, coll_nodes: list[fx.Node]) -> None:
226+
if self.fsdp_param_module_order:
227+
coll_nodes = sorted(coll_nodes, key=self._param_order_key)
228+
return super()._bucket_group(coll_nodes)
229+
230+
196231
class JointManualOverlapScheduler(ManualOverlapScheduler):
197232
"""Manual overlap scheduler for joint forward+backward graphs.
198233
@@ -228,6 +263,7 @@ def __init__(
228263
is_backward_fn: Callable[[fx.Node], bool],
229264
module_stack_fn: Callable[[fx.Node], list[tuple[str, type[Any]]]],
230265
bucket_mode: BucketMode | None = None,
266+
fsdp_param_module_order: dict[str, int] | None = None,
231267
) -> None:
232268
super().__init__(
233269
gm,
@@ -237,6 +273,14 @@ def __init__(
237273
bucket_mode=bucket_mode,
238274
)
239275
self._is_backward_fn = is_backward_fn
276+
effective_bucket_mode = self.bucketer.bucket_mode
277+
self.bucketer = FSDPParamOrderBucketer(
278+
graph=self.graph,
279+
collective_info=self.collective_info,
280+
scheduled=OrderedSet(self.graph.nodes),
281+
bucket_mode=effective_bucket_mode,
282+
fsdp_param_module_order=fsdp_param_module_order,
283+
)
240284

241285
def _manual_bucket_collectives(self) -> None:
242286
"""Bucket per module, splitting by direction to keep fwd/bwd buckets disjoint."""
@@ -402,6 +446,7 @@ def joint_transformer_block_bucketing_reordering_pass(
402446
module_bucket_plans: list[list[str] | str],
403447
insert_overlap_deps: bool = False,
404448
bucket_mode: BucketMode | None = None,
449+
fsdp_param_module_order: dict[str, int] | None = None,
405450
) -> torch.fx.GraphModule:
406451
"""Run joint-graph manual bucketing and reordering.
407452
@@ -423,6 +468,8 @@ def joint_transformer_block_bucketing_reordering_pass(
423468
``preserve_node_ordering`` after the topological sort.
424469
bucket_mode: bucket mode forwarded to the underlying bucketer;
425470
defaults to ``"custom_ops"`` via the parent class.
471+
fsdp_param_module_order: module order derived from traced parameter
472+
FQNs, used to pack FSDP buckets like Eager FSDP2.
426473
"""
427474

428475
def _stack_fn(node: torch.fx.Node) -> list[tuple[str, type]]:
@@ -438,6 +485,7 @@ def _stack_fn(node: torch.fx.Node) -> list[tuple[str, type]]:
438485
is_backward_fn=_is_backward_node,
439486
module_stack_fn=_stack_fn,
440487
bucket_mode=bucket_mode,
488+
fsdp_param_module_order=fsdp_param_module_order,
441489
).run()
442490
overlapped_gm.recompile()
443491
return overlapped_gm

‎torchtitan/experiments/graph_trainer/passes.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@
4242
tlparse_log_graph_pass,
4343
)
4444
from torchtitan.experiments.graph_trainer.fsdp_passes import (
45+
get_fsdp_param_module_order,
4546
joint_transformer_block_bucketing_reordering_pass,
4647
reassign_collective_pgs_pass,
4748
)
@@ -148,6 +149,11 @@ def compile_time_passes(
148149
functools.partial(
149150
joint_transformer_block_bucketing_reordering_pass,
150151
module_bucket_plans=get_default_transformer_block_buckets(n_layers),
152+
# FSDP2 packs buckets in managed parameter order. The traced state
153+
# FQNs preserve that registration order, unlike graph execution order.
154+
fsdp_param_module_order=get_fsdp_param_module_order(
155+
traced_result.state_fqns
156+
),
151157
),
152158
]
153159
if config.parallelism.enable_async_tensor_parallel:

0 commit comments

Comments
 (0)