Guard the cross-attention FSDP2 wrap for pipeline stage modules - #187
Merged
Merged
Conversation
Codecov Report✅ All modified and coverable lines are covered by tests.
🚀 New features to boost your workflow:
|
amazloumi
previously approved these changes
Aug 26, 2026
amazloumi
left a comment
Member
There was a problem hiding this comment.
LGTM. none of comments are blockers.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
apply_fsdp2receives aPipelineStageModuleunder pipeline parallelism, but _fsdp_wrap_transformer_blocks read transformer.cross_attention_layers unconditionally. Only Transformer defines that attribute, so every pp > 1 run died with AttributeError before the first step.getattr.apply_fsdp2receives aPipelineStageModuleunder PP, which has nocross_attention_layers, so everypp > 1run died withAttributeErrorbefore the first step.Transformer | PipelineStageModule. With the annotation widened and the guard removed,pyrightflags the offending line, so CI catches the next one.apply_fsdp2on a stage module. CPU only, runs in CI today, fails on the unpatched tree.Testing
uv run ruff check kempnerforge/ tests/passesuv run ruff format --check kempnerforge/ tests/ scripts/passesuv run pyright kempnerforge/passes (0 errors)uv run pytest tests/unit/ -v --timeout=60passesuv run torchrun --nproc_per_node=4 -m pytest tests/distributed/ -vuv run pytest tests/e2e/ --e2e -vCloses #185