From 20bd66f49f514dc000c372b5ff2aac4da4530a45 Mon Sep 17 00:00:00 2001 From: Scott Roy <161522778+metascroy@users.noreply.github.com> Date: Fri, 14 Nov 2025 03:08:21 -0800 Subject: [PATCH] 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> --- unsloth/models/_utils.py | 182 +++++++++++++++++++++++++++++++-------- 1 file changed, 145 insertions(+), 37 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index ba168325d4..99e41c13c8 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -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