From 20a55adec51253537442171ca05da1ca480b6421 Mon Sep 17 00:00:00 2001 From: lvliang-intel Date: Wed, 17 Jun 2026 15:25:47 +0800 Subject: [PATCH 1/7] Support DiffusionGemma model Signed-off-by: lvliang-intel --- auto_round/calibration/llm.py | 3 +- auto_round/inference/convert_model.py | 19 +++++- auto_round/special_model_handler.py | 7 +- auto_round/utils/common.py | 33 ++++++++++ auto_round/utils/model.py | 95 +++++++++++++++++++++++++++ 5 files changed, 153 insertions(+), 4 deletions(-) diff --git a/auto_round/calibration/llm.py b/auto_round/calibration/llm.py index 522399f47..ae0af5e87 100644 --- a/auto_round/calibration/llm.py +++ b/auto_round/calibration/llm.py @@ -38,6 +38,7 @@ hook_ngram_embeddings_on_cpu, is_quantized_input_module, mv_module_from_gpu, + safe_tie_weights, to_device, to_dtype, ) @@ -135,7 +136,7 @@ def collect(self, block_names, nsamples, layer_names=None, last_cache_name=None) no_split_module_classes=no_split_modules, ) if hasattr(c.model_context.model, "tie_weights"): - c.model_context.model.tie_weights() + safe_tie_weights(c.model_context.model) device_map = infer_auto_device_map( c.model_context.model, max_memory=new_max_memory, diff --git a/auto_round/inference/convert_model.py b/auto_round/inference/convert_model.py index 7d19b1e3e..834076577 100644 --- a/auto_round/inference/convert_model.py +++ b/auto_round/inference/convert_model.py @@ -47,6 +47,7 @@ is_transformers_version_greater_or_equal_5, set_module, ) +from auto_round.utils.model import prune_stale_tied_weights_keys supported_devices = ("cpu", "hpu", "xpu", "cuda", "mps") @@ -62,14 +63,24 @@ def flatten_list(nested_list): def skip_not_convert_modules(model, quantization_config, layer_names, layer_configs): + user_specified = bool(getattr(quantization_config, "modules_to_not_convert", None)) modules_to_not_convert = getattr(quantization_config, "modules_to_not_convert", []) try: # transformers new api modules_to_not_convert = get_modules_to_not_convert(model, modules_to_not_convert, add_default_skips=True) except: modules_to_not_convert = _get_modules_to_not_convert(model, modules_to_not_convert) + + if modules_to_not_convert and not user_specified: + _DEFAULT_SKIP_KEYWORDS = ("embed", "embed_tokens", "lm_head", "output_embed", "norm") + modules_to_not_convert = [ + name for name in modules_to_not_convert if any(key in name for key in _DEFAULT_SKIP_KEYWORDS) + ] + if modules_to_not_convert: + # Pre-compile patterns once instead of recompiling them for every layer name. + compiled_patterns = [re.compile(n) for n in modules_to_not_convert] for layer_name in layer_names: - if any([re.search(re.compile(n), layer_name) for n in modules_to_not_convert]): + if any(pattern.search(layer_name) for pattern in compiled_patterns): layer_configs[layer_name] = {"bits": 16} return layer_configs @@ -367,8 +378,10 @@ def get_layer_config(model, quantization_config): modules_in_block_to_quantize = flatten_list( quantization_config.modules_in_block_to_quantize ) # Flatten the list + # Pre-compile patterns once instead of recompiling them for every layer name. + compiled_modules_in_block = [re.compile(n) for n in modules_in_block_to_quantize] for layer_name in layer_names: - if not any([re.search(re.compile(n), layer_name) is not None for n in modules_in_block_to_quantize]): + if not any(pattern.search(layer_name) is not None for pattern in compiled_modules_in_block): extra_config[layer_name] = {"bits": 16} # Default to 16-bit for unquantized layers # Expand GPTQ 'dynamic' config (regex-based) @@ -873,6 +886,8 @@ def convert_hf_model(model: nn.Module, target_device: str = "cpu") -> tuple[nn.M layer_configs = get_layer_config(model, quantization_config) used_backends = _replace_by_quant_layers(model, layer_configs, backend, target_device, packing_format) + prune_stale_tied_weights_keys(model) + rotation_config = getattr(quantization_config, "rotation_config", None) if rotation_config is not None and rotation_config: from auto_round.algorithms.transforms.rotation.apply import apply_rotation_transform diff --git a/auto_round/special_model_handler.py b/auto_round/special_model_handler.py index 54b7a69a4..d54d5e5e7 100644 --- a/auto_round/special_model_handler.py +++ b/auto_round/special_model_handler.py @@ -20,6 +20,7 @@ from auto_round.formats import OutputFormat from auto_round.modeling.fused_moe.replace_modules import apply_replacements, release_original_module_ from auto_round.utils import is_moe_model_via_config, logger +from auto_round.utils.model import prune_stale_tied_weights_keys mllms_with_limited_bs = ( "llava", @@ -401,6 +402,7 @@ def update_module( if cleanup_original: release_original_module_(model) + prune_stale_tied_weights_keys(model) return model @@ -1270,7 +1272,10 @@ def _nextstep_pipeline_fn(pipe, prompts, guidance_scale=7.5, num_inference_steps return pipe, model -_PRE_DEFINED_FIXED_ATTR = {"gemma4_unified": {"has_variable_block_shape": True}} +_PRE_DEFINED_FIXED_ATTR = { + "gemma4_unified": {"has_variable_block_shape": True}, + "diffusion_gemma": {"has_variable_block_shape": True}, +} def get_predefined_fixed_attr(model: torch.nn.Module) -> dict | None: diff --git a/auto_round/utils/common.py b/auto_round/utils/common.py index 5e7a50e6a..812013d78 100644 --- a/auto_round/utils/common.py +++ b/auto_round/utils/common.py @@ -23,6 +23,7 @@ import torch import transformers +from transformers.models.diffusion_gemma.modeling_diffusion_gemma import DiffusionGemmaModel from packaging import version from auto_round.export.export_to_gguf.config import GGUF_CONFIG @@ -238,6 +239,36 @@ def _tensor_get_dtype(self): torch.Tensor.get_dtype = _tensor_get_dtype +def _patch_diffusion_gemma_tied_weights(): + """Install a ``__init__`` hook on ``DiffusionGemmaModel`` to prune stale tied-weights keys. + + AutoRound unfuses ``DiffusionGemmaTextExperts.gate_up_proj`` / + ``down_proj`` (fused 3D ``nn.Parameter``) into per-expert + ``gate_proj / up_proj / down_proj`` ``nn.Linear`` modules at quantize + time. The encoder/decoder weight-tying map declared in + ``DiffusionGemmaModel._tied_weights_keys`` still references the original + fused parameter names (``gate_up_proj``, ``down_proj``). + The fix is to prune the stale patterns at the moment the model is constructed. + """ + if getattr(DiffusionGemmaModel, "_ar_tied_prune_patched", False): + return + original_init = DiffusionGemmaModel.__init__ + + from auto_round.utils.model import prune_stale_tied_weights_keys + + def _patched_init(self, *args, **kwargs): + original_init(self, *args, **kwargs) + try: + prune_stale_tied_weights_keys(self) + except Exception as exc: # noqa: BLE001 + logger.warning( + f"[DiffusionGemma] prune_stale_tied_weights_keys during __init__ failed: {exc}" + ) + + DiffusionGemmaModel.__init__ = _patched_init + DiffusionGemmaModel._ar_tied_prune_patched = True + + def _patch_default_rope_init(): """Restore legacy ``rope_type='default'`` support for older remote-code models. @@ -351,6 +382,8 @@ def monkey_patch_transformers(): # transformers 5.3.0 calls tensor.get_dtype() on plain torch.Tensor objects # while loading pre-quantized checkpoints. _patch_tensor_get_dtype_for_prequantized_loading() + if parsed_version >= version.parse("5.11.0"): + _patch_diffusion_gemma_tied_weights() _patch_default_rope_init() _patch_rotary_embedding_init_for_legacy_remote_code() if parsed_version >= version.parse("4.56.0"): diff --git a/auto_round/utils/model.py b/auto_round/utils/model.py index 94ff1b199..bede35eec 100644 --- a/auto_round/utils/model.py +++ b/auto_round/utils/model.py @@ -23,6 +23,7 @@ import psutil import torch import transformers +from transformers import PreTrainedModel from packaging import version from auto_round import envs @@ -76,6 +77,100 @@ def resolve_model_type(model): from auto_round.schemes import QuantizationScheme +def prune_stale_tied_weights_keys(model: torch.nn.Module) -> int: + """Drop ``_tied_weights_keys`` regex patterns that no longer match any parameter. + + AutoRound unfuses MoE experts before calibration, splitting fused projections such as + ``gate_up_proj`` ``[num_experts, 2*inter, hidden]`` into per-expert ``gate_proj`` / + ``up_proj`` linear weights. The model's declared ``_tied_weights_keys`` still reference + the original fused name (e.g. DiffusionGemma ties ``encoder...gate_up_proj`` to + ``decoder...gate_up_proj``), which now matches zero parameters. + + Removing only the now-empty patterns lets the remaining ties be established normally. + The split ``gate_proj`` / ``up_proj`` weights stay tied through the generic ``*.weight`` pattern + that already covers every ``.weight`` parameter under the layers. Already-expanded + plain parameter names are kept untouched. + + Args: + model (torch.nn.Module): model whose stale tie patterns should be removed. + + Returns: + int: number of stale tie entries removed across all submodels. + """ + common_case = re.compile(r"^[A-Za-z0-9_\.]+(weight)|(bias)$") + removed = 0 + + for _, submodule in model.named_modules(remove_duplicate=False): + if not isinstance(submodule, PreTrainedModel): + continue + tied = getattr(submodule, "_tied_weights_keys", None) + if not isinstance(tied, dict) or not tied: + continue + if not getattr(submodule.config, "tie_word_embeddings", False): + continue + if all(common_case.match(target) and common_case.match(source) for target, source in tied.items()): + continue + names = {k for k, _ in submodule.named_parameters(remove_duplicate=False)} | { + k for k, _ in submodule.named_buffers(remove_duplicate=False) + } + pruned = {} + for target, source in tied.items(): + if common_case.match(target) and common_case.match(source): + pruned[target] = source + continue + source_params = [n for n in names if re.search("^" + source, n)] + target_params = [n for n in names if re.search("^" + target, n)] + if len(source_params) > 0 and len(target_params) > 0 and len(target_params) % len(source_params) == 0: + pruned[target] = source + else: + removed += 1 + logger.trace( + f"Removing stale tie pattern '{target}' -> '{source}' from " + f"{type(submodule).__name__}: it no longer matches any parameter." + ) + if len(pruned) != len(tied): + submodule._tied_weights_keys = pruned + + for _, submodule in model.named_modules(remove_duplicate=False): + if not isinstance(submodule, PreTrainedModel): + continue + cached = getattr(submodule, "all_tied_weights_keys", None) + if not isinstance(cached, dict) or not cached: + continue + existing = {n for n, _ in submodule.named_parameters(remove_duplicate=False)} | { + n for n, _ in submodule.named_buffers(remove_duplicate=False) + } + cached_pruned = { + target: source for target, source in cached.items() if target in existing and source in existing + } + if len(cached_pruned) != len(cached): + removed += len(cached) - len(cached_pruned) + logger.trace( + f"Pruned {len(cached) - len(cached_pruned)} stale entries from " + f"{type(submodule).__name__}.all_tied_weights_keys cache " + f"(targets/sources no longer resolve after unfuse/quantization)." + ) + submodule.all_tied_weights_keys = cached_pruned + + return removed + + +def safe_tie_weights(model: torch.nn.Module) -> None: + """Call ``model.tie_weights()`` defensively. + + Args: + model (torch.nn.Module): model whose weights should be tied if supported. + """ + tie_fn = getattr(model, "tie_weights", None) + if not callable(tie_fn): + return + prune_stale_tied_weights_keys(model) + try: + tie_fn() + except ValueError as e: + logger.warning(f"model.tie_weights() raised ValueError, skipping weight tying: {e}") + + def clean_module_parameter(submodule: torch.nn.Module, param_name: str) -> None: """This function is recommended to be used instead of module.weight = None. For models like `tie_word_embeddings`, setting the embedding weight to None From 0991c1842ef4c22a1e4e2daaa9ea07b28630b4ed Mon Sep 17 00:00:00 2001 From: lvliang-intel Date: Wed, 17 Jun 2026 21:59:48 +0800 Subject: [PATCH 2/7] add ut Signed-off-by: lvliang-intel --- .../utils/test_diffusiongemma_support.py | 610 ++++++++++++++++++ 1 file changed, 610 insertions(+) create mode 100644 test/test_cpu/utils/test_diffusiongemma_support.py diff --git a/test/test_cpu/utils/test_diffusiongemma_support.py b/test/test_cpu/utils/test_diffusiongemma_support.py new file mode 100644 index 000000000..ce6387f3a --- /dev/null +++ b/test/test_cpu/utils/test_diffusiongemma_support.py @@ -0,0 +1,610 @@ +# Copyright (c) 2026 Intel Corporation +# +# 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. +"""Unit tests for DiffusionGemma model changes.""" + +import types +from typing import Iterable +from unittest.mock import patch + +import pytest +import torch +import torch.nn as nn + +from transformers import PretrainedConfig, PreTrainedModel + + +# --------------------------------------------------------------------------- +# Test helpers +# --------------------------------------------------------------------------- + + +def _make_pretrained_submodule( + name: str, + params: Iterable[str], + tied: dict | None = None, + all_tied: dict | None = None, + tie_word_embeddings: bool = True, + parent: nn.Module | None = None, +): + """Build a minimal ``PreTrainedModel`` submodule with controllable tie state. + + Args: + name: attribute name on ``parent`` (or local var if ``parent`` is None). + params: iterable of leaf names — each becomes an ``nn.Linear`` of + shape ``(out=2, in=2)`` attached directly to the fake model. + ``named_parameters()`` therefore surfaces them as just the leaf + name (e.g. ``"weight"``) which matches ``common_case`` regexes. + tied: value to assign to ``_tied_weights_keys``. + all_tied: value to assign to ``all_tied_weights_keys``. + tie_word_embeddings: value for ``config.tie_word_embeddings``. + parent: if provided, register the module as ``parent.name``. + """ + + class _FakePT(PreTrainedModel): + config_class = PretrainedConfig + + def __init__(self, cfg): + super().__init__(cfg) + self.dummy = nn.Parameter(torch.zeros(1)) + + def _init_weights(self, module): + pass + + cfg = PretrainedConfig(tie_word_embeddings=tie_word_embeddings) + module = _FakePT(cfg) + for pname in params: + # Use bare ``nn.Parameter`` so ``named_parameters`` surfaces the + # leaf name directly (matches what the production model looks like + # after MoE unfuse: a ``gate_proj`` / ``up_proj`` weight). + # ``nn.Module.__setattr__`` registers ``nn.Parameter`` as a parameter + # when the name contains no "."; nested names would need explicit + # ``register_parameter`` on the sub-module. + assert "." not in pname, "test fixture only supports leaf parameter names" + setattr(module, pname, nn.Parameter(torch.zeros(2, 2))) + + if tied is not None: + module._tied_weights_keys = tied + if all_tied is not None: + module.all_tied_weights_keys = all_tied + + if parent is not None: + setattr(parent, name, module) + return parent + return module + + +# --------------------------------------------------------------------------- +# 1. Tests for prune_stale_tied_weights_keys +# --------------------------------------------------------------------------- + + +class TestPruneStaleTiedWeightsKeys: + """Unit tests for ``prune_stale_tied_weights_keys``.""" + + def test_drops_pattern_with_no_matching_target_or_source(self): + """Tie pattern referencing a name that no longer exists is dropped.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + # Only ``gate`` exists; ``gate_up_proj`` was unfused away. + _make_pretrained_submodule( + "text", + params=["gate"], + tied={"encoder.gate_up_proj": "decoder.gate_up_proj"}, + parent=wrapper, + ) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 1 + assert wrapper.text._tied_weights_keys == {} + + def test_keeps_pattern_that_resolves_to_real_params(self): + """Tie patterns whose target/source both resolve to existing parameters are kept.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + _make_pretrained_submodule( + "text", + params=["src", "dst"], + tied={"dst": "src"}, + parent=wrapper, + ) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 0 + assert wrapper.text._tied_weights_keys == {"dst": "src"} + + def test_keeps_common_case_patterns_even_if_unused(self): + """``*.weight`` / ``*.bias`` style patterns are kept untouched even if no params match.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + _make_pretrained_submodule( + "text", + params=[], + tied={"model.embed_tokens.weight": "lm_head.weight"}, + parent=wrapper, + ) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 0 + assert wrapper.text._tied_weights_keys == { + "model.embed_tokens.weight": "lm_head.weight" + } + + def test_keeps_pattern_when_target_source_count_is_multiple_of_source(self): + """Pattern is kept when target param count is a non-zero multiple of source param count.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + # 1 source, 2 targets -> 2 % 1 == 0 -> keep. + _make_pretrained_submodule( + "text", + params=["src", "t0", "t1"], + tied={r"^t[0-9]$": r"^src$"}, + parent=wrapper, + ) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 0 + assert r"^t[0-9]$" in wrapper.text._tied_weights_keys + + def test_drops_pattern_when_target_count_not_multiple_of_source(self): + """Pattern is dropped when target param count is not a multiple of source.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + # 2 sources, 3 targets -> 3 % 2 != 0 -> drop. + _make_pretrained_submodule( + "text", + params=["src0", "src1", "t0", "t1", "t2"], + tied={r"^t[0-9]$": r"^src[0-9]$"}, + parent=wrapper, + ) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 1 + assert r"^t[0-9]$" not in wrapper.text._tied_weights_keys + + def test_skips_submodule_without_tied_weights_keys(self): + """Submodules without ``_tied_weights_keys`` are skipped.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + _make_pretrained_submodule("text", params=["w"], tied=None, parent=wrapper) + # PreTrainedModel base class does not define ``_tied_weights_keys`` by default. + assert getattr(wrapper.text, "_tied_weights_keys", None) is None + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 0 + + def test_skips_submodule_when_tie_word_embeddings_false(self): + """Submodules with ``tie_word_embeddings=False`` are skipped entirely.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + _make_pretrained_submodule( + "text", + params=[], + tied={"a": "b"}, + tie_word_embeddings=False, + parent=wrapper, + ) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 0 + # Untouched. + assert wrapper.text._tied_weights_keys == {"a": "b"} + + def test_skips_non_pretrained_submodules(self): + """Plain ``nn.Module`` children are never inspected.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + wrapper.plain = nn.Linear(2, 2) + wrapper.plain._tied_weights_keys = {"a": "b"} # would never appear on nn.Module + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 0 + assert wrapper.plain._tied_weights_keys == {"a": "b"} + + def test_prunes_all_tied_weights_keys_cache(self): + """Stale entries in ``all_tied_weights_keys`` are pruned (target/source both must exist).""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + _make_pretrained_submodule( + "text", + params=["real"], + tied=None, + all_tied={ + "real": "real", + "stale": "missing", + }, + parent=wrapper, + ) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 1 + assert wrapper.text.all_tied_weights_keys == {"real": "real"} + + def test_prunes_stale_entry_even_if_only_target_missing(self): + """Cached entry whose target no longer exists is dropped.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + _make_pretrained_submodule( + "text", + params=["keep"], + tied=None, + all_tied={ + "missing": "keep", + "keep": "keep", + }, + parent=wrapper, + ) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 1 + assert wrapper.text.all_tied_weights_keys == {"keep": "keep"} + + def test_returns_zero_when_nothing_to_prune(self): + """No stale entries -> returns 0.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + _make_pretrained_submodule("text", params=[], tied=None, parent=wrapper) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 0 + + def test_handles_multiple_pretrained_submodules(self): + """Multiple ``PreTrainedModel`` submodules are pruned independently.""" + from auto_round.utils.model import prune_stale_tied_weights_keys + + wrapper = nn.Module() + _make_pretrained_submodule( + "a", + params=[], + tied={"stale_a": "stale_a"}, + parent=wrapper, + ) + _make_pretrained_submodule( + "b", + params=["real"], + tied={"real": "real"}, + parent=wrapper, + ) + + removed = prune_stale_tied_weights_keys(wrapper) + + assert removed == 1 + assert wrapper.a._tied_weights_keys == {} + assert wrapper.b._tied_weights_keys == {"real": "real"} + + +# --------------------------------------------------------------------------- +# 2. Tests for safe_tie_weights +# --------------------------------------------------------------------------- + + +class TestSafeTieWeights: + """Unit tests for ``safe_tie_weights``.""" + + def test_calls_tie_weights_when_supported(self): + """Calls ``model.tie_weights()`` when present.""" + from auto_round.utils.model import safe_tie_weights + + called = [] + + class _M(nn.Module): + def tie_weights(self): + called.append(True) + + m = _M() + safe_tie_weights(m) + assert called == [True] + + def test_noop_when_no_tie_weights_method(self): + """Models without ``tie_weights`` are silently ignored.""" + from auto_round.utils.model import safe_tie_weights + + m = nn.Module() # no tie_weights attribute + safe_tie_weights(m) # must not raise + + def test_drops_stale_entries_before_calling_tie(self): + """Stale tie patterns are pruned before ``tie_weights`` runs.""" + from auto_round.utils.model import safe_tie_weights + + wrapper = nn.Module() + _make_pretrained_submodule( + "text", + params=["real"], + tied={"stale": "real"}, + parent=wrapper, + ) + # tie_weights must be called on the wrapper, not on `text` directly. + wrapper.tie_weights = lambda: None + + safe_tie_weights(wrapper) + + assert wrapper.text._tied_weights_keys == {} + + def test_swallows_value_error_from_tie_weights(self): + """``ValueError`` raised by ``tie_weights`` is caught and warned, not raised.""" + from auto_round.utils.model import safe_tie_weights + + class _M(nn.Module): + def tie_weights(self): + raise ValueError("simulated tie failure") + + m = _M() + # Must not raise. + safe_tie_weights(m) + + def test_non_value_error_is_not_swallowed(self): + """Other exception types propagate to the caller.""" + from auto_round.utils.model import safe_tie_weights + + class _M(nn.Module): + def tie_weights(self): + raise RuntimeError("boom") + + m = _M() + with pytest.raises(RuntimeError, match="boom"): + safe_tie_weights(m) + + +# --------------------------------------------------------------------------- +# 3. Tests for _patch_diffusion_gemma_tied_weights +# --------------------------------------------------------------------------- + + +class TestPatchDiffusionGemmaTiedWeights: + """Unit tests for ``_patch_diffusion_gemma_tied_weights``.""" + + def test_idempotent_when_called_twice(self): + """Calling the patcher twice does not re-wrap ``__init__``.""" + from auto_round.utils import common as ar_common + + # Reset any prior patching state. + ar_common.DiffusionGemmaModel._ar_tied_prune_patched = False + original = ar_common.DiffusionGemmaModel.__init__ + + ar_common._patch_diffusion_gemma_tied_weights() + patched = ar_common.DiffusionGemmaModel.__init__ + ar_common._patch_diffusion_gemma_tied_weights() + second = ar_common.DiffusionGemmaModel.__init__ + + assert patched is second, "second call should be a no-op" + # Restore the original to avoid side effects on other tests. + ar_common.DiffusionGemmaModel.__init__ = original + ar_common.DiffusionGemmaModel._ar_tied_prune_patched = False + + def test_patched_init_calls_prune(self): + """Patched ``__init__`` invokes ``prune_stale_tied_weights_keys`` after the original.""" + from auto_round.utils import common as ar_common + + # Ensure a clean slate. + ar_common.DiffusionGemmaModel._ar_tied_prune_patched = False + original = ar_common.DiffusionGemmaModel.__init__ + + prune_called = [] + + def _fake_prune(model): + prune_called.append(model) + + # The patcher imports ``prune_stale_tied_weights_keys`` lazily from + # ``auto_round.utils.model``; patch it there. + import auto_round.utils.model as ar_model + + with patch.object(ar_model, "prune_stale_tied_weights_keys", _fake_prune): + def _original_init(self, *args, **kwargs): + pass + + ar_common.DiffusionGemmaModel.__init__ = _original_init + ar_common._patch_diffusion_gemma_tied_weights() + + sentinel = types.SimpleNamespace() + ar_common.DiffusionGemmaModel.__init__(sentinel) + + assert len(prune_called) == 1 + assert prune_called[0] is sentinel + + # Restore the original __init__ and clear the patched flag. + ar_common.DiffusionGemmaModel.__init__ = original + ar_common.DiffusionGemmaModel._ar_tied_prune_patched = False + + def test_patched_init_logs_warning_on_prune_failure(self, caplog): + """If ``prune_stale_tied_weights_keys`` raises, the patched init logs a warning.""" + from auto_round.utils import common as ar_common + from auto_round.logger import logger + + ar_common.DiffusionGemmaModel._ar_tied_prune_patched = False + original = ar_common.DiffusionGemmaModel.__init__ + + def _boom(_model): + raise RuntimeError("prune exploded") + + def _original_init(self, *args, **kwargs): + pass + + ar_common.DiffusionGemmaModel.__init__ = _original_init + + # Patch the symbol as imported in ``auto_round.utils.common`` — the + # patcher references it via module attribute, so we monkey-patch at + # that location. + with patch.object(ar_common, "prune_stale_tied_weights_keys", _boom, create=True): + ar_common._patch_diffusion_gemma_tied_weights() + + caplog.set_level("WARNING", logger=logger.name) + # Make sure propagation is enabled so caplog can capture records. + prev_propagate = logger.propagate + logger.propagate = True + try: + sentinel = types.SimpleNamespace() + # Must not raise. + ar_common.DiffusionGemmaModel.__init__(sentinel) + finally: + logger.propagate = prev_propagate + + assert any("prune_stale_tied_weights_keys" in rec.message for rec in caplog.records) + + # Restore. + ar_common.DiffusionGemmaModel.__init__ = original + ar_common.DiffusionGemmaModel._ar_tied_prune_patched = False + + +# --------------------------------------------------------------------------- +# 4. Tests for skip_not_convert_modules default keyword filter +# --------------------------------------------------------------------------- + + +class TestSkipNotConvertModulesDefaultKeywords: + """When the user did not specify ``modules_to_not_convert``, the default + skip set is filtered to the standard keywords: embed / embed_tokens / + lm_head / output_embed / norm.""" + + @staticmethod + def _call(model, user_modules, mock_default): + """Invoke ``skip_not_convert_modules`` with a mocked transformers default-skip helper.""" + from auto_round.inference import convert_model + + cfg = types.SimpleNamespace(modules_to_not_convert=user_modules) + layer_names = [ + "model.embed_tokens.weight", + "model.layers.0.self_attn.q_proj.weight", + "model.layers.0.mlp.gate_up_proj.weight", + "model.norm.weight", + "lm_head.weight", + ] + layer_configs = {n: {"bits": 4} for n in layer_names} + with patch.object(convert_model, "get_modules_to_not_convert", return_value=mock_default): + return convert_model.skip_not_convert_modules(model, cfg, layer_names, layer_configs) + + def test_user_unspecified_keeps_only_default_keywords(self): + """Default-skip keywords are kept; everything else is filtered out.""" + from auto_round.inference import convert_model + + model = nn.Module() + result = self._call( + model, + user_modules=None, + mock_default=["model.embed_tokens", "lm_head", "model.norm"], + ) + # Only default-keyword layers were marked as 16-bit; MLP layers unchanged. + assert result["model.embed_tokens.weight"]["bits"] == 16 + assert result["lm_head.weight"]["bits"] == 16 + assert result["model.norm.weight"]["bits"] == 16 + assert result["model.layers.0.self_attn.q_proj.weight"]["bits"] == 4 + assert result["model.layers.0.mlp.gate_up_proj.weight"]["bits"] == 4 + + def test_user_specified_passes_full_list_through(self): + """When the user explicitly listed modules, no default filtering is applied.""" + from auto_round.inference import convert_model + + model = nn.Module() + result = self._call( + model, + user_modules=["model.layers.0.mlp"], + mock_default=["model.embed_tokens", "model.layers.0.mlp"], + ) + # Both default and user-supplied keywords reach the matcher. + assert result["model.embed_tokens.weight"]["bits"] == 16 + assert result["model.layers.0.self_attn.q_proj.weight"]["bits"] == 4 + # `model.layers.0.mlp` matches `model.layers.0.mlp.gate_up_proj.weight`. + assert result["model.layers.0.mlp.gate_up_proj.weight"]["bits"] == 16 + + def test_empty_everything_no_op(self): + """No skip patterns -> no layer is touched.""" + from auto_round.inference import convert_model + + model = nn.Module() + result = self._call(model, user_modules=None, mock_default=[]) + for name, cfg in result.items(): + assert cfg["bits"] == 4, f"{name} unexpectedly set to 16-bit" + + def test_filter_keeps_keywords_embeds_norm_lm_head_output_embed(self): + """Each of the five default keywords is preserved by the filter.""" + from auto_round.inference import convert_model + + model = nn.Module() + result = self._call( + model, + user_modules=None, + mock_default=[ + "model.embed_tokens", + "model.embed", + "lm_head", + "output_embed", + "model.norm", + "model.layers.0.self_attn", # should be filtered out + "model.layers.0.mlp", # should be filtered out + ], + ) + # The two "model.layers.*" entries should have been filtered out. + assert result["model.layers.0.self_attn.q_proj.weight"]["bits"] == 4 + assert result["model.layers.0.mlp.gate_up_proj.weight"]["bits"] == 4 + # The five default keywords should still apply. + assert result["model.embed_tokens.weight"]["bits"] == 16 + assert result["model.norm.weight"]["bits"] == 16 + assert result["lm_head.weight"]["bits"] == 16 + + +# --------------------------------------------------------------------------- +# 5. Tests for _PRE_DEFINED_FIXED_ATTR registration +# --------------------------------------------------------------------------- + + +class TestPredefinedFixedAttrDiffusionGemma: + """``_PRE_DEFINED_FIXED_ATTR`` exposes DiffusionGemma-specific fixed attrs.""" + + def test_diffusion_gemma_registered(self): + from auto_round.special_model_handler import ( + _PRE_DEFINED_FIXED_ATTR, + get_predefined_fixed_attr, + ) + + assert "diffusion_gemma" in _PRE_DEFINED_FIXED_ATTR + assert _PRE_DEFINED_FIXED_ATTR["diffusion_gemma"] == {"has_variable_block_shape": True} + + def test_get_predefined_fixed_attr_returns_diffusion_gemma_attrs(self): + from auto_round.special_model_handler import get_predefined_fixed_attr + + model = nn.Module() + model.config = types.SimpleNamespace(model_type="diffusion_gemma") + assert get_predefined_fixed_attr(model) == {"has_variable_block_shape": True} + + def test_get_predefined_fixed_attr_returns_none_for_unknown_model(self): + from auto_round.special_model_handler import get_predefined_fixed_attr + + model = nn.Module() + model.config = types.SimpleNamespace(model_type="llama") + assert get_predefined_fixed_attr(model) is None + + def test_get_predefined_fixed_attr_returns_none_when_no_config(self): + from auto_round.special_model_handler import get_predefined_fixed_attr + + assert get_predefined_fixed_attr(nn.Module()) is None From f413a48adc4ae6c3b2a657ae4f83117f7de2f690 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 17 Jun 2026 14:00:57 +0000 Subject: [PATCH 3/7] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- auto_round/inference/convert_model.py | 2 +- auto_round/utils/common.py | 6 ++---- auto_round/utils/model.py | 2 +- test/test_cpu/utils/test_diffusiongemma_support.py | 9 +++------ 4 files changed, 7 insertions(+), 12 deletions(-) diff --git a/auto_round/inference/convert_model.py b/auto_round/inference/convert_model.py index 13d4c6120..afc806ef0 100644 --- a/auto_round/inference/convert_model.py +++ b/auto_round/inference/convert_model.py @@ -71,7 +71,7 @@ def skip_not_convert_modules(model, quantization_config, layer_names, layer_conf modules_to_not_convert = _get_modules_to_not_convert(model, modules_to_not_convert) if modules_to_not_convert and not user_specified: - _DEFAULT_SKIP_KEYWORDS = ("embed", "embed_tokens", "lm_head", "output_embed", "norm") + _DEFAULT_SKIP_KEYWORDS = ("embed", "embed_tokens", "lm_head", "output_embed", "norm") modules_to_not_convert = [ name for name in modules_to_not_convert if any(key in name for key in _DEFAULT_SKIP_KEYWORDS) ] diff --git a/auto_round/utils/common.py b/auto_round/utils/common.py index 812013d78..f4a2e8290 100644 --- a/auto_round/utils/common.py +++ b/auto_round/utils/common.py @@ -23,8 +23,8 @@ import torch import transformers -from transformers.models.diffusion_gemma.modeling_diffusion_gemma import DiffusionGemmaModel from packaging import version +from transformers.models.diffusion_gemma.modeling_diffusion_gemma import DiffusionGemmaModel from auto_round.export.export_to_gguf.config import GGUF_CONFIG from auto_round.logger import logger @@ -261,9 +261,7 @@ def _patched_init(self, *args, **kwargs): try: prune_stale_tied_weights_keys(self) except Exception as exc: # noqa: BLE001 - logger.warning( - f"[DiffusionGemma] prune_stale_tied_weights_keys during __init__ failed: {exc}" - ) + logger.warning(f"[DiffusionGemma] prune_stale_tied_weights_keys during __init__ failed: {exc}") DiffusionGemmaModel.__init__ = _patched_init DiffusionGemmaModel._ar_tied_prune_patched = True diff --git a/auto_round/utils/model.py b/auto_round/utils/model.py index 3bbf76dd8..9ee2885cd 100644 --- a/auto_round/utils/model.py +++ b/auto_round/utils/model.py @@ -23,8 +23,8 @@ import psutil import torch import transformers -from transformers import PreTrainedModel from packaging import version +from transformers import PreTrainedModel from auto_round import envs from auto_round.export.export_to_gguf.config import ModelType diff --git a/test/test_cpu/utils/test_diffusiongemma_support.py b/test/test_cpu/utils/test_diffusiongemma_support.py index ce6387f3a..1c7287123 100644 --- a/test/test_cpu/utils/test_diffusiongemma_support.py +++ b/test/test_cpu/utils/test_diffusiongemma_support.py @@ -20,10 +20,8 @@ import pytest import torch import torch.nn as nn - from transformers import PretrainedConfig, PreTrainedModel - # --------------------------------------------------------------------------- # Test helpers # --------------------------------------------------------------------------- @@ -142,9 +140,7 @@ def test_keeps_common_case_patterns_even_if_unused(self): removed = prune_stale_tied_weights_keys(wrapper) assert removed == 0 - assert wrapper.text._tied_weights_keys == { - "model.embed_tokens.weight": "lm_head.weight" - } + assert wrapper.text._tied_weights_keys == {"model.embed_tokens.weight": "lm_head.weight"} def test_keeps_pattern_when_target_source_count_is_multiple_of_source(self): """Pattern is kept when target param count is a non-zero multiple of source param count.""" @@ -421,6 +417,7 @@ def _fake_prune(model): import auto_round.utils.model as ar_model with patch.object(ar_model, "prune_stale_tied_weights_keys", _fake_prune): + def _original_init(self, *args, **kwargs): pass @@ -439,8 +436,8 @@ def _original_init(self, *args, **kwargs): def test_patched_init_logs_warning_on_prune_failure(self, caplog): """If ``prune_stale_tied_weights_keys`` raises, the patched init logs a warning.""" - from auto_round.utils import common as ar_common from auto_round.logger import logger + from auto_round.utils import common as ar_common ar_common.DiffusionGemmaModel._ar_tied_prune_patched = False original = ar_common.DiffusionGemmaModel.__init__ From 6e666ecae679f5c550bb4a8820dede53a94483ee Mon Sep 17 00:00:00 2001 From: lvliang-intel Date: Thu, 18 Jun 2026 11:46:31 +0800 Subject: [PATCH 4/7] fix ci Signed-off-by: lvliang-intel --- auto_round/utils/common.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/auto_round/utils/common.py b/auto_round/utils/common.py index 812013d78..74a4c2002 100644 --- a/auto_round/utils/common.py +++ b/auto_round/utils/common.py @@ -23,12 +23,17 @@ import torch import transformers -from transformers.models.diffusion_gemma.modeling_diffusion_gemma import DiffusionGemmaModel from packaging import version from auto_round.export.export_to_gguf.config import GGUF_CONFIG from auto_round.logger import logger +try: + from transformers.models.diffusion_gemma.modeling_diffusion_gemma import DiffusionGemmaModel +except ImportError: + # diffusion_gemma is only available in transformers >= 5.11.0. + DiffusionGemmaModel = None # type: ignore[assignment] + def download_audiocaps_csv(): """Download AudioCaps train.csv and return the local cache path. @@ -250,6 +255,9 @@ def _patch_diffusion_gemma_tied_weights(): fused parameter names (``gate_up_proj``, ``down_proj``). The fix is to prune the stale patterns at the moment the model is constructed. """ + if DiffusionGemmaModel is None: + # diffusion_gemma is unavailable (transformers < 5.11.0); nothing to patch. + return if getattr(DiffusionGemmaModel, "_ar_tied_prune_patched", False): return original_init = DiffusionGemmaModel.__init__ From e0bad097e737ef4b671f76ebe8ed44b92ca4409d Mon Sep 17 00:00:00 2001 From: lvliang-intel Date: Thu, 18 Jun 2026 14:01:22 +0800 Subject: [PATCH 5/7] remove redundant line Signed-off-by: lvliang-intel --- auto_round/utils/common.py | 1 - 1 file changed, 1 deletion(-) diff --git a/auto_round/utils/common.py b/auto_round/utils/common.py index 53d92b09d..9e08c803f 100644 --- a/auto_round/utils/common.py +++ b/auto_round/utils/common.py @@ -24,7 +24,6 @@ import torch import transformers from packaging import version -from transformers.models.diffusion_gemma.modeling_diffusion_gemma import DiffusionGemmaModel from auto_round.export.export_to_gguf.config import GGUF_CONFIG from auto_round.logger import logger From d4c1a294acf636cb5f5e2ced3b20e73661d1444a Mon Sep 17 00:00:00 2001 From: lvliang-intel Date: Thu, 25 Jun 2026 15:45:03 +0800 Subject: [PATCH 6/7] fix comments Signed-off-by: lvliang-intel --- auto_round/inference/convert_model.py | 17 ++++------------- 1 file changed, 4 insertions(+), 13 deletions(-) diff --git a/auto_round/inference/convert_model.py b/auto_round/inference/convert_model.py index 72c3fdf9b..c65beb7d6 100644 --- a/auto_round/inference/convert_model.py +++ b/auto_round/inference/convert_model.py @@ -63,24 +63,14 @@ def flatten_list(nested_list): def skip_not_convert_modules(model, quantization_config, layer_names, layer_configs): - user_specified = bool(getattr(quantization_config, "modules_to_not_convert", None)) modules_to_not_convert = getattr(quantization_config, "modules_to_not_convert", []) try: # transformers new api modules_to_not_convert = get_modules_to_not_convert(model, modules_to_not_convert, add_default_skips=True) except: modules_to_not_convert = _get_modules_to_not_convert(model, modules_to_not_convert) - - if modules_to_not_convert and not user_specified: - _DEFAULT_SKIP_KEYWORDS = ("embed", "embed_tokens", "lm_head", "output_embed", "norm") - modules_to_not_convert = [ - name for name in modules_to_not_convert if any(key in name for key in _DEFAULT_SKIP_KEYWORDS) - ] - if modules_to_not_convert: - # Pre-compile patterns once instead of recompiling them for every layer name. - compiled_patterns = [re.compile(n) for n in modules_to_not_convert] for layer_name in layer_names: - if any(pattern.search(layer_name) for pattern in compiled_patterns): + if any([re.search(re.compile(n), layer_name) for n in modules_to_not_convert]): layer_configs[layer_name] = {"bits": 16} return layer_configs @@ -396,8 +386,9 @@ def get_layer_config(model, quantization_config): model=model, ) - # AWQ format: exclude specified modules - extra_config = skip_not_convert_modules(model, quantization_config, layer_names, extra_config) + # AWQ format: exclude specified modules. + if "awq" in (getattr(quantization_config, "quant_method", "") or "").lower(): + extra_config = skip_not_convert_modules(model, quantization_config, layer_names, extra_config) # Expand auto_round regex configs (regex-based) extra_config = _expand_regex_config( From 2fa108550c518228d474baa342b4e1a715a19b6d Mon Sep 17 00:00:00 2001 From: lvliang-intel Date: Thu, 25 Jun 2026 22:54:46 +0800 Subject: [PATCH 7/7] fix accuracy issue with vllm Signed-off-by: lvliang-intel --- auto_round/special_model_handler.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/auto_round/special_model_handler.py b/auto_round/special_model_handler.py index d54d5e5e7..d94628095 100644 --- a/auto_round/special_model_handler.py +++ b/auto_round/special_model_handler.py @@ -1140,6 +1140,16 @@ def get_bagel_ignore_layers(model) -> list[str]: ], ) +# diffusion_gemma +register_ignore_layers( + matchers=[ + ArchitectureMatcher(r"DiffusionGemma", mode="in"), + ], + ignore_layers=[ + "router.proj", + ], +) + def get_predefined_ignore_layers(model: torch.nn.Module) -> list[str]: layers = []