Skip to content

Revert "Always Pre-Split Microbatches for PP" - #4042

Open
wconstab wants to merge 1 commit into
mainfrom
revert-3856-gh/sanketpurandare/3/head
Open

Revert "Always Pre-Split Microbatches for PP"#4042
wconstab wants to merge 1 commit into
mainfrom
revert-3856-gh/sanketpurandare/3/head

Conversation

@wconstab

@wconstab wconstab commented Jul 31, 2026

Copy link
Copy Markdown
Contributor

Revert #3856 because it broke PP model tests by passing pre-split microbatches through pp_schedule.step() as arg_mbs / kwarg_mbs / target_mbs.

step() is the public whole-batch API and re-splits its inputs internally. The pre-split arguments are only accepted by the private _step_microbatches() path, so the scheduler sees one
microbatch and fails with errors like ValueError: Expecting 8 arg_mbs but got 1.

This restores main while we work out a proper pre-split PP API/path for varlen metadata.

@sanketpurandare can you confirm? Possibly this was coordinated with an upstream pytorch change that we're not using yet?

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Meta Open Source bot. label Jul 31, 2026
@sanketpurandare

Copy link
Copy Markdown
Contributor

@wconstab the upstream change had already been landed: pytorch/pytorch#188500

wconstab added a commit that referenced this pull request Aug 1, 2026
Fixes the shared GraphTrainer failure seen on main and the PP revert PR:

- Main GraphTrainer 8 GPU Integration Tests:
https://github.com/pytorch/torchtitan/actions/runs/30675722021/job/91302425570
- Revert PR #4042 GraphTrainer 8 GPU Integration Tests:
https://github.com/pytorch/torchtitan/actions/runs/30672227677/job/91292165456

Both fail in
`TestMetadataPropagation::test_backward_nodes_have_stack_trace` with
`AssertionError: 23 != 24`, while `bwd_nodes_missing_stack_trace = []`.

The exact number of eligible FX/autograd nodes is not the behavior this
test needs to lock down; it can change as PyTorch
tracing/decomposition/autograd internals change. The invariant is that
backward nodes corresponding to forward nodes with stack traces also
have stack traces. This PR keeps that assertion and only replaces the
brittle exact graph-shape count with a sanity check that the test
actually examined at least one backward node.

Validation:
- `python3 -m py_compile
torchtitan/experiments/graph_trainer/tests/test_trace_module.py`

Note: local `pytest` was unavailable in the default Python environment
on my host, so CI should be used for the targeted runtime check. This PR
does not address the separate H100 GraphTrainer integration failures.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

ci-no-td ciflow/8gpu CLA Signed This label is managed by the Meta Open Source bot.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants