Update _utils.py

This commit is contained in:
Daniel Han 2025-10-16 04:45:27 -07:00
commit bfc3b13564

View file

@ -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