Skip to content

Fix (brevitas_examples/llm): more robust BF16 rotation training - #1583

Open
Giuseppe5 wants to merge 4 commits into
Xilinx:masterfrom
Giuseppe5:feat-rotation-training-stability
Open

Giuseppe5 wants to merge 4 commits into
Xilinx:masterfrom
Giuseppe5:feat-rotation-training-stability

Conversation

@Giuseppe5

Copy link
Copy Markdown
Collaborator

Summary

Improve the robustness of BF16 rotation fine-tuning by supporting FP32 rotation master parameters and correctly honoring per-group CaileySGD compute dtypes.

This PR fixes a configuration gap where optimizer_dtype was supplied through optimizer_scheduler_args as a parameter-group value, but CaileySGD only read a constructor-level dtype. As a result, FP32 optimizer configuration could be silently ignored.

Motivation

BF16 model execution should not require BF16 optimizer state for trainable rotation matrices. Keeping the rotation master and manifold update in FP32 preserves the Stiefel constraint while the forward path continues to use BF16 tensors.

Changes

  • Make CaileySGD consume and validate its per-group dtype setting.
    • Supports None, dtype names such as "float32", and torch.dtype values.
    • Rejects unknown and non-floating-point dtypes.
    • Uses the configured dtype for the Stiefel master/shadow parameter, gradient conversion, and momentum state.
    • Rejects mixed-dtype non-Stiefel usage rather than silently ignoring the requested master dtype.
  • Add generic FP32 rotation master storage.
    • Introduce rotation_parameter_dtype in RotationTrainingArguments.
    • Preserve Parameter identity and sharing while changing storage dtype.
    • Keep the model/rotation forward in the model dtype; only the trainable master storage is promoted.
  • Add a pre-compilation fine-tuning preparation lifecycle.
    • Resolve trainer class and training args after quantization/post-processing.
    • Invoke optional prepare_model_for_training() before calibration forward and compile_quant().
    • Pass the resolved trainer and args to apply_fine_tuning() to avoid duplicate parsing/preparation.
  • Consolidate duplicate dtype resolution logic.
    • Add resolve_torch_dtype() in brevitas.utils.torch_utils.
    • Reuse it from CaileySGD and generic rotation preparation.
  • Add CaileySGD coverage for:
    • FP32 per-group master/shadow behavior with low-precision parameters.
    • Invalid dtype rejection.
    • Rejection of unsupported non-Stiefel master-dtype configuration.

Notable difference with optimizer_dtype

The main difference happens when using gradient_accumulation_steps>1

With the shadow approach, gradients accumulate in the BF16 parameter’s .grad buffer before CaileySGD converts them to FP32. With FP32 parameter storage, gradients accumulate in FP32. Therefore the two configurations are not mathematically identical whenever there are multiple microbatches, gradient clipping, hooks, or other code operating on p.grad.

@Giuseppe5 Giuseppe5 changed the title Fix: more robust BF16 spinquant training Fix(brevitas_examples/llm): more robust BF16 spinquant training Aug 14, 2026
@Giuseppe5 Giuseppe5 changed the title Fix(brevitas_examples/llm): more robust BF16 spinquant training Fix(brevitas_examples/llm): more robust BF16 rotation training Aug 14, 2026
@Giuseppe5 Giuseppe5 changed the title Fix(brevitas_examples/llm): more robust BF16 rotation training Fix (brevitas_examples/llm): more robust BF16 rotation training Aug 14, 2026
@Giuseppe5

Copy link
Copy Markdown
Collaborator Author

Simplify this, we don't need the resolve dtype function

@Giuseppe5 Giuseppe5 self-assigned this Aug 26, 2026
@Giuseppe5
Giuseppe5 force-pushed the feat-rotation-training-stability branch from 44fe022 to 982c0ed Compare September 2, 2026 09:54

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