Skip to content

Feat (brevitas_examples/llm): support for FSDP + rotation - #1610

Open
Giuseppe5 wants to merge 19 commits into
Xilinx:masterfrom
Giuseppe5:fsdp_rotation_experimental
Open

Giuseppe5 wants to merge 19 commits into
Xilinx:masterfrom
Giuseppe5:fsdp_rotation_experimental

Conversation

@Giuseppe5

@Giuseppe5 Giuseppe5 commented Sep 6, 2026 •

Copy link
Copy Markdown
Collaborator

FSDP2 + trainable rotations: replicated RotationBank and a robust unshard guard

Stack:

Summary

This PR makes trainable rotation optimization work under FSDP2 for large models.
It adds two independent pieces that together make
the combination correct and stable:

  1. A root-owned, FSDP-ignored RotationBank that owns each logical rotation
    exactly once and keeps it replicated (not sharded) across data-parallel ranks,
    with a single packed gradient all-reduce and deterministic optimization.
  2. A universal FSDP2 unshard synchronization guard (fsdp_sync_all_unshards)
    that removes a ROCm unshard/consume stream race which crashed training with large models.

The end result: rotation matrices are trained data-parallel with identical replicas
on every rank, while quantized weights remain FSDP-sharded, and the run no longer
takes a nondeterministic SIGSEGV during forward, backward, or gradient-checkpoint
recompute.

All new behaviour is opt-in through configuration; the default paths are unchanged.


Motivation

Trainable rotations (SpinQuant/QuaRot-style) share one rotation matrix across many
layers. Under FSDP2 this creates two problems:

  • Ownership / sharding. A rotation is a small orthogonal matrix that must stay
    replicated and identical on every rank; sharding it (or letting each FSDP unit
    hold an aliased copy) breaks the math and makes replicas drift. FSDP's per-module
    parameter handling and per-alias registration fight this.
  • Stream lifetime races. FSDP2 all-gathers a parameter's unsharded storage, a
    consumer reads it, and the post-forward/backward reshard frees it. On ROCm the
    consume/free ordering was not enforced for every consumer, producing
    nondeterministic SIGSEGV crashes. These surfaced at 14B/32B in the quantized-weight
    forward and in the backward gradient-checkpoint recompute (crash stacks contained
    no rotation frames — the hazard is generic to FSDP unshard/consume, not
    rotation-specific).

Part 1 — Replicated RotationBank

Ownership model

src/brevitas/utils/parametrization_utils.py

  • RotationBank (nn.Module): a single root-owned module holding each logical
    rotation once in a ParameterDict (rotation_0000, rotation_0001, ...), exposed
    under state_dict key _brevitas_rotation_bank.rotations.rotation_XXXX.
  • _RotationBankHandle: a plain, non-registering holder that each consumer
    parametrization keeps instead of re-registering the shared nn.Parameter. This is
    what prevents FSDP from seeing N aliases of the same rotation.
  • RotationWeightParametrization now resolves rot_mat through the bank handle
    (__getattr__ / bind_rotation_bank), so whole-model .to(), deepcopy, and
    state-dict load all operate through the single registered bank parameter.
  • ensure_rotation_bank / get_rotation_bank / prune_rotation_bank: build the
    bank in deterministic module-traversal order, reject a parameter that appears in
    multiple logical groups, and prune fused/dead rotations after
    fuse_parametrizations.

Coordinator

src/brevitas_examples/llm/llm_quant/fsdp_rotation.py — FSDPRotationCoordinator

  • FSDP-ignores the bank by appending it to the plugin's ignored_modules
    (supporting both list and regex forms) and pins it to the accelerator device, so
    FSDP never sharded the rotations.
  • Broadcasts the initial bank from rank 0 so all replicas start identical.
  • consolidate_gradients(): at Accelerator.sync_gradients, flattens all bank
    gradients in a fixed insertion order, packs a per-rotation presence vector into
    the same buffer, does one all_reduce, averages by world size, and restores
    grad=None for globally-unused rotations. One collective for the whole bank, no
    per-step owner broadcasts.
  • clip_grad_norm_(): a mixed-collection clip that combines FSDP DTensor
    gradients and replicated plain-tensor rotation gradients into one correct global
    norm.
  • check_replica_consistency() (opt-in, diagnostic): all-gathers each bank
    parameter and asserts bit-for-bit equality across ranks, raising with the max abs
    difference and offending rank on divergence.

Deterministic optimization

src/brevitas/optim/cailey_sgd.py

  • CaileySGD uses a deterministic periodic QR retraction (qr_retraction,
    every 100 steps) so rank-local RNG state cannot make replicas diverge — a
    prerequisite for keeping the replicated bank identical without per-step broadcasts.

Supporting fixes

  • src/brevitas_examples/common/accelerate_utils/accelerate.py: remove_hooks now
    deletes the stale hf_device_map (including the bank entry) so Accelerate FSDP
    preparation does not choke; the bank is pinned to LOCAL_RANK during pre-FSDP
    dispatch.
  • src/brevitas/graph/equalize.py: rotation group IDs and bank pruning in
    fuse_parametrizations.

Part 2 — Universal FSDP2 unshard synchronization

src/brevitas_examples/llm/llm_quant/fsdp_workarounds.py — enable_fsdp_unshard_sync

What it does

Wraps the single chokepoint every unshard flows through,
FSDPParamGroup.wait_for_unshard, and after the all-gather copy-out
unconditionally synchronizes the compute stream for every unshard in every
training state
(forward, pre-forward, pre-backward, and backward
gradient-checkpoint recompute). This closes the consume/free race for any
consumer — quantized-weight forward, rotation, checkpoint recompute — without
enumerating consumers or training states.

Why "universal" instead of selective

The earlier selective fences (fsdp_sync_pre_backward_unshard /
fsdp_sync_forward_unshard) synchronized only at enumerated _training_state
values. That was whack-a-mole: each larger model surfaced a new unsynced unshard
site (e.g. the backward checkpoint recompute, whose state is neither FORWARD nor
PRE_BACKWARD) that needed yet another flag. The universal guard removes the state
guessing entirely.

Modes and wiring

src/brevitas_examples/llm/llm_quant/trainer_utils.py

  • New TrainingArguments flag fsdp_sync_all_unshards (preferred). The two
    selective flags remain for comparison/diagnostics and are documented as legacy.
    sync_all supersedes them.
  • Installed once, lazily, on the first training_step, via
    _install_fsdp_unshard_sync, which applies the guard to every FSDP-managed
    model in the step
    : the student and, when use_distillation_loss is set, the
    separately FSDP2-prepared teacher (its compute_loss forward unshards too).
  • Idempotent: re-install OR-merges flags on already-patched groups (safe if
    student/teacher share groups); no double-wrapping.
  • Observability: prints, on the main process, how many parameter groups were
    patched per model, e.g.
    FSDP unshard synchronization (universal) installed on N student parameter group(s).
  • Only barriers when an all-gather is actually pending, so no-op unshards cost
    nothing.

Memory / performance characteristics

The guard allocates no tensors and does not extend unsharded-storage lifetime;
FSDP resharding is unchanged, so FSDP's memory-sharding benefit is preserved. In
the measured 14B/32B short runs, peak allocated/reserved memory matched the
selective-fence baseline exactly. The cost is reduced compute/communication overlap
(throughput), not VRAM. (Note: strict peak-memory invariance is not guaranteed in
general, since a host barrier changes prefetch/allocator scheduling; the doc states
the measured result rather than an absolute guarantee.)

Known limitations / follow-ups

  • The universal fence relies on private FSDP2 internals (_get_fsdp_state,
    _fsdp_param_groups, _all_gather_result, _training_state, device_handle); it
    fails loudly if these are absent. Effective support is the validated PyTorch/ROCm
    stack; a version gate is a follow-up.
  • The fence is a host-side barrier: it establishes unshard completion before the
    consumer, but does not by itself prove ordering for consumers that run on custom
    side streams. The durable fix is an upstream event/record_stream change; a
    minimal reproduction is a follow-up.
  • Install timing is lazy at the first training_step (pre-training/eval forwards are
    not covered); moving it to immediately post-wrap is a candidate improvement.

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.

1 participant