Describe the bug
#1716 fixed a real bug (#1713): casting every floating parameter to the requested dtype during boot_transformers was corrupting FP8 scale tensors on quantized checkpoints (e.g. MXFP4's float8_e8m0fnu scales getting silently upcast to bf16, breaking the weight/scale pairing). The fix was implemented at two levels: a narrow per-parameter guard in cast_floating_params_to_dtype (skip any floating tensor with itemsize < 2, i.e. FP8/narrow floats), and a call-site guard in the new maybe_cast_floating_params that skips the cast for the entire model whenever quantization_method(model.config) returns non-None.
What's directly proven (function-level, verified by running the actual merged code, not inferred from reading the diff): maybe_cast_floating_params -- the exact function boot() calls -- skips all dtype casting for a model the moment any quantization_config is present, regardless of any individual parameter's actual dtype. The narrower per-parameter guard already in cast_floating_params_to_dtype is sufficient on its own to reproduce #1713's fix (protects the FP8 scale) while still correctly casting an ordinary floating parameter on the same model -- see the repro below.
What's inferred, not directly checked against a live checkpoint (real-world impact): quantization methods (bitsandbytes, GPTQ, AWQ, MXFP4, ...) typically only replace specific layers (usually nn.Linear); embeddings, layernorms, and often lm_head commonly stay in their original, non-quantized dtype. If that holds for a given checkpoint, boot_transformers(repo_id, dtype=torch.bfloat16, ...) would silently leave those non-quantized parameters in whatever dtype they loaded in rather than the caller's requested dtype, with no error or warning. This part is standard quantization practice and consistent with cast_floating_params_to_dtype existing in boot() since #1459 (replacing an even older raw per-parameter cast loop, i.e. it has been the relied-upon mechanism for honoring dtype=, not a redundant safety net) -- but I don't have a live quantized checkpoint on hand to confirm end-to-end against boot_transformers itself.
Code example
import torch
import torch.nn as nn
from transformer_lens.utilities.multi_gpu import cast_floating_params_to_dtype, maybe_cast_floating_params
class TinyModel(nn.Module):
def __init__(self):
super().__init__()
# A layernorm-like param that is NOT quantizer-owned -- realistic for
# any bnb/GPTQ/AWQ/MXFP4 checkpoint, which quantize Linear layers only.
self.layernorm_weight = nn.Parameter(torch.ones(8, dtype=torch.float32))
self.fp8_scale = nn.Parameter(torch.ones(4, dtype=torch.float8_e8m0fnu)) # #1713's exact dtype
class FakeConfig:
class quantization_config:
quant_method = "bitsandbytes"
model = TinyModel()
model.config = FakeConfig()
# The narrow per-parameter guard alone already protects the FP8 scale AND
# correctly casts the non-quantized param:
cast_floating_params_to_dtype(model, torch.bfloat16)
print(model.layernorm_weight.dtype) # bfloat16 (correct)
print(model.fp8_scale.dtype) # float8_e8m0fnu, untouched (correct)
# The actual boot() call path uses maybe_cast_floating_params instead, which
# skips the WHOLE model:
model2 = TinyModel()
model2.config = FakeConfig()
maybe_cast_floating_params(model2, torch.bfloat16)
print(model2.layernorm_weight.dtype) # float32 -- the dtype=bfloat16 request was silently ignored
The existing test suite currently encodes this as intended behavior: tests/unit/utilities/test_multi_gpu_unit.py::TestMaybeCastFloatingParams::test_skips_quantized_model asserts a plain fp32 nn.Linear.weight (not an FP8 scale, not a packed quantized tensor) stays uncast solely because some quantization_config is present anywhere on the model -- that's the exact behavior this report is about.
System Info
- Installed from source,
dev-4.x
- OS: macOS (arm64), CPU-only (mechanism is device-independent, reproduces on a plain
nn.Module)
Expected behaviour & fix pointers
Remove (or substantially narrow) the call-site guard in maybe_cast_floating_params; the per-parameter itemsize < 2 guard already in cast_floating_params_to_dtype is sufficient to protect quantizer-owned narrow-float scale tensors (proven above) while still correctly normalizing every other floating parameter to the caller's requested dtype, matching the behavior for unquantized models. Flagging the mechanism with confidence; flagging the real-world severity as a likely-but-unconfirmed consequence, since I couldn't test against an actual quantized checkpoint here.
Acceptance:
Checklist
Describe the bug
#1716 fixed a real bug (#1713): casting every floating parameter to the requested dtype during
boot_transformerswas corrupting FP8 scale tensors on quantized checkpoints (e.g. MXFP4'sfloat8_e8m0fnuscales getting silently upcast to bf16, breaking the weight/scale pairing). The fix was implemented at two levels: a narrow per-parameter guard incast_floating_params_to_dtype(skip any floating tensor withitemsize < 2, i.e. FP8/narrow floats), and a call-site guard in the newmaybe_cast_floating_paramsthat skips the cast for the entire model wheneverquantization_method(model.config)returns non-None.What's directly proven (function-level, verified by running the actual merged code, not inferred from reading the diff):
maybe_cast_floating_params-- the exact functionboot()calls -- skips all dtype casting for a model the moment anyquantization_configis present, regardless of any individual parameter's actual dtype. The narrower per-parameter guard already incast_floating_params_to_dtypeis sufficient on its own to reproduce #1713's fix (protects the FP8 scale) while still correctly casting an ordinary floating parameter on the same model -- see the repro below.What's inferred, not directly checked against a live checkpoint (real-world impact): quantization methods (bitsandbytes, GPTQ, AWQ, MXFP4, ...) typically only replace specific layers (usually
nn.Linear); embeddings, layernorms, and oftenlm_headcommonly stay in their original, non-quantized dtype. If that holds for a given checkpoint,boot_transformers(repo_id, dtype=torch.bfloat16, ...)would silently leave those non-quantized parameters in whatever dtype they loaded in rather than the caller's requested dtype, with no error or warning. This part is standard quantization practice and consistent withcast_floating_params_to_dtypeexisting inboot()since #1459 (replacing an even older raw per-parameter cast loop, i.e. it has been the relied-upon mechanism for honoringdtype=, not a redundant safety net) -- but I don't have a live quantized checkpoint on hand to confirm end-to-end againstboot_transformersitself.Code example
The existing test suite currently encodes this as intended behavior:
tests/unit/utilities/test_multi_gpu_unit.py::TestMaybeCastFloatingParams::test_skips_quantized_modelasserts a plain fp32nn.Linear.weight(not an FP8 scale, not a packed quantized tensor) stays uncast solely because somequantization_configis present anywhere on the model -- that's the exact behavior this report is about.System Info
dev-4.xnn.Module)Expected behaviour & fix pointers
Remove (or substantially narrow) the call-site guard in
maybe_cast_floating_params; the per-parameteritemsize < 2guard already incast_floating_params_to_dtypeis sufficient to protect quantizer-owned narrow-float scale tensors (proven above) while still correctly normalizing every other floating parameter to the caller's requested dtype, matching the behavior for unquantized models. Flagging the mechanism with confidence; flagging the real-world severity as a likely-but-unconfirmed consequence, since I couldn't test against an actual quantized checkpoint here.Acceptance:
lm_head, ...) are cast to the caller's requesteddtypeduringboot_transformers, same as an unquantized modelboot_transformerscasts quantizer-owned scale parameters and breaks MXFP4 checkpoints #1713's exact repro (FP8/narrow-float scale tensors) remains protectedtest_skips_quantized_modelis updated to reflect the corrected, narrower scope (or replaced with a test using a realistic mixed-parameter module, matchingtest_cast_skips_fp8_scales_even_if_called's pattern)make unit-testanduv run mypy .passChecklist