Skip to content

[RFC] Adopt PyTorch per-parameter mixed precision in AutoModel FSDP2 #4021

Description

@HuiyingLi

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_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.

Current AutoModel behavior

fully_shard_by_dtype is our own wrapper around PyTorch's fully_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:

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_bias and 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_fn API selects compute dtype per parameter within one FSDP unit.

At the inspected revision:

  • 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

  1. 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.
  2. 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.
  3. 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.
  4. 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.
  5. 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:

  • Wait for [fsdp2][mp] per-param mixed_precision policy pytorch/pytorch#196626 to merge, then pin a PyTorch build containing the merged change and confirm the final supported API and restrictions.
  • 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.
  • Preserve module-visible input/output dtypes, FP32-sensitive math, parameter names/identity, and checkpoint save/load/resume behavior.
  • 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

  1. Is a Kimi K3 KDA pilot the right first integration target?
  2. After upstream merge, which minimum PyTorch release or post-merge nightly should AutoModel support for this integration?
  3. After validation, should eligible groups select the native path automatically, or should it remain explicitly opt-in?
  4. 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.

Activity

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

Metadata

Metadata

Assignees

Labels

enhancementNew feature or request

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions