From 086627b17dc48a8b96d3fb29451644bf63f23e3e Mon Sep 17 00:00:00 2001 From: titaiwangms Date: Tue, 29 Sep 2026 00:58:22 +0000 Subject: [PATCH] Allow independent QKV precision in PyTorch quantization passes Preserve per-projection RTN, KQuant, and GPTQ overrides behind an explicit opt-in while retaining the existing shared-QKV default for packed-weight exporters. Cover fresh checkpoints, locked checkpoint merges, serialization, and tiny Qwen3-MoE fused experts without changing SMP or downstream export semantics. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: titaiwangms --- docs/source/features/quantization.md | 14 ++ olive/passes/pytorch/gptq.py | 2 +- olive/passes/pytorch/kquant.py | 4 +- olive/passes/pytorch/quant_utils.py | 52 +++-- olive/passes/pytorch/rtn.py | 4 +- test/passes/pytorch/test_independent_qkv.py | 200 ++++++++++++++++++++ 6 files changed, 257 insertions(+), 19 deletions(-) create mode 100644 test/passes/pytorch/test_independent_qkv.py diff --git a/docs/source/features/quantization.md b/docs/source/features/quantization.md index 4a02f39f4f..c3860c5fb3 100644 --- a/docs/source/features/quantization.md +++ b/docs/source/features/quantization.md @@ -258,6 +258,20 @@ Re-run the `Rtn` pass on the original full-precision model to regenerate a check raises a clear `ValueError` at export time rather than silently producing an incorrect graph; 2-bit quantization remains usable for PyTorch-only workflows. +### Independent Q/K/V settings + +`Rtn`, `KQuant`, and the native PyTorch `Gptq` pass accept +`independent_qkv: true` to preserve separate quantization settings for split +attention Q/K/V projections. For example, with `bits: 4`, set +`overrides: {"re:.*\\.self_attn\\.v_proj": {"bits": 8}}` to quantize Q/K at +4 bits and V at 8 bits. This flag does not select V precision on its own. +By default (`false`), Q/K/V settings are promoted to a shared config for +packed-QKV consumers. Set the flag again on each follow-up quantization pass +that should keep independent settings; it is not stored in the checkpoint. +Existing quantized weights are not re-quantized. This option does not change +`SelectiveMixedPrecision` score allocation or make packed-QKV ONNX exporters +compatible with mixed widths; use an exporter that keeps Q/K/V projections separate. + ## PyTorch Native KQuant The `KQuant` pass is a calibration-free weight quantizer that applies llama.cpp's diff --git a/olive/passes/pytorch/gptq.py b/olive/passes/pytorch/gptq.py index 778a521d9a..900b6459e8 100644 --- a/olive/passes/pytorch/gptq.py +++ b/olive/passes/pytorch/gptq.py @@ -45,7 +45,7 @@ class Gptq(Pass): @classmethod def _default_config(cls, accelerator_spec: AcceleratorSpec) -> dict[str, PassConfigParam]: return { - **get_quantizer_config(allow_moe=True), + **get_quantizer_config(allow_moe=True, allow_independent_qkv=True), "damp_percent": PassConfigParam( type_=float, default_value=0.01, diff --git a/olive/passes/pytorch/kquant.py b/olive/passes/pytorch/kquant.py index 74712f6df9..6db51a70a3 100644 --- a/olive/passes/pytorch/kquant.py +++ b/olive/passes/pytorch/kquant.py @@ -289,7 +289,9 @@ class KQuant(Pass): @classmethod def _default_config(cls, accelerator_spec: AcceleratorSpec) -> dict[str, PassConfigParam]: - config = get_quantizer_config(allow_embeds=True, allow_moe=True, auto_component_targets=True) + config = get_quantizer_config( + allow_embeds=True, allow_moe=True, auto_component_targets=True, allow_independent_qkv=True + ) config["group_size"] = PassConfigParam( type_=int, default_value=32, diff --git a/olive/passes/pytorch/quant_utils.py b/olive/passes/pytorch/quant_utils.py index 689f975f49..a0e515bcfc 100644 --- a/olive/passes/pytorch/quant_utils.py +++ b/olive/passes/pytorch/quant_utils.py @@ -50,6 +50,7 @@ def get_quantizer_config( allow_embeds: bool = False, allow_moe: bool = False, auto_component_targets: bool = False, + allow_independent_qkv: bool = False, ) -> dict[str, PassConfigParam]: return { "bits": PassConfigParam( @@ -129,6 +130,21 @@ def get_quantizer_config( if allow_moe else {} ), + **( + { + "independent_qkv": PassConfigParam( + type_=bool, + default_value=False, + description=( + "Preserve independent Q/K/V quantization settings rather than promoting split " + "attention projections to a shared config. Default is False for packed-QKV " + "compatibility. Must be set again on follow-up quantization passes." + ), + ) + } + if allow_independent_qkv + else {} + ), "modules_to_not_convert": PassConfigParam( type_=list, default_value=None, @@ -916,12 +932,14 @@ def prepare_model( excluded_attn_inputs = _collect_excluded_attn_inputs(wrapper) if exclude_attn_inputs else set() selected_module_names = {_root_module_name(name, name_prefix) for name, _ in wrapper.model.named_modules()} - fresh_qcfg = normalize_qkv_quant_config( - wrapper, - get_quant_config(model, config, existing_qcfg), - module_names=selected_module_names, - name_prefix=name_prefix, - ) + fresh_qcfg = get_quant_config(model, config, existing_qcfg) + if not getattr(config, "independent_qkv", False): + fresh_qcfg = normalize_qkv_quant_config( + wrapper, + fresh_qcfg, + module_names=selected_module_names, + name_prefix=name_prefix, + ) originally_tied_embeddings = getattr(wrapper.config, "tie_word_embeddings", False) wrapper.olive_originally_tied_embeddings = originally_tied_embeddings @@ -1076,13 +1094,14 @@ def _iter_component_quant_targets( fresh_qcfg, "quantize_vision", False ) qcfg = OliveHfQuantizationConfig(**merged) - qcfg = normalize_qkv_quant_config( - wrapper, - qcfg, - locked_modules=already_quantized, - module_names=selected_module_names, - name_prefix=name_prefix, - ) + if not getattr(config, "independent_qkv", False): + qcfg = normalize_qkv_quant_config( + wrapper, + qcfg, + locked_modules=already_quantized, + module_names=selected_module_names, + name_prefix=name_prefix, + ) else: qcfg = fresh_qcfg @@ -1170,9 +1189,10 @@ def _iter_component_quant_targets( # Drop overrides for modules that won't be quantized this pass. Pre-existing (on-disk) # overrides are preserved verbatim since they describe already-quantized weights. - # QKV-group overrides for modules excluded from this pass are not kept: when the - # follow-up pass runs, the quantized members in the group will be locked and pull the - # remaining members back into the shared config via ``normalize_qkv_quant_config``. + # QKV-group overrides for modules excluded from this pass are not kept: a + # follow-up pass with default normalization pulls remaining members into the + # shared config from locked quantized members. Opt-in independent passes can + # instead provide fresh overrides for those remaining members. for name in list(qcfg.overrides or {}): # ``re:`` keys aren't tied to a specific module, so leave them in place. if name.startswith("re:"): diff --git a/olive/passes/pytorch/rtn.py b/olive/passes/pytorch/rtn.py index b4b8fdbd04..7f93d6da51 100644 --- a/olive/passes/pytorch/rtn.py +++ b/olive/passes/pytorch/rtn.py @@ -27,7 +27,9 @@ class Rtn(Pass): @classmethod def _default_config(cls, accelerator_spec: AcceleratorSpec) -> dict[str, PassConfigParam]: - return get_quantizer_config(allow_embeds=True, allow_moe=True, auto_component_targets=True) + return get_quantizer_config( + allow_embeds=True, allow_moe=True, auto_component_targets=True, allow_independent_qkv=True + ) @torch.no_grad() def _run_for_config( diff --git a/test/passes/pytorch/test_independent_qkv.py b/test/passes/pytorch/test_independent_qkv.py new file mode 100644 index 0000000000..02caa8ff61 --- /dev/null +++ b/test/passes/pytorch/test_independent_qkv.py @@ -0,0 +1,200 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Invocation-local independent Q/K/V quantization on tiny offline models.""" + +from pathlib import Path + +import pytest +from pydantic import ValidationError + +from olive.common.quant.tensor import QuantTensor +from olive.passes.olive_pass import create_pass_from_dict +from olive.passes.pytorch.autoclip import AutoClip +from olive.passes.pytorch.gptq import Gptq +from olive.passes.pytorch.kquant import KQuant +from olive.passes.pytorch.rtn import Rtn +from test.passes.pytorch.test_quantization_utils import ( + assert_packed_quant_module, + assert_saved_quant_tensor_matches, + make_local_calibration_data_config, + make_local_tiny_dense_llama, +) +from test.passes.pytorch.test_rtn import _make_local_tiny_qwen3_moe + +QKV = "model.layers.0.self_attn" +V = f"{QKV}.v_proj" +Q = f"{QKV}.q_proj" +K = f"{QKV}.k_proj" + + +def _run(pass_type, input_model, output_path: Path, **options): + if pass_type is Gptq: + options["data_config"] = make_local_calibration_data_config(seq_len=8, max_samples=2) + options["desc_act"] = False + return create_pass_from_dict(pass_type, {"group_size": 16, **options}, disable_search=True).run( + input_model, str(output_path) + ) + + +def _assert_qkv(output, path: Path, expected: tuple[int, int, int], *, group_size=16, symmetric=False): + loaded = output.load_model() + qcfg = loaded.config.quantization_config + assert "independent_qkv" not in qcfg.to_dict() + for proj, bits in zip(("q_proj", "k_proj", "v_proj"), expected): + name = f"{QKV}.{proj}" + assert qcfg.get_qlinear_init_args(name) == { + "bits": bits, + "symmetric": symmetric, + "group_size": group_size, + } + tensor = assert_packed_quant_module( + loaded.get_submodule(name), bits=bits, group_size=group_size, symmetric=symmetric + ) + assert_saved_quant_tensor_matches(path, name, tensor) + return loaded + + +def test_only_native_quantizers_expose_independent_qkv(): + for pass_type in (Rtn, KQuant, Gptq): + quantizer = create_pass_from_dict(pass_type, {}, disable_search=True) + assert quantizer.config.independent_qkv is False + assert "independent_qkv" not in AutoClip._default_config(None) + with pytest.raises(ValidationError, match="independent_qkv"): + create_pass_from_dict(Rtn, {"independent_qkv": "not-a-bool"}, disable_search=True) + + +@pytest.mark.parametrize("pass_type", [Rtn, KQuant, Gptq]) +@pytest.mark.parametrize("flag", [None, False]) +def test_default_still_promotes_fresh_qkv(tmp_path: Path, pass_type, flag): + model = make_local_tiny_dense_llama(tmp_path / "input") + options = {"bits": 4, "overrides": {V: {"bits": 8}}} + if flag is not None: + options["independent_qkv"] = flag + path = tmp_path / "default" + output = _run(pass_type, model, path, **options) + _assert_qkv(output, path, (8, 8, 8)) + + +@pytest.mark.parametrize( + ("pass_type", "v_override"), + [ + (Rtn, V), + (Rtn, r"re:.*\.self_attn\.v_proj"), + (KQuant, r"re:.*\.self_attn\.v_proj"), + (Gptq, r"re:.*\.self_attn\.v_proj"), + ], +) +def test_opt_in_materializes_independent_qkv(tmp_path: Path, pass_type, v_override): + model = make_local_tiny_dense_llama(tmp_path / "input") + path = tmp_path / "independent" + output = _run( + pass_type, + model, + path, + bits=4, + independent_qkv=True, + overrides={v_override: {"bits": 8}}, + ) + _assert_qkv(output, path, (4, 4, 8)) + + +def test_qwen3_moe_independent_qkv_preserves_fused_experts(tmp_path: Path): + model = _make_local_tiny_qwen3_moe(tmp_path / "input") + path = tmp_path / "qwen3_moe" + output = _run( + Rtn, + model, + path, + bits=4, + group_size=-1, + moe=True, + independent_qkv=True, + overrides={V: {"bits": 8}}, + ) + loaded = _assert_qkv(output, path, (4, 4, 8), group_size=-1) + experts = loaded.model.layers[0].mlp.experts + assert isinstance(experts.gate_up_proj.data, QuantTensor) + assert experts.gate_up_proj.data.bits == 4 + assert isinstance(experts.down_proj.data, QuantTensor) + assert experts.down_proj.data.bits == 4 + + +@pytest.mark.parametrize("pass_type", [Rtn, KQuant]) +@pytest.mark.parametrize("opt_in", [False, True]) +def test_locked_v_from_checkpoint_defaults(tmp_path: Path, pass_type, opt_in): + model = make_local_tiny_dense_llama(tmp_path / "input") + # V is physically INT8 but has no explicit checkpoint override. + first_path = tmp_path / "first" + first = _run( + Rtn, + model, + first_path, + bits=8, + modules_to_not_convert=[Q, K], + ) + original_v = _assert_v(first, first_path) + path = tmp_path / "second" + second = _run( + pass_type, + first, + path, + bits=4, + independent_qkv=opt_in, + ) + loaded = _assert_qkv(second, path, (4, 4, 8) if opt_in else (8, 8, 8)) + assert loaded.get_submodule(V)._parameters["weight"].qweight.equal(original_v.qweight) + + +def _assert_v(output, path): + loaded = output.load_model() + assert V not in (loaded.config.quantization_config.overrides or {}) + tensor = assert_packed_quant_module(loaded.get_submodule(V), bits=8, group_size=16, symmetric=False) + assert_saved_quant_tensor_matches(path, V, tensor) + return tensor + + +@pytest.mark.parametrize("pass_type", [Rtn, KQuant]) +def test_conflicting_locked_qv_and_excluded_k(tmp_path: Path, pass_type): + model = make_local_tiny_dense_llama(tmp_path / "input") + first_path = tmp_path / "first" + first = _run( + Rtn, + model, + first_path, + bits=4, + independent_qkv=True, + overrides={V: {"bits": 8}}, + modules_to_not_convert=[K], + ) + first_loaded = first.load_model() + q_weight = first_loaded.get_submodule(Q)._parameters["weight"].qweight.clone() + v_weight = first_loaded.get_submodule(V)._parameters["weight"].qweight.clone() + path = tmp_path / "second" + second = _run(pass_type, first, path, bits=4, independent_qkv=True) + loaded = _assert_qkv(second, path, (4, 4, 8)) + assert loaded.get_submodule(Q)._parameters["weight"].qweight.equal(q_weight) + assert loaded.get_submodule(V)._parameters["weight"].qweight.equal(v_weight) + + +def test_independent_qkv_keeps_exclusions_and_other_quant_settings(tmp_path: Path): + model = make_local_tiny_dense_llama(tmp_path / "input") + path = tmp_path / "excluded" + output = _run( + Rtn, + model, + path, + bits=4, + sym=True, + independent_qkv=True, + overrides={V: {"bits": 8, "group_size": 32, "symmetric": False}}, + modules_to_not_convert=[K], + ) + loaded = output.load_model() + qcfg = loaded.config.quantization_config + assert_packed_quant_module(loaded.get_submodule(Q), bits=4, group_size=16, symmetric=True) + assert_packed_quant_module(loaded.get_submodule(V), bits=8, group_size=32, symmetric=False) + assert_saved_quant_tensor_matches(path, V, loaded.get_submodule(V)._parameters["weight"]) + assert not hasattr(loaded.get_submodule(K)._parameters["weight"].data, "qweight") + assert qcfg.get_qlinear_init_args(V) == {"bits": 8, "group_size": 32, "symmetric": False}