Skip to content

Commit 9362f01

Browse files
authored
Clarify logical rules handling in pipeline parallelism
1 parent b5bb27b commit 9362f01

1 file changed

Lines changed: 2 additions & 0 deletions

File tree

src/maxtext/layers/attentions.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -556,6 +556,8 @@ def __init__(
556556
)
557557

558558
def _logical_to_mesh_axes(self, logical_name):
559+
# Pipeline parallelism uses context managers for logical rules instead of the config,
560+
# so pass None to ensure `logical_to_mesh_axes` defers to using the current Flax context manager
559561
logical_rules = None if self.config.using_pipeline_parallelism else self.config.logical_axis_rules
560562
return logical_to_mesh_axes(logical_name, mesh=self.mesh, rules=logical_rules)
561563

0 commit comments

Comments
 (0)