Skip to content

[Bug Report] maybe_cast_floating_params skips dtype normalization for an entire quantized model, not just quantizer-owned tensors #1743

Description

@LightWork666

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:

  • A quantized model's non-quantizer-owned floating parameters (layernorms, embeddings, unquantized lm_head, ...) are cast to the caller's requested dtype during boot_transformers, same as an unquantized model
  • [Bug Report] boot_transformers casts quantizer-owned scale parameters and breaks MXFP4 checkpoints #1713's exact repro (FP8/narrow-float scale tensors) remains protected
  • test_skips_quantized_model is updated to reflect the corrected, narrower scope (or replaced with a test using a realistic mixed-parameter module, matching test_cast_skips_fp8_scales_even_if_called's pattern)
  • make unit-test and uv run mypy . pass

Checklist

  • I have checked that there is no similar issue in the repo (required)

Activity

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

Metadata

Metadata

Assignees

Labels

TransformerBridgeBug specific to the new TransformerBridge systembugSomething isn't workingcomplexity-moderateModerately complicated issues for people who have intermediate experience with the code

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions