diff --git a/generator/model.py b/generator/model.py index b91eeca..40172f6 100644 --- a/generator/model.py +++ b/generator/model.py @@ -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)