You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
Timing / dependency: This RFC is the AutoModel adoption plan after pytorch/pytorch#196626 merges. Design discussion can happen now; implementation, native-path validation, and performance experiments start after the upstream change is merged and a PyTorch build containing it is available. The merged API and behavior must be rechecked before implementation.
Direction: adopt PyTorch's per-parameter compute-dtype policy for eligible FSDP2 groups after the upstream merge. This RFC defines the integration and migration plan, while retaining AutoModel's fully_shard_by_dtype for unsupported cases.
Today, a block with BF16 computation plus a few FP32 parameters can require multiple FSDP units. The proposed upstream API could keep those parameters in one unit, reducing dtype-driven wrapping and collective launches. The adoption direction is established; discussion should focus on implementation, rollout, and fallback boundaries. Model-specific decisions about which parameters compute in FP32 remain unchanged.
Upstream: pytorch/pytorch#196626, still open at the time of this proposal. Scope below is based on head dbbdc84256a7d44f43d7c4d48207a2d20258a63b, not just its PR description.
It resolves compute dtype using model-declared FP32 requirements, recorded HF compute-dtype metadata, and the default mixed-precision policy. It then groups by (original/storage dtype, compute dtype) and wraps suitable submodules separately. Distinct-dtype parameters that cannot be isolated into supported submodules produce an explicit error.
A simplified example, with all master weights stored in FP32:
Parameters
Compute dtype
Current arrangement
Proposed eligible arrangement
Projection weights
BF16
FSDP unit A
One shared FSDP unit
Precision-sensitive decay parameters
FP32
FSDP unit B
The same unit, with a per-parameter override
The existing path also handles boundary semantics: FP32 children own any required input conversion, and internal child wrappers preserve natural output dtype. Those behaviors are part of the compatibility contract.
Kimi K3's KimiKDAFp32Params illustrates why this matters: its holder owns A_log and dt_biasand computes the decay gate. Removing a dtype-specific FSDP wrapper must not imply removing that model-owned computation or renaming its parameters.
An override can retain a parameter's original dtype or use the group's default param_dtype; it is not an arbitrary per-parameter dtype map.
Trainable parameters in a group must still have a uniform original dtype.
Mixed compute dtypes require a common effective gradient-reduction dtype, such as FP32.
Forward input casting still follows the group's default param_dtype; preserving an FP32 parameter alone does not make every operation using it execute in FP32.
Overrides must resolve consistently across ranks.
These limits are visible in the group validation. The PR description includes an earlier multi-reduction-dtype / mixed-packing design that has since been narrowed. We should not treat those earlier benchmark numbers as evidence of an AutoModel speedup.
Proposed approach
Preserve precision policy ownership. Reuse the existing compute-dtype resolution rules. Model-specific precision requirements remain in the owning model package; the FSDP integration only applies them.
After upstream merge, add an explicitly gated native path. Use a validated PyTorch build containing the merged API and check eligibility before applying FSDP. Start with uniform FP32 master weights, BF16 bulk computation, FP32 exceptions, and an already-common FP32 reduction policy.
Keep the existing fallback. Older PyTorch builds, mixed original dtypes, unsupported overrides, or incompatible per-group reduction policies continue using the existing helper where supported. Do not silently upcast stored weights or change gradient-reduction precision merely to make a group eligible.
Preserve model and training contracts. Retain holder modules initially, parameter names/identities, checkpoint mappings, ignored-parameter ownership, and input/output casting behavior. Do not change EP/TP ownership boundaries or group together parameters merely because their dtypes match.
After upstream merge, prototype one representative model first. Compare both paths on a small Kimi K3 KDA block, then a reduced-depth training configuration. Expand rollout once correctness and performance compatibility are demonstrated. Keep existing recipe defaults until that evidence and an upstream compatibility policy are agreed.
No new recipe flag or public configuration schema is proposed yet. Configuration belongs in the existing typed FSDP2 strategy configuration if an opt-in surface is needed.
Expected benefits and tradeoffs
Potential benefit: fewer dtype-specific FSDP units and collective launches. For the two-unit example, a compatible single unit can use one all-gather per unshard and one reduce-scatter per synchronized reduction. Backward re-gathers and gradient-accumulation settings still determine total step counts.
Not a communication-volume guarantee: combining launches does not inherently reduce the parameter/gradient payload.
Not a memory or speed guarantee: larger shared units can change parameter residency, prefetch overlap, packing cost, and peak memory. A model-level benchmark is required.
Incomplete replacement: mixed original storage dtypes and different reduction requirements can still require separate units. The native API also does not replace model-specific precision selection.
Validation and acceptance criteria
The following work is proposed for after the upstream merge, not completed:
Compare old and native paths against an independent tiny reference: outputs, input and per-parameter gradients, global gradient norm, and parameters after an optimizer step, with justified dtype-specific tolerances.
Cover FP32 master weights with BF16/FP32 computation; verify original, compute, gradient-accumulation, and reduction dtypes independently.
Verify fallback for mixed storage dtypes, incompatible reduction policies, and unsupported builds without changing user-requested precision.
Exercise one rank and real two-rank FSDP, accumulation with and without synchronization, full/selective activation checkpointing, and resharding on/off. Validate tied/frozen parameters and parameters excluded because another parallelism unit owns them.
Benchmark the same model/configuration on both paths: collective counts/bytes, packing time, peak allocated/reserved memory, and synchronized step time. Include warm-up, repeated measurements, exact versions, and hardware. Start with H100; validate GB200 before claiming benefits there.
Keep initial PR functional coverage within two GPUs. Validate supported HSDP/TP/EP compositions before extending rollout, using scheduled tests where more ranks are necessary.
Implementation and rollout questions
Is a Kimi K3 KDA pilot the right first integration target?
After upstream merge, which minimum PyTorch release or post-merge nightly should AutoModel support for this integration?
After validation, should eligible groups select the native path automatically, or should it remain explicitly opt-in?
What correctness, end-to-end performance, and model-coverage gates should define each rollout stage and retirement of dtype-specific wrapping that the native path makes redundant?
Related: router precision audit #3650. That issue determines model-specific precision contracts; this RFC concerns how FSDP preserves those contracts.
Summary
Timing / dependency: This RFC is the AutoModel adoption plan after pytorch/pytorch#196626 merges. Design discussion can happen now; implementation, native-path validation, and performance experiments start after the upstream change is merged and a PyTorch build containing it is available. The merged API and behavior must be rechecked before implementation.
Direction: adopt PyTorch's per-parameter compute-dtype policy for eligible FSDP2 groups after the upstream merge. This RFC defines the integration and migration plan, while retaining AutoModel's
fully_shard_by_dtypefor unsupported cases.Today, a block with BF16 computation plus a few FP32 parameters can require multiple FSDP units. The proposed upstream API could keep those parameters in one unit, reducing dtype-driven wrapping and collective launches. The adoption direction is established; discussion should focus on implementation, rollout, and fallback boundaries. Model-specific decisions about which parameters compute in FP32 remain unchanged.
Upstream: pytorch/pytorch#196626, still open at the time of this proposal. Scope below is based on head
dbbdc84256a7d44f43d7c4d48207a2d20258a63b, not just its PR description.Current AutoModel behavior
fully_shard_by_dtypeis our own wrapper around PyTorch'sfully_shard.It resolves compute dtype using model-declared FP32 requirements, recorded HF compute-dtype metadata, and the default mixed-precision policy. It then groups by (original/storage dtype, compute dtype) and wraps suitable submodules separately. Distinct-dtype parameters that cannot be isolated into supported submodules produce an explicit error.
A simplified example, with all master weights stored in FP32:
The existing path also handles boundary semantics: FP32 children own any required input conversion, and internal child wrappers preserve natural output dtype. Those behaviors are part of the compatibility contract.
Kimi K3's
KimiKDAFp32Paramsillustrates why this matters: its holder ownsA_loganddt_biasand computes the decay gate. Removing a dtype-specific FSDP wrapper must not imply removing that model-owned computation or renaming its parameters.What the current upstream proposal supports
The proposed
param_dtype_override_fnAPI selects compute dtype per parameter within one FSDP unit.At the inspected revision:
param_dtype; it is not an arbitrary per-parameter dtype map.param_dtype; preserving an FP32 parameter alone does not make every operation using it execute in FP32.These limits are visible in the group validation. The PR description includes an earlier multi-reduction-dtype / mixed-packing design that has since been narrowed. We should not treat those earlier benchmark numbers as evidence of an AutoModel speedup.
Proposed approach
No new recipe flag or public configuration schema is proposed yet. Configuration belongs in the existing typed FSDP2 strategy configuration if an opt-in surface is needed.
Expected benefits and tradeoffs
Validation and acceptance criteria
The following work is proposed for after the upstream merge, not completed:
Implementation and rollout questions
Related: router precision audit #3650. That issue determines model-specific precision contracts; this RFC concerns how FSDP preserves those contracts.