Skip to content

[spmd_types] rename to VocabParallelEmbedding - #3631

Open
pianpwk wants to merge 11 commits into
gh/pianpwk/31/basefrom
gh/pianpwk/31/head
Open

pianpwk wants to merge 11 commits into
gh/pianpwk/31/basefrom
gh/pianpwk/31/head

Conversation

@pianpwk

@pianpwk pianpwk commented Jun 11, 2026 •

Copy link
Copy Markdown
Contributor

[ghstack-poisoned]
@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Meta Open Source bot. label Jun 11, 2026
@pianpwk pianpwk changed the title Rename Embedding to VocabParallelEmbedding [spmd_types] rename to VocabParallelEmbedding Jun 11, 2026
pianpwk added 6 commits June 10, 2026 22:22
[ghstack-poisoned]
[ghstack-poisoned]
[ghstack-poisoned]
[ghstack-poisoned]
[ghstack-poisoned]
[ghstack-poisoned]
pianpwk added a commit that referenced this pull request Jun 11, 2026
Main infra PR for spmd_types + titan.

`spmd_types.py`:
- Adds a thread-local DeviceMesh stack, and current_mesh(),
set_current_mesh() helpers, to be set and accessed at model init (weight
parallelize) and runtime (PG access for collectives)
- Various helpers for spmd_types : module-boundary redistributions, spmd
-> DTensor placement translation (for full_dtensor backend),
`mesh_size(axis_name)` helper that returns > 1 when mesh is set & axis
is active. Some of this will be moved to spmd_types in near-term, see
comments.

`parallel_dims.py`:
- For spmd backend, we need 2 views over the world mesh: [pp, dp, cp,
tp] for typechecking, and [pp, dp_replicate, dp_shard, cp, tp] to pass
the unfolded fsdp axes to fully_shard, via DataParallelMeshDims. So we
hold both the full-DTensor-style dense mesh, as well as the
"typechecking" mesh.

`module.py`: module.parallelize() paths for spmd backend: weight init,
local SPMD drop in, input/output redistribution.

`trainer.py/utils.py`: typechecking context, input annotation, spmd
set_current_spmd_mesh, PP + typechecking raises a hard error.

Stack from [ghstack](https://github.com/ezyang/ghstack/tree/0.12.0)
(oldest at bottom):
* #3278
* #3472
* #3632
* #3631
* #3468
* #3471
* __->__ #3253
@pianpwk pianpwk mentioned this pull request Jun 11, 2026
pianpwk added a commit that referenced this pull request Jun 11, 2026
Landed #3253 into wrong branch

Stack from [ghstack](https://github.com/ezyang/ghstack/tree/0.12.0)
(oldest at bottom):
* #3278
* #3472
* #3632
* #3631
* #3468
* #3471
* __->__ #3641
pianpwk added 2 commits June 11, 2026 17:49
[ghstack-poisoned]
[ghstack-poisoned]
pianpwk added a commit that referenced this pull request Jun 12, 2026
switches decoder_sharding.py and llama3/sharding.py to spmd.* types. 
- Fill in src placements to be explicit, where previously we implicitly
relied on DTensor
- `LocalMapConfig(in_grad_placements=...)` carries info only used by
default/full_dtensor backends; spmd_types backend just checks for
presence of config, to switch to local SPMD
- PartitionSpec used once for SP activation: CP/TP shard seq-dim
- SP=off activations are I@TP
- unsharded weights (e.g. norm) are R@TP when SP on, FSDP handles the
gradient AR in DTensor: pytorch/pytorch#181519


Stack from [ghstack](https://github.com/ezyang/ghstack/tree/0.12.0)
(oldest at bottom):
* #3278
* #3472
* #3632
* #3631
* #3468
* __->__ #3471
[ghstack-poisoned]
pianpwk added a commit that referenced this pull request Jun 12, 2026
[ghstack-poisoned]
pianpwk added a commit that referenced this pull request Jun 12, 2026
Adds a custom vocab-parallel Embedding module, to be wired for
spmd_types and DTensor backend. Uses a local SPMD / local_map region,
avoiding DTensor dependence on MaskPartial.

Stack from [ghstack](https://github.com/ezyang/ghstack/tree/0.12.0)
(oldest at bottom):
* #3278
* #3472
* #3632
* #3631
* __->__ #3468

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

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.

3 participants