bettertransformer deprecation issue is fixed

This commit is contained in:
alpertunga-bile
2025-04-19 15:33:09 +03:00
parent 1c8436daef
commit 11ba792a35
+32 -37
View File
@@ -2,7 +2,6 @@ from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
)
from optimum.bettertransformer import BetterTransformer
from optimum.onnxruntime import ORTModelForCausalLM
from torch import bfloat16 as torch_bfloat16
@@ -24,6 +23,7 @@ from .utility import (
QuantizationPackage,
get_quantization_package,
is_base_model,
check_torch_version_is_enough
)
from peft import PeftModel
@@ -60,26 +60,29 @@ def get_bitsandbytes_config(type: QuantizationType):
return bnb_config
def get_model_from_base(model_name: str, required_torch_dtype, type: QuantizationType):
def get_model_from_base(model_name: str, required_torch_dtype, type: QuantizationType, is_acceleration: bool):
model_configs = {}
model_configs["device_map"] = "auto"
model_configs["torch_dtype"] = required_torch_dtype
quant_pack = get_quantization_package()
if quant_pack == QuantizationPackage.BITSANDBYTES:
quant_conf = get_bitsandbytes_config(type)
model = AutoModelForCausalLM.from_pretrained(
model_name,
device_map="auto",
torch_dtype=required_torch_dtype,
quantization_config=quant_conf,
)
elif quant_pack == QuantizationPackage.QUANTO:
from optimum.quanto import qfloat8, qint8, qint4, quantize, freeze
model_configs["quantization_config"] = quant_conf
model = AutoModelForCausalLM.from_pretrained(
model_name,
device_map="auto",
torch_dtype=required_torch_dtype,
)
if is_acceleration and check_torch_version_is_enough(2, 1):
model_configs["attn_implementation"] = "sdpa"
model = AutoModelForCausalLM.from_pretrained(
model_name,
**model_configs
)
if quant_pack == QuantizationPackage.QUANTO:
from optimum.quanto import qfloat8, qint8, qint4, quantize, freeze
if type == QuantizationType.EightBit:
quantize(model, weights=qint8)
@@ -89,18 +92,12 @@ def get_model_from_base(model_name: str, required_torch_dtype, type: Quantizatio
quantize(model, weights=qint4)
freeze(model)
elif quant_pack == QuantizationPackage.NONE:
model = AutoModelForCausalLM.from_pretrained(
model_name,
device_map="auto",
torch_dtype=required_torch_dtype,
)
return model
def get_model_from_lora(model_name: str, required_torch_dtype, type: QuantizationType):
model = get_model_from_base(model_name, required_torch_dtype, type)
def get_model_from_lora(model_name: str, required_torch_dtype, type: QuantizationType, is_acceleration: bool):
model = get_model_from_base(model_name, required_torch_dtype, type, is_acceleration)
model.config.forced_decoder_ids = None
model.config.suppress_tokens = []
@@ -110,16 +107,18 @@ def get_model_from_lora(model_name: str, required_torch_dtype, type: Quantizatio
return model
def get_model(model_name: str, type: QuantizationType):
def get_model(model_name: str, type: QuantizationType, is_acceleration: bool):
req_torch_dtype = get_torch_dtype()
if is_base_model(model_name):
model = get_model_from_base(model_name, req_torch_dtype, type)
model = get_model_from_base(model_name, req_torch_dtype, type, is_acceleration)
else:
model = get_model_from_lora(model_name, req_torch_dtype, type)
model = get_model_from_lora(model_name, req_torch_dtype, type, is_acceleration)
# torch.compile is supported only in Linux
# has to be tested though
"""
torch.compile is supported only in Linux
has to be tested though
"""
if system() == "Linux":
torch_compile(model)
@@ -138,18 +137,14 @@ def get_tokenizer(model_name: str):
def get_model_tokenizer(model_path: str, type: ModelType, quant_type: QuantizationType):
"""
transformers -> ONNX operation brokes often
transformers -> ONNX operation brokes often
"""
if type == ModelType.ONNX:
model = ORTModelForCausalLM.from_pretrained(model_path)
elif type == ModelType.BETTERTRANSFORMER or ModelType.DEFAULT:
model = get_model(model_path, quant_type)
if type == ModelType.BETTERTRANSFORMER:
try:
model = BetterTransformer.transform(model)
except:
pass
elif type == ModelType.BETTERTRANSFORMER:
model = get_model(model_path, quant_type, is_acceleration=True)
elif type == ModelType.DEFAULT:
model = get_model(model_path, quant_type, is_acceleration=False)
tokenizer = get_tokenizer(model_path)