From bfc3b135646adfd4cc74c4dd1e19592a41298d43 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 16 Oct 2025 04:45:27 -0700 Subject: [PATCH] Update _utils.py --- unsloth/models/_utils.py | 105 +++++++++++++++++++++++---------------- 1 file changed, 61 insertions(+), 44 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 4616bd3f91..09bcbda28a 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -82,6 +82,7 @@ platform_system = platform_system() import numpy as np import contextlib import re +from dataclasses import dataclass import functools import warnings, subprocess, re, inspect, psutil, os, math from unsloth_zoo.utils import Version @@ -1651,15 +1652,23 @@ def error_out_no_vllm(*args, **kwargs): raise NotImplementedError("Unsloth: vLLM is not yet supported for fast inference for this model! Please use `.generate` instead") -class TorchAOMetadata: - def __init__(self, qat_scheme, base_config, filter_fn, group_size): - self.qat_scheme = qat_scheme - self.base_config = base_config - self.filter_fn = filter_fn - self.group_size = group_size +from torchao.core.config import AOBaseConfig +try: + from torchao.quantization import Int4WeightOnlyConfig +except: + print("Unsloth: TorchAO changed `torchao.quantization.Int4WeightOnlyConfig`") + Int4WeightOnlyConfig = None pass -def _prepare_model_for_qat(model: torch.nn.Module, qat_scheme: str) -> torch.nn.Module: +@dataclass +class TorchAOConfig: + qat_scheme : str = "int4" + base_config : AOBaseConfig = Int4WeightOnlyConfig + filter_fn : Callable = lambda m, _: isinstance(m, torch.nn.Linear) and m.in_features >= group_size + group_size : int = 128 +pass + +def _prepare_model_for_qat(model: torch.nn.Module, qat_scheme: Union[str, TorchAOConfig]) -> torch.nn.Module: """ Transform a model for Quantization-Aware Training (QAT) during fine-tuning. @@ -1670,48 +1679,56 @@ def _prepare_model_for_qat(model: torch.nn.Module, qat_scheme: str) -> torch.nn. QAT can be optionally combined with LoRA fine-tuning to for additional throughput improvement. For more details: https://dev-discuss.pytorch.org/t/speeding-up-qat-by-1-89x-with-lora/2700 """ - from torchao.quantization import ( - Float8DynamicActivationInt4WeightConfig, - Float8DynamicActivationFloat8WeightConfig, - Int8DynamicActivationIntxWeightConfig, - Int4WeightOnlyConfig, - PerRow, - quantize_, - ) + from torchao.quantization import PerRow, quantize_ from torchao.quantization.granularity import PerGroup from torchao.quantization.qat import QATConfig - filter_fn = None - group_size = None - base_config = None - if qat_scheme == "fp8-int4": - group_size = 128 - base_config = Float8DynamicActivationInt4WeightConfig() - filter_fn = lambda m, _: isinstance(m, torch.nn.Linear) and m.in_features >= group_size - elif qat_scheme == "fp8-fp8": - base_config = Float8DynamicActivationFloat8WeightConfig(granularity=PerRow()) - elif qat_scheme == "int8-int4": - group_size = 32 - base_config = Int8DynamicActivationIntxWeightConfig(weight_dtype=torch.int4, weight_granularity=PerGroup(group_size)) - filter_fn = lambda m, _: isinstance(m, torch.nn.Linear) and m.in_features >= group_size - elif qat_scheme == "int4": - group_size = 128 - base_config = Int4WeightOnlyConfig(group_size=group_size) - filter_fn = lambda m, _: isinstance(m, torch.nn.Linear) and m.in_features >= group_size + + if isinstance(qat_scheme, TorchAOConfig): + filter_fn = None + group_size = None + base_config = None + if qat_scheme == "fp8-int4": + from torchao.quantization import Float8DynamicActivationInt4WeightConfig + group_size = 128 + base_config = Float8DynamicActivationInt4WeightConfig() + filter_fn = lambda m, _: isinstance(m, torch.nn.Linear) and m.in_features >= group_size + 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)) + filter_fn = lambda m, _: isinstance(m, torch.nn.Linear) and m.in_features >= group_size + elif qat_scheme == "int4": + from torchao.quantization import Int4WeightOnlyConfig + group_size = 128 + base_config = Int4WeightOnlyConfig(group_size=group_size) + filter_fn = lambda m, _: isinstance(m, torch.nn.Linear) and m.in_features >= group_size + else: + raise ValueError(f"Unexpected QAT scheme {qat_scheme}") + pass + # Save TorchAO schemes + torchao_config = TorchAOConfig( + qat_scheme = qat_scheme, + base_config = base_config, + filter_fn = filter_fn, + group_size = group_size, + ) else: - raise ValueError(f"Unexpected QAT scheme {qat_scheme}") - pass - # Save TorchAO schemes - torchao_metadata = TorchAOMetadata( - qat_scheme = qat_scheme, - base_config = base_config, - filter_fn = filter_fn, - group_size = group_size, - ) + torchao_config = qat_scheme + qat_scheme = torchao_config.qat_scheme + base_config = torchao_config.base_config + filter_fn = torchao_config.filter_fn + group_size = torchao_config.group_size + + # Save Torchao metadata everywhere inner_model = model while hasattr(inner_model, "model"): - inner_model._torchao_metadata = torchao_metadata + inner_model._torchao_config = torchao_config inner_model = model.model - inner_model._torchao_metadata = torchao_metadata - quantize_(model, QATConfig(base_config, step="prepare"), filter_fn=filter_fn) + inner_model._torchao_config = torchao_config + # Quantize with TorchAO + quantize_(model, QATConfig(base_config, step = "prepare"), filter_fn = filter_fn) return model pass