Files
Model-Optimizer/modelopt/torch/export/trtllm/model_config_utils.py
T
Shengliang Xu d19925e446 simple refactor(export): split TensorRT-LLM-only code into modelopt/torch/export/trtllm (#2365)
### What does this PR do?

Type of change: refactor.

**The TensorRT-LLM checkpoint export format is deprecated.** Per
`docs/source/deployment/1_tensorrt_llm.rst`: *"The
`export_tensorrt_llm_checkpoint` API will be deprecated in future
releases. Users are encouraged to transition to the unified HF export
API, which provides enhanced functionality and flexibility for exporting
models to multiple inference frameworks including TensorRT-LLM, vLLM,
and SGLang."*

That deprecated code was not sitting off to one side — it was
**interleaved with the export path we actually want to grow.**
`modelopt/torch/export` mixed the deprecated TensorRT-LLM checkpoint
logic with the framework-agnostic HF/Megatron export code, in the same
modules:

- `layer_utils.py` was 1,986 lines, of which ~1,600 were TensorRT-LLM
`build_*_config` builders. The HF path imports this module for five
small predicates (`is_moe`, `is_quantlinear`, …) and dragged the whole
deprecated builder set in with them.
- `model_config.py` held the TensorRT-LLM `ModelConfig` dataclasses
*and* the `QUANTIZATION_*` / `KV_CACHE_*` constants that every backend
needs, so all of HF export imported the deprecated checkpoint schema to
get a format name string.
- `quant_utils.py` carried two helpers whose only caller is the
deprecated `postprocess.py`.

**This PR isolates the deprecated format so it stops polluting the
HuggingFace export path.** Everything reachable only from
`export_tensorrt_llm_checkpoint` now lives under
`modelopt/torch/export/trtllm/`, and the dependency is **one-way**:
`trtllm/` reaches into the parent through `quant_format`, `quant_utils`
and `layer_utils`, and **no implementation module in the parent imports
`trtllm/`.** The single exception is the deprecation re-export in
`modelopt/torch/export/__init__.py` described below, which is scheduled
for deletion in 0.49.0.

That one-way edge is the property worth protecting in review. It means
the deprecated format can be evolved, frozen, or eventually removed
without touching HF export, and HF export can no longer accidentally
grow a dependency on it.

### Deprecation handling

The format has carried a deprecation notice in the deployment docs since
`bc546943b4` (2025-10-08, first shipped in 0.39.0) — about 11 months.
But the deprecation policy in `README.md` also specifies *how* a
deprecation is communicated: a changelog entry, a source statement of
timing, and a runtime warning on use. **None of those existed**; only
one docs page ever said anything. So 0.48.0 is the first release that
gives users a signal they can act on, and this PR treats it as the
*start* of the migration period rather than the end:

- Both entry points now emit a `DeprecationWarning` naming 0.48.0 and
the 0.49.0 removal.
- `export_tensorrt_llm_checkpoint` and
`torch_to_tensorrt_llm_checkpoint` **remain importable from
`modelopt.torch.export`** for this release only, so existing callers
keep working *and* actually receive the warning. Removing the path in
the same release that first warns would mean callers hit `ImportError`
and never see it.
- The 0.49.0 removal date is stated in all four channels the policy
names: the runtime warning, the source (`.. deprecated:: 0.48.0` plus a
comment), the changelog, and the deployment doc.

The deeper module paths (`modelopt.torch.export.model_config_export`,
`modelopt.torch.export.model_config`) are **not** forwarded. Neither
appeared in a docs example, and `model_config.py` never declared
`__all__`, so by the `__all__` convention in `CONTRIBUTING.md` they were
never part of the public surface.

Eight modules had no non-TRT-LLM importer and moved whole:
`model_config_export`, `model_config_utils`, `postprocess`,
`distribute`, `tensorrt_llm_utils`, `tensorrt_llm_type`,
`hf_config_map`, `mcore_config_map`.

Three were genuinely mixed and were split by call-graph analysis rather
than by file:

| module | stayed shared (HF path) | moved to `trtllm/` (deprecated) |
|---|---|---|
| `model_config.py` | `QUANTIZATION_*`, `KV_CACHE_*`,
`FUSION_FREE_FORMATS` → new leaf module `quant_format.py` | the
`ModelConfig` dataclasses + `LINEAR_*`/`LAYERNORM_*` checkpoint-layout
constants |
| `layer_utils.py` | 9 module-shape predicates and MoE quantizer helpers
(`is_moe`, `is_quantlinear`, `get_experts_list`,
`sync_moe_gate_up_amax`, …) | the 39 `build_*_config` builders and
enc/dec helpers |
| `quant_utils.py` | everything else | `get_scaling_factor_from_weight`,
`resmooth_and_get_scale` (only caller is `trtllm/postprocess.py`) |

`adjust_attn_amax_values` was deliberately left in the shared
`quant_utils.py`: it has no production caller at all (only a test), so
"used only by TRT-LLM export" is not demonstrable for it.

Nothing was added or removed. `export_tensorrt_llm_checkpoint` behaves
exactly as before, just from a new import path and with a warning
attached.

### Usage

```python
# Deprecated TensorRT-LLM checkpoint export — new home, and warns on call
from modelopt.torch.export.trtllm import (
    export_tensorrt_llm_checkpoint,
    torch_to_tensorrt_llm_checkpoint,
)
from modelopt.torch.export.trtllm.model_config import ModelConfig

# The pre-0.48 path still works for one release, and warns — removed in 0.49.0
from modelopt.torch.export import export_tensorrt_llm_checkpoint

# Shared format constants — new home, still re-exported from the top level
from modelopt.torch.export.quant_format import QUANTIZATION_NVFP4, KV_CACHE_FP8
from modelopt.torch.export import QUANTIZATION_NVFP4  # still works

# The recommended path — unchanged
from modelopt.torch.export import export_hf_checkpoint, get_model_type
```

### Testing

- `pre-commit` on all changed files: passes (ruff, ruff-format,
**mypy**, bandit, markdownlint). mypy caught one implicit re-export of
`is_layernorm`, now imported from the shared module directly.
- `tests/unit/torch/export`: **189 passed**. With the new `trtllm/` test
dir: **193 passed**.
- Full `tests/unit/torch`: **2367 passed, 0 export failures**. The 45
failures are pre-existing environment issues — a deepspeed circular
import and a read-only HF cache — confirmed by reading their error text,
not assumed.
- `pytest tests/gpu/torch/export --collect-only`: 172 items, no
collection error.
- In-repo consumers updated and re-verified by an AST scan that imports
every `modelopt.torch.export*` module referenced anywhere in the tree
and checks each imported name still resolves: `hf_ptq.py`,
`export_trtllm_ckpt.py`, `deepseek_v3/ptq.py`, the AutoQuantize
notebook, `hf_ptq/README.md`, 2 docs pages, 4 tests.
- **Deprecation contract is covered by committed tests** (3 new, in the
`trtllm/` test dir): the pre-0.48 top-level import still resolves to the
same objects, `torch_to_tensorrt_llm_checkpoint` warns *at call time*
rather than on first `next()` (it returns a generator, so a naive
`warnings.warn` in the body would fire late or never), and one
`export_tensorrt_llm_checkpoint` call emits exactly one warning rather
than two. The first of these makes closing the migration window early a
test failure rather than a silent regression. `pyproject.toml` sets no
`filterwarnings = error`, so no suite fails on the new warning.
- **After merging `main`** (4 commits, incl. a 180-line rewrite of
`unified_export_megatron.py` that touches a file this PR also edits):
merged with no conflicts, then re-verified rather than trusted — import
scan clean across 24 export modules, `ruff check` clean repo-wide, 193
export tests passing, GPU collection still clean.

**Not run: the GPU suites** (`tests/gpu/torch/export`,
`tests/gpu_trtllm`) — no GPU in my environment.
`tests/gpu/torch/export/test_export.py` had its imports retargeted, so
it is the one most worth a GPU run before merge.

### Reviewer note: the deprecated path has no test coverage

Worth knowing before reviewing. **No test in the repo — including
`tests/examples/` — calls `export_tensorrt_llm_checkpoint`,
`torch_to_tensorrt_llm_checkpoint`, any `build_*_config`,
`convert_to_tensorrt_llm_config`, or `postprocess_model_config`.** So
~4,600 moved lines have no direct tests, and this refactor is validated
by import-graph reasoning, lint and mypy rather than by tests exercising
the moved code.

Given the format is deprecated and scheduled for removal in 0.49.0, **no
new coverage is planned for the conversion path itself** — writing fresh
tests for an API being removed next release isn't a good use of effort.
The gap is documented so reviewers can weigh the risk, not as a TODO.
(The deprecation *mechanism* is tested; see Testing.)

One caveat on how the gap was established: a runtime check showing all
12 `trtllm` modules in `sys.modules` after the export suites is *not*
evidence of coverage — importing any submodule runs
`trtllm/__init__.py`, which star-imports `model_config_export` and pulls
in the rest. Real line coverage could not be measured (`coverage`'s
tracer is incompatible with this venv's torch build: `ValueError: module
functions cannot set METH_CLASS or METH_STATIC`, on both the C tracer
and `sysmon`). The claim rests on a call-site audit generated from the
actual public symbols of those modules.

### Before your PR is "*Ready for review*"

- Is this change backward compatible?: ✅ for the public API —
`export_tensorrt_llm_checkpoint` and `torch_to_tensorrt_llm_checkpoint`
remain importable from `modelopt.torch.export` through the 0.49.0
migration period, now with a `DeprecationWarning`. The undocumented
submodule paths `modelopt.torch.export.model_config_export` and
`.model_config` did move; see **Usage**.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A — no new
code or dependencies; existing code relocated.
- Did you write any new necessary tests?: ✅ — 3 tests covering the
deprecation contract (old import path, call-time warning, exactly-one
warning). One existing test also moved to mirror the source split. No
new coverage for the deprecated conversion path itself; see the note
above.
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ — under 0.48.0 **Deprecations**, covering both the runtime warning and
the new import location.
- Did you get Claude approval on this PR?: ❌ — not yet run.

### Additional Information

Git detected the moves, so the diff stays reviewable: 8 files show as
pure renames (100%), the three split files as rename/copy at 94–99%
similarity, and only `layer_utils.py` as a 79% rewrite — expected, since
it shed 1,616 lines to `trtllm/`.

`examples/hf_ptq/hf_ptq.py` and
`examples/llm_sparsity/weight_sparsity/export_trtllm_ckpt.py` still call
the deprecated API, so those examples now print the warning. That is the
intended nudge, but happy to silence or migrate them if preferred. They
import from the new `.trtllm` path already, so they need no change at
0.49.0.

Two incidental changes, easy to revert if unwanted:
- `modelopt/torch/export/layer_utils.py` mode `100755 → 100644` (it was
needlessly executable).
- The new test is named `test_trtllm_quant_utils.py`, not
`test_quant_utils.py`: these directories have no `__init__.py`, so
pytest derives the module name from the bare filename and the shorter
name fails collection with `import file mismatch` against the existing
`test_quant_utils.py` one level up.


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

- **New Features**
- Added shared quantization and KV-cache format definitions for export
workflows.
- Added expanded TensorRT-LLM export support, including broader model
architecture and quantization handling.
- Added distributed export utilities for coordinating checkpoint data
across processes.

- **Deprecation**
- TensorRT-LLM checkpoint export now emits a warning and is scheduled
for removal in version 0.49.0.
- Use the documented export module and save optimized model state
explicitly when needed.

- **Documentation**
- Updated guides and examples with new import paths and deprecation
guidance.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Shengliang Xu <shengliangx@nvidia.com>
2026-09-10 10:44:04 -07:00

391 lines
15 KiB
Python
Executable File

# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Common utils for the ModelConfig."""
import dataclasses
import math
from types import UnionType
from typing import Union, get_args, get_origin
import numpy as np
import torch
from ..quant_format import QUANTIZATION_FP8_PC_PT, QUANTIZATION_INT4_AWQ, QUANTIZATION_W4A8_AWQ
from ..quant_utils import to_quantized_weight
from .model_config import (
DecoderLayerConfig,
LayernormConfig,
LinearConfig,
MLPConfig,
ModelConfig,
MOEConfig,
QKVConfig,
)
# numpy doesn't know bfloat16, define abstract binary type instead
np_bfloat16 = np.dtype("V2", metadata={"dtype": "bfloat16"})
def _numpy_to_torch(x):
"""Convert numpy array to torch tensor."""
if isinstance(x, torch.Tensor):
return x
if x.dtype != np_bfloat16:
return torch.tensor(x)
return torch.tensor(x.view(np.int16)).view(torch.bfloat16)
def model_config_to_dict(model_config: ModelConfig) -> dict:
"""Converts the instance to a python dict."""
assert model_config is not None, "model_config is None"
def _to_dict(obj):
if dataclasses.is_dataclass(obj):
return {k: _to_dict(v) for k, v in vars(obj).items()}
elif isinstance(obj, dict):
return {k: _to_dict(v) for k, v in obj.items()}
elif isinstance(obj, list):
return [_to_dict(v) for v in obj]
return obj
return _to_dict(model_config)
def split_config_and_weights(
config,
weights: dict[str, torch.tensor],
prefix: str = "transformer",
layer_config_dict: dict = {},
):
"""Util function to split the weights or any torch.Tensor in nested config to weights.
A weight id starts with transformers or lm_head will also be generated to link the original key to the weights dict.
The weights in the weights dict are contiguous.
layer_config_dict: A dictionary containing layerwise quantization format information and awq_block_size information
when relevant. It is used to export quantization.json for auto_quant checkpoint.
"""
if isinstance(config, dict):
for k, v in config.items():
if k == "lm_head" and "medusa_heads" not in prefix:
# lm_head is not part of the transformer.
array_key = k
elif k == "experts":
# Omit the 'experts' key that is not in the model name
array_key = prefix
elif k == "medusa_heads":
# medusa_heads is not part of the transformer
array_key = k
elif "rel_attn_table" in prefix:
# rel_attn_table is treated as a weight for quantize loop whereas in TRTLLM it is a Tensor
array_key = prefix
else:
array_key = f"{prefix}.{k}"
# TensorRT-LLM 0.17 uses the kv_cache_scaling_factor,
# the following code will be deprecated after the 0.18 upgrade.
if "v_cache_scaling_factor" in array_key:
k_cache_scaling_factor_key = array_key.replace(
"v_cache_scaling_factor", "k_cache_scaling_factor"
)
if k_cache_scaling_factor_key in weights:
weights[
array_key.replace("v_cache_scaling_factor", "kv_cache_scaling_factor")
] = torch.maximum(v, weights[k_cache_scaling_factor_key])
# Construct per_layer quantization dictionary, with block size information
if array_key != "transformer.quantization" and (
"quantization" in array_key or "awq_block_size" in array_key
):
layer_config_dict[array_key] = v
if isinstance(v, torch.Tensor):
weights[array_key] = v
config[k] = f"{array_key}"
else:
split_config_and_weights(v, weights, array_key, layer_config_dict)
elif isinstance(config, list):
for i, v in enumerate(config):
array_key = f"{prefix}.{i}"
if isinstance(v, torch.Tensor):
weights[array_key] = v
config[i] = f"{array_key}"
else:
split_config_and_weights(v, weights, array_key, layer_config_dict)
def _unified_weights_key(k: str) -> str:
"""Try to unify the weights dict key between old npz and the new safetensors format."""
prefixes = ["transformer.", "_np:"]
for prefix in prefixes:
k = k.removeprefix(prefix)
k = k.replace("final_layernorm", "ln_f")
return k.replace(":", ".")
def _restore_model_config(model_config, weights: dict[str, np.ndarray | torch.Tensor]):
def _is_tensor_key(k):
return isinstance(k, str) and _unified_weights_key(k) in weights
if isinstance(model_config, dict):
for k, v in model_config.items():
if _is_tensor_key(v):
model_config[k] = _numpy_to_torch(weights[_unified_weights_key(v)])
else:
_restore_model_config(v, weights)
if isinstance(model_config, list):
for i, v in enumerate(model_config):
if _is_tensor_key(v):
model_config[i] = _numpy_to_torch(weights[_unified_weights_key(v)])
else:
_restore_model_config(v, weights)
def restore_model_config(model_config, weights: dict[str, np.ndarray | torch.Tensor]):
"""Recursively restores the model_config from json and loads np.ndarray or torch.Tensor weights from weights."""
unified_key_weights = {}
for k, v in weights.items():
unified_key_weights[_unified_weights_key(k)] = v
_restore_model_config(model_config, unified_key_weights)
def _from_dict(class_type, data):
"""Helper function to load the data as a class_type. class_type must be a dataclass."""
if data is None:
return None
# class_type of quantization is str | None, which is catergrorized as Union
if class_type != str | None and get_origin(class_type) in [Union, UnionType]:
# Handle QKV
if all(key in data for key in ["q", "k", "v"]):
# splitted qkv case
class_type = QKVConfig
elif all(key in data for key in ["router", "experts"]):
# moe
class_type = MOEConfig
elif all(key in data for key in ["fc", "gate", "proj"]):
# mlp
class_type = MLPConfig
else:
# merged qkv case
assert "linear_type" in data, f"{data} is not a valid LinearConfig"
class_type = LinearConfig
if dataclasses.is_dataclass(class_type):
fieldtypes = {f.name: f.type for f in dataclasses.fields(class_type)}
fields_map = {}
for k, v in data.items():
if k in fieldtypes:
# We only handle keys available in the fields.
# Deprecated fields in the checkpoint will be ignored.
fields_map[k] = _from_dict(fieldtypes[k], v)
return class_type(**fields_map)
elif get_origin(class_type) is list and dataclasses.is_dataclass(get_args(class_type)[0]):
list_value = []
for child in data:
child_class_type = get_args(class_type)[0]
list_value.append(_from_dict(child_class_type, child))
return list_value
else:
return data
def model_config_from_dict(d: dict) -> ModelConfig:
"""Load a dict to a `ModelConfig` instance."""
config_type = ModelConfig
config_type_map = {}
for t in [ModelConfig, DecoderLayerConfig, LayernormConfig, LinearConfig]:
config_type_map[t.__name__] = t
if "__name__" in d:
config_name = d.pop("__name__")
try:
config_type = config_type_map[config_name]
except Exception as e:
raise NotImplementedError(f"{config_name} not supported") from e
return _from_dict(config_type, d)
def pad_weights(weights, tp_size):
"""Returns the padded weights to tp_size."""
assert len(weights.shape) > 1
def _pad_size(original_size, tp_size):
return int(math.ceil(original_size / tp_size) * tp_size)
original_size = weights.shape[0]
padded_size = _pad_size(original_size, tp_size)
if original_size != padded_size:
pad_width = padded_size - original_size
return torch.nn.functional.pad(weights, (0, 0, 0, pad_width), "constant", value=0)
return weights
def merge_qkv(model_config):
"""Merges the qkv fields in model_config from QKVConfig to a single LinearConfig."""
for decoder_config in model_config.layers:
for attention_key in ["attention", "self_attention", "cross_attention"]:
attention = getattr(decoder_config, attention_key, None)
if attention and isinstance(attention.qkv, QKVConfig):
splitted_qkv = attention.qkv
attention.qkv = LinearConfig()
attention.qkv.weight = splitted_qkv.weight
attention.qkv.bias = splitted_qkv.bias
attention.qkv.activation_scaling_factor = splitted_qkv.activation_scaling_factor
attention.qkv.weights_scaling_factor = splitted_qkv.weights_scaling_factor
attention.qkv.weights_scaling_factor_2 = splitted_qkv.weights_scaling_factor_2
attention.qkv.prequant_scaling_factor = splitted_qkv.prequant_scaling_factor
attention.qkv.awq_block_size = splitted_qkv.awq_block_size
# Assert q,k,v have same quantization formats before merging
assert (
splitted_qkv.q.quantization
== splitted_qkv.k.quantization
== splitted_qkv.v.quantization
), "Quantization formats of q,k,v must be the same."
attention.qkv.quantization = splitted_qkv.q.quantization
# Collect GPU memory from the deleted tensors
del splitted_qkv
def merge_gate_fc(model_config):
"""Postprocess the MLP config for TensorRT-LLM export."""
for decoder_config in model_config.layers:
mlp = None
if isinstance(decoder_config.mlp, MLPConfig):
mlp = decoder_config.mlp
elif (
isinstance(decoder_config.mlp, MOEConfig)
and decoder_config.mlp.shared_expert is not None
):
mlp = decoder_config.mlp.shared_expert
if mlp is not None and mlp.merge_gate_fc and mlp.gate is not None and mlp.fc is not None:
mlp.fc.weight = torch.cat(
[
mlp.gate.weight,
mlp.fc.weight,
],
dim=0,
)
if (
mlp.fc.weights_scaling_factor is not None
and mlp.fc.weights_scaling_factor.numel() > 1
):
mlp.fc.weights_scaling_factor = torch.cat(
[mlp.gate.weights_scaling_factor, mlp.fc.weights_scaling_factor], dim=0
)
mlp.gate = None
def pack_linear_weights(model_config: ModelConfig):
"""Packs the quantized linear weights in the model_config to the quantized format."""
def _linear_layer_to_quantized_weight(linear_layers):
for linear_layer in linear_layers:
# Check if quantization of the layer is None to support auto_quant
if isinstance(linear_layer, LinearConfig) and (
linear_layer.weights_scaling_factor is not None
and linear_layer.quantization is not None
):
# Quantize on CPU if we are short of GPU memory.
# Using 2x of the tensor size as a threshold.
if linear_layer.weight.is_cuda:
free_mem, _ = torch.cuda.mem_get_info(linear_layer.weight.device)
if (
free_mem
< 2 * linear_layer.weight.element_size() * linear_layer.weight.nelement()
):
linear_layer.weight = linear_layer.weight.cpu()
# Save the quantize layer weights to cpu and save gpu memory.
if linear_layer.weight.element_size() > 1:
linear_layer.weight = to_quantized_weight(
linear_layer.weight,
linear_layer.weights_scaling_factor,
linear_layer.quantization,
linear_layer.weights_scaling_factor_2,
linear_layer.awq_block_size,
).cpu()
# Convert to int8 if to make the checkpoint compatible with the latest TensorRT-LLM release.
# The future TensorRT-LLM release will use uint8 weights instead.
if linear_layer.quantization in [QUANTIZATION_INT4_AWQ, QUANTIZATION_W4A8_AWQ]:
linear_layer.weight = linear_layer.weight.view(torch.int8)
# TensorRT-LLM uses per_channel_scale for FP8 per channel weight quantization.
if linear_layer.quantization == QUANTIZATION_FP8_PC_PT:
linear_layer.per_channel_scale = linear_layer.weights_scaling_factor.cpu()
linear_layer.weights_scaling_factor = None
else:
linear_layer.weights_scaling_factor = linear_layer.weights_scaling_factor.cpu()
if not model_config.quantization:
return
def _find_linear_configs_recursive(model_config):
linear_configs = []
# Base case - not a dataclass
if not dataclasses.is_dataclass(model_config):
return linear_configs
# Check if current object is a LinearConfig
if isinstance(model_config, LinearConfig):
linear_configs.append(model_config)
return linear_configs
# Recursively check all fields
for field in dataclasses.fields(model_config):
value = getattr(model_config, field.name)
if isinstance(value, list):
for item in value:
linear_configs.extend(_find_linear_configs_recursive(item))
elif isinstance(value, dict):
for item in value.values():
linear_configs.extend(_find_linear_configs_recursive(item))
# Handle nested dataclasses
elif dataclasses.is_dataclass(value):
linear_configs.extend(_find_linear_configs_recursive(value))
return linear_configs
linear_layers = _find_linear_configs_recursive(model_config)
_linear_layer_to_quantized_weight(linear_layers)
if model_config.medusa_heads is not None:
linear_layers = []
for head in model_config.medusa_heads:
linear_layers.append(head.lm_head)
for layer in head.medusa_layers:
linear_layers.append(layer.linear)
_linear_layer_to_quantized_weight(linear_layers)