bettertransformer deprecation issue is fixed
This commit is contained in:
+32
-37
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user