Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions docs/source/features/quantization.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion olive/passes/pytorch/gptq.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
4 changes: 3 additions & 1 deletion olive/passes/pytorch/kquant.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
52 changes: 36 additions & 16 deletions olive/passes/pytorch/quant_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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:"):
Expand Down
4 changes: 3 additions & 1 deletion olive/passes/pytorch/rtn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
200 changes: 200 additions & 0 deletions test/passes/pytorch/test_independent_qkv.py
Original file line number Diff line number Diff line change
@@ -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}
Loading