Repository navigation
Conversation
[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 resolved9c05508 to
80d83d1
Compare
|
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`.
80d83d1 to
87f176f
Compare
vasqu
left a comment
There was a problem hiding this comment.
Some small initial questions, not too familiar with the new system so wanna make sure I understand
|
@vasqu decided to rely on dtensor which is cleaner |
vasqu
left a comment
There was a problem hiding this comment.
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() |
There was a problem hiding this comment.
Oh this is really neat, just to be sure it's safe for torch 2.6+? Got bitten too many times recently
…e/transformers into fix-warmup-byte-count-ep-plan
CI recapDashboard: View test results in Grafana |
ArthurZucker
left a comment
There was a problem hiding this comment.
ty, and yeah thanks for pushing for DTensor its a great design for this!
| # 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. |
There was a problem hiding this comment.
| # 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 |
the
get_total_byte_countfunction is used to calculate the total bytes count needed to load the model on each device by reading thetp_plan. This is useful forcaching_allocator_warmupas 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.