Update _utils.py
This commit is contained in:
parent
92e281f2ee
commit
bfc3b13564
1 changed files with 61 additions and 44 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue