Extend TorchAOConfig to support mobile usecases (#3587)

* up

* up

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
Scott Roy 2025-11-14 03:08:21 -08:00 committed by GitHub
commit 20bd66f49f

View file

@ -74,7 +74,7 @@ __all__ = [
]
import torch
from typing import Union, Optional, List, Any, Callable, Tuple
from typing import Union, Optional, List, Any, Callable, Tuple, Iterator
from platform import system as platform_system
platform_system = platform_system()
@ -2013,18 +2013,110 @@ except:
@dataclass
class TorchAOConfig:
qat_scheme: str = "int4"
base_config: AOBaseConfig = field(
default_factory = lambda: Int4WeightOnlyConfig(group_size = 128)
)
group_size: int = 128
filter_fn: Optional[Callable] = None
def __post_init__(self):
if self.filter_fn is None:
self.filter_fn = (
# Each (config, filter_fn) pair defines a quantization rule
base_config_and_filter_fns: List[
Tuple["AOBaseConfig", Optional[Callable[[torch.nn.Module, str], bool]]]
] = field(
default_factory = lambda: [
(
Int4WeightOnlyConfig(group_size = 128),
lambda m, _: isinstance(m, torch.nn.Linear)
and m.in_features >= self.group_size
)
and getattr(m, "in_features", 0) >= 128,
),
]
)
# Optional transformation to apply before quantization setup
prequantization_transform: Optional[Callable[[torch.nn.Module], None]] = None
def _untie_input_output_embeddings(model: torch.nn.Module) -> None:
"""
Utility to untie input/output embeddings in a HuggingFace model.
This is useful if we want to quantize the input/ouput embeddings differently.
Model is modified in-place.
"""
# 1) Persist setting in config
if hasattr(model.config, "tie_word_embeddings"):
model.config.tie_word_embeddings = False
# 2) Find input and output embeddings
in_emb = model.get_input_embeddings()
out_proj = model.get_output_embeddings() or getattr(model, "lm_head", None)
if out_proj is None:
raise AttributeError("Couldn't locate output projection (lm_head).")
# (Optional) sanity: shapes should match [vocab, hidden]
assert (
out_proj.weight.shape == in_emb.weight.shape
), f"Shape mismatch: out_proj {out_proj.weight.shape} vs in_emb {in_emb.weight.shape}"
# 3) Only clone if they are actually tied (shared storage)
if out_proj.weight.data_ptr() == in_emb.weight.data_ptr():
with torch.no_grad():
W = in_emb.weight.detach().clone()
out_proj.weight = torch.nn.Parameter(W) # new storage, keeps dtype/device
# 4) Prevent future automatic re-tying
def _no_tie(self):
return
model.tie_weights = _no_tie.__get__(model, model.__class__)
# 5) Verify no shared storage
assert (
out_proj.weight.data_ptr() != in_emb.weight.data_ptr()
), "Embeddings still tied!"
def _filter_fn_to_fqns(
model: torch.nn.Module,
filter_fn: Callable[[torch.nn.Module, str], bool],
) -> Iterator[str]:
"""
Given a model and a filter function (m, fqn) -> bool,
yield fully qualified names (FQNs) of modules that match.
"""
for fqn, module in model.named_modules():
if filter_fn(module, fqn):
yield fqn
def _convert_torchao_model(model):
from transformers import TorchAoConfig
from torchao.quantization import quantize_, ModuleFqnToConfig
from torchao.quantization.qat import QATConfig
from torchao.utils import TorchAOBaseTensor
module_to_fqn_dict = {}
for base_config, filter_fn in model._torchao_config.base_config_and_filter_fns:
quantize_(model, QATConfig(base_config, step = "convert"), filter_fn = filter_fn)
# Default filter function used for quantize_
if filter_fn is None:
if "_default" in module_to_fqn_dict:
raise ValueError("Cannot use multiple default quantization configs")
module_to_fqn_dict["_default"] = base_config
else:
for fqn in _filter_fn_to_fqns(model, filter_fn):
if fqn in module_to_fqn_dict:
raise ValueError(f"Found multiple quantization configs for {fqn}")
module_to_fqn_dict[fqn] = base_config
in_emb = model.get_input_embeddings()
out_proj = model.get_output_embeddings() or getattr(model, "lm_head", None)
kwargs = {}
if isinstance(in_emb.weight, TorchAOBaseTensor) or (
out_proj is not None and isinstance(out_proj.weight, TorchAOBaseTensor)
):
kwargs["include_input_output_embeddings"] = True
kwargs["modules_to_not_convert"] = []
quant_config = ModuleFqnToConfig(module_to_fqn_dict)
quantization_config = TorchAoConfig(quant_type = quant_config, **kwargs)
model.config.quantization_config = quantization_config
def _prepare_model_for_qat(
@ -2041,13 +2133,11 @@ def _prepare_model_for_qat(
For more details: https://dev-discuss.pytorch.org/t/speeding-up-qat-by-1-89x-with-lora/2700
"""
from torchao.quantization import PerRow, quantize_
from torchao.quantization.granularity import PerGroup
from torchao.quantization.granularity import PerGroup, PerAxis
from torchao.quantization.qat import QATConfig
if not isinstance(qat_scheme, TorchAOConfig):
filter_fn = None
group_size = None
base_config = None
torchao_config: Optional[TorchAOConfig] = None
if qat_scheme == "fp8-int4":
from torchao.quantization import Float8DynamicActivationInt4WeightConfig
@ -2057,22 +2147,42 @@ def _prepare_model_for_qat(
lambda m, _: isinstance(m, torch.nn.Linear)
and m.in_features >= group_size
)
torchao_config = TorchAOConfig(
qat_scheme = qat_scheme,
base_config_and_filter_fns = [(base_config, filter_fn)],
)
elif qat_scheme == "fp8-fp8":
from torchao.quantization import Float8DynamicActivationFloat8WeightConfig
base_config = Float8DynamicActivationFloat8WeightConfig(
granularity = PerRow()
)
elif qat_scheme == "int8-int4":
from torchao.quantization import Int8DynamicActivationIntxWeightConfig
group_size = 32
base_config = Int8DynamicActivationIntxWeightConfig(
weight_dtype = torch.int4, weight_granularity = PerGroup(group_size)
torchao_config = TorchAOConfig(
qat_scheme = qat_scheme, base_config_and_filter_fns = [(base_config, None)]
)
filter_fn = (
lambda m, _: isinstance(m, torch.nn.Linear)
and m.in_features >= group_size
elif qat_scheme == "int8-int4":
from torchao.quantization import (
Int8DynamicActivationIntxWeightConfig,
IntxWeightOnlyConfig,
)
torchao_config = TorchAOConfig(
qat_scheme = qat_scheme,
base_config_and_filter_fns = [
(
IntxWeightOnlyConfig(
weight_dtype = torch.int8, granularity = PerAxis(0)
),
lambda m, fqn: isinstance(m, torch.nn.Embedding),
),
(
Int8DynamicActivationIntxWeightConfig(
weight_dtype = torch.int4, weight_granularity = PerGroup(32)
),
None,
),
],
prequantization_transform = _untie_input_output_embeddings,
)
elif qat_scheme == "int4":
from torchao.quantization import Int4WeightOnlyConfig
@ -2083,21 +2193,15 @@ def _prepare_model_for_qat(
lambda m, _: isinstance(m, torch.nn.Linear)
and m.in_features >= group_size
)
torchao_config = TorchAOConfig(
qat_scheme = qat_scheme,
base_config_and_filter_fns = [(base_config, filter_fn)],
)
else:
raise ValueError(f"Unexpected QAT scheme {qat_scheme}")
# Save TorchAO schemes
torchao_config = TorchAOConfig(
qat_scheme = qat_scheme,
base_config = base_config,
group_size = group_size,
filter_fn = filter_fn,
)
assert torchao_config is not None, f"TorchAOConfig was not set for {qat_scheme}"
else:
torchao_config = qat_scheme
qat_scheme = torchao_config.qat_scheme
base_config = torchao_config.base_config
group_size = torchao_config.group_size
filter_fn = torchao_config.filter_fn
# Save Torchao metadata everywhere
inner_model = model
@ -2105,8 +2209,12 @@ def _prepare_model_for_qat(
inner_model._torchao_config = torchao_config
inner_model = inner_model.model
inner_model._torchao_config = torchao_config
# Quantize with TorchAO
quantize_(model, QATConfig(base_config, step = "prepare"), filter_fn = filter_fn)
if torchao_config.prequantization_transform is not None:
torchao_config.prequantization_transform(model)
for base_config, filter_fn in torchao_config.base_config_and_filter_fns:
quantize_(model, QATConfig(base_config, step = "prepare"), filter_fn = filter_fn)
return model