Skip to content

Count ep and fsdp in caching_allocator_warmup byte calculation - #49384

Open
3outeille wants to merge 6 commits into
mainfrom
fix-warmup-byte-count-ep-plan
Open

3outeille wants to merge 6 commits into
mainfrom
fix-warmup-byte-count-ep-plan

Conversation

@3outeille

@3outeille 3outeille commented Oct 7, 2026 •

Copy link
Copy Markdown
Member

CPU CI GPU run-slow

the get_total_byte_count function is used to calculate the total bytes count needed to load the model on each device by reading the tp_plan. This is useful for caching_allocator_warmup as we want to know how much cache we need to pre-allocate.

Since we support more parallelism now, we can rely on dtensor to get the local tensor and thus its local shape after being distributed.

@3outeille 3outeille changed the title [distributed] Count EP-sharded experts in the allocator warmup, and keep the model's plans resolved [distributed] Count EP-sharded experts in the allocator warmup, and keep the model's plans resolved Oct 7, 2026
@3outeille
3outeille requested a review from vasqu October 7, 2026 10:17
@3outeille 3outeille changed the title [distributed] Count EP-sharded experts in the allocator warmup, and keep the model's plans resolved Count ep_plan in caching_allocator_warmup byte calculation Oct 7, 2026
@3outeille
3outeille force-pushed the fix-warmup-byte-count-ep-plan branch from 9c05508 to 80d83d1 Compare October 7, 2026 10:25
@HuggingFaceDocBuilderDev

Copy link
Copy Markdown

The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update.

…eep the model's plans resolved

`get_total_byte_count` divided a parameter's size only when it matched `model.tp_plan`,
always by the world size. Since the TP and EP plans were decoupled (#48859), EP-sharded
experts no longer show up there unless the TP plan happens to list them too, so an EP-only
layout (`tp_plan={}` or a model without a `base_model_tp_plan`, e.g. DeepSeek V4) counted
every expert at full size on every rank and `caching_allocator_warmup` tried to reserve the
whole model per GPU.

`resolve_parallel_plans` now writes the resolved plans back to `model.tp_plan` and
`model.ep_plan`: after loading they hold the rules that were applied, each plan is empty
when its parallel size is 1, and no parameter is named by both. The warmup then divides a
parameter by `tp_size` when the TP plan names it and by `ep_size` when the EP plan does,
rather than by the world size, which also makes the estimate right on 2-D meshes and under
token dispatch with `tp_size=1`.
@3outeille
3outeille force-pushed the fix-warmup-byte-count-ep-plan branch from 80d83d1 to 87f176f Compare October 7, 2026 10:26

@vasqu vasqu left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Some small initial questions, not too familiar with the new system so wanna make sure I understand

Comment thread src/transformers/modeling_utils.py Outdated
Comment thread src/transformers/distributed/tensor_parallel.py Outdated
@3outeille
3outeille requested a review from ArthurZucker October 7, 2026 12:21
@3outeille
3outeille marked this pull request as draft October 7, 2026 13:25
@3outeille 3outeille changed the title Count ep_plan in caching_allocator_warmup byte calculation Count ep and fsdp in caching_allocator_warmup byte calculation Oct 7, 2026
@3outeille
3outeille marked this pull request as ready for review October 8, 2026 00:56
@3outeille

Copy link
Copy Markdown
Member Author

@vasqu decided to rely on dtensor which is cleaner

@3outeille
3outeille requested a review from vasqu October 8, 2026 00:59

@vasqu vasqu left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Much cleaner now thanks 🫡

total_byte_count[device] += param_byte_count
# Parallelism is applied before loading, so a sharded parameter is already a DTensor placeholder
# Therefore, we can use the local tensor size to determine the total byte count.
numel = param._local_tensor.numel() if is_dtensor(param) else param.numel()

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Oh this is really neat, just to be sure it's safe for torch 2.6+? Got bitten too many times recently

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

yeah

@github-actions

github-actions Bot commented Oct 8, 2026

Copy link
Copy Markdown
Contributor

CI recap

Dashboard: View test results in Grafana
Latest run: 37753906457:1
Result: success | Jobs: 16 | Tests: 192,702 | Failures: 0 | Duration: 15h 19m

@ArthurZucker ArthurZucker left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

ty, and yeah thanks for pushing for DTensor its a great design for this!

Comment on lines +482 to +483
# Mixtral's default plans would also work, since EP rules take precedence over TP rules on shared keys, but the
# plans are spelled out so that readers can easily understand the intent.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
# Mixtral's default plans would also work, since EP rules take precedence over TP rules on shared keys, but the
# plans are spelled out so that readers can easily understand the intent.
# set this to allows default plan to change

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

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants