Conversation
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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:
RotationBankthat owns each logical rotationexactly once and keeps it replicated (not sharded) across data-parallel ranks,
with a single packed gradient all-reduce and deterministic optimization.
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:
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.
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
RotationBankOwnership model
src/brevitas/utils/parametrization_utils.pyRotationBank(nn.Module): a single root-owned module holding each logicalrotation once in a
ParameterDict(rotation_0000,rotation_0001, ...), exposedunder
state_dictkey_brevitas_rotation_bank.rotations.rotation_XXXX._RotationBankHandle: a plain, non-registering holder that each consumerparametrization keeps instead of re-registering the shared
nn.Parameter. This iswhat prevents FSDP from seeing N aliases of the same rotation.
RotationWeightParametrizationnow resolvesrot_matthrough the bank handle(
__getattr__/bind_rotation_bank), so whole-model.to(),deepcopy, andstate-dict load all operate through the single registered bank parameter.
ensure_rotation_bank/get_rotation_bank/prune_rotation_bank: build thebank 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—FSDPRotationCoordinatorignored_modules(supporting both list and regex forms) and pins it to the accelerator device, so
FSDP never sharded the rotations.
consolidate_gradients(): atAccelerator.sync_gradients, flattens all bankgradients in a fixed insertion order, packs a per-rotation presence vector into
the same buffer, does one
all_reduce, averages by world size, and restoresgrad=Nonefor globally-unused rotations. One collective for the whole bank, noper-step owner broadcasts.
clip_grad_norm_(): a mixed-collection clip that combines FSDPDTensorgradients and replicated plain-tensor rotation gradients into one correct global
norm.
check_replica_consistency()(opt-in, diagnostic): all-gathers each bankparameter 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.pyCaileySGDuses 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_hooksnowdeletes the stale
hf_device_map(including the bank entry) so Accelerate FSDPpreparation does not choke; the bank is pinned to
LOCAL_RANKduring pre-FSDPdispatch.
src/brevitas/graph/equalize.py: rotation group IDs and bank pruning infuse_parametrizations.Part 2 — Universal FSDP2 unshard synchronization
src/brevitas_examples/llm/llm_quant/fsdp_workarounds.py—enable_fsdp_unshard_syncWhat it does
Wraps the single chokepoint every unshard flows through,
FSDPParamGroup.wait_for_unshard, and after the all-gather copy-outunconditionally 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_statevalues. That was whack-a-mole: each larger model surfaced a new unsynced unshard
site (e.g. the backward checkpoint recompute, whose state is neither
FORWARDnorPRE_BACKWARD) that needed yet another flag. The universal guard removes the stateguessing entirely.
Modes and wiring
src/brevitas_examples/llm/llm_quant/trainer_utils.pyTrainingArgumentsflagfsdp_sync_all_unshards(preferred). The twoselective flags remain for comparison/diagnostics and are documented as legacy.
sync_allsupersedes them.training_step, via_install_fsdp_unshard_sync, which applies the guard to every FSDP-managedmodel in the step: the student and, when
use_distillation_lossis set, theseparately FSDP2-prepared teacher (its
compute_lossforward unshards too).student/teacher share groups); no double-wrapping.
patched per model, e.g.
FSDP unshard synchronization (universal) installed on N student parameter group(s).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
_get_fsdp_state,_fsdp_param_groups,_all_gather_result,_training_state,device_handle); itfails loudly if these are absent. Effective support is the validated PyTorch/ROCm
stack; a version gate is a follow-up.
consumer, but does not by itself prove ordering for consumers that run on custom
side streams. The durable fix is an upstream event/
record_streamchange; aminimal reproduction is a follow-up.
training_step(pre-training/eval forwards arenot covered); moving it to immediately post-wrap is a candidate improvement.