quantization is added

This commit is contained in:
alpertunga-bile
2024-05-04 16:27:34 +03:00
parent 939771d7d3
commit bd3adec818
5 changed files with 264 additions and 109 deletions
+18
View File
@@ -3,6 +3,7 @@ from os.path import dirname, exists
from subprocess import run
from importlib.util import find_spec
from platform import system
from torch import __version__ as torch_version
os_name = system()
@@ -44,3 +45,20 @@ check_package("onnxruntime", "optimum[onnxruntime-gpu]")
# use_fast for tokenizers used this
check_package("sentencepiece", "transformers[sentencepiece]")
def check_torch_version_is_enough(min_major: int, min_minor: int) -> bool:
torch_version_splitted = torch_version.split(".")
torch_version_major = int(torch_version_splitted[0])
torch_version_minor = int(torch_version_splitted[1])
if torch_version_major >= min_major and torch_version_minor >= min_minor:
return True
else:
return False
if os_name == "Linux":
check_package("bitsandbytes", "bitsandbytes")
elif check_torch_version_is_enough(2, 2):
check_package("quanto", "quanto")
+62 -35
View File
@@ -1,19 +1,22 @@
from dataclasses import dataclass
from transformers import Pipeline
from generator.model import (
get_default_pipeline,
get_onnx_pipeline,
get_bettertransformer_pipeline,
from comfy.model_management import get_torch_device
from generator.model import get_model_tokenizer
from generator.utility import (
get_accelerator_type,
get_variable_dictionary,
str_to_quant_type,
)
from generator.utility import get_accelerator_type, get_variable_dictionary
from generator.preprocess import preprocess
from .utility import ModelType
@dataclass
class GenerateArgs:
num_return_sequences: int = 1
return_full_text: bool = False
min_new_tokens: int = 0
max_new_tokens: int = 50
early_stopping: bool = False
@@ -33,30 +36,28 @@ class GenerateArgs:
@dataclass
class Generator:
pipe: Pipeline = None
model = None
tokenizer = None
dev = None
def __init__(
self, model_path: str, is_accelerate: bool, model_quant_type: str
) -> None:
quantize_type = str_to_quant_type(model_quant_type)
def __init__(self, model_path: str, is_accelerate: bool) -> None:
if is_accelerate is False:
self.pipe = get_default_pipeline(model_path)
return
accelerator_type = get_accelerator_type(model_path)
if accelerator_type == "onnx":
self.pipe = get_onnx_pipeline(model_name=model_path, is_native=True)
elif accelerator_type == "bettertransformer":
# onnx pipeline can broke easily so try without onnx pipeline
self.try_wo_onnx_pipeline(model_path)
self.model, self.tokenizer = get_model_tokenizer(
model_path, ModelType.DEFAULT, quantize_type
)
else:
raise ValueError(
"Cant define the accelerator type by folder. Can't find .onnx file for onnx, .bin or .safetensors for bettertransformer and default pipeline. Please check your model"
accelerator_type = get_accelerator_type(model_path)
is_onnx_native = True if accelerator_type == ModelType.ONNX else False
self.model, self.tokenizer = get_model_tokenizer(
model_path, accelerator_type, quantize_type, is_onnx_native
)
# try with bettertransformer first then transformers
def try_wo_onnx_pipeline(self, model_path: str):
try:
self.pipe = get_bettertransformer_pipeline(model_name=model_path)
except:
self.pipe = get_default_pipeline(model_path)
self.dev = get_torch_device()
# generate single output
def generate_text(
@@ -64,14 +65,27 @@ class Generator:
input: str,
args: GenerateArgs = GenerateArgs(),
) -> str:
if self.pipe is None:
raise RuntimeError("Pipeline is NONE. Please check your model path")
if self.model and self.tokenizer is None:
raise RuntimeError(
"Model and tokenizer is NONE. Please check your model path"
)
args.num_return_sequences = 1
args = get_variable_dictionary(args)
output = self.pipe(input, **args)
return output[0]["generated_text"]
inputs = self.tokenizer(input, padding=True, return_tensors="pt").to(self.dev)
generated_ids = self.model.generate(
**inputs,
**get_variable_dictionary(args),
pad_token_id=self.tokenizer.eos_token_id,
)
output = self.tokenizer.decode(
generated_ids[0], skip_special_tokens=True, cleanup_tokenization_spaces=True
)
# output = self.pipe(input, **args)
# return output[0]["generated_text"]
return output
# generate 5 outputs
def generate_multiple_texts(
@@ -79,14 +93,27 @@ class Generator:
input: str,
args: GenerateArgs = GenerateArgs(),
) -> list[str]:
if self.pipe is None:
raise RuntimeError("Pipeline is NONE. Please check your model path")
if self.model and self.tokenizer is None:
raise RuntimeError(
"Model and tokenizer is NONE. Please check your model path"
)
args.num_return_sequences = 5
args = get_variable_dictionary(args)
outputs = self.pipe(input, **args)
return [output["generated_text"] for output in outputs]
inputs = self.tokenizer(input, padding=True, return_tensors="pt").to(self.dev)
generated_ids = self.model.generate(
**inputs,
**get_variable_dictionary(args),
pad_token_id=self.tokenizer.eos_token_id,
)
outputs = self.tokenizer.batch_decode(
generated_ids, skip_special_tokens=True, cleanup_tokenization_spaces=True
)
# outputs = self.pipe(input, **args)
# return [output["generated_text"] for output in outputs]
return outputs
# first generating 5 outputs
+100 -62
View File
@@ -1,30 +1,32 @@
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
Pipeline,
)
from transformers import pipeline as tf_pipe
from optimum.pipelines import pipeline as opt_pipe
from optimum.bettertransformer import BetterTransformer
from optimum.onnxruntime import ORTModelForCausalLM
from torch import bfloat16 as torch_bfloat16
from torch import float16 as torch_float16
from torch import float32 as torch_float32
from torch import compile as torch_compile
from platform import system
from comfy.model_management import (
get_torch_device,
should_use_fp16,
should_use_bf16,
is_device_mps,
)
from torch import bfloat16 as torch_bfloat16
from torch import float16 as torch_float16
from torch import float32 as torch_float32
from torch import compile as torch_compile
from .utility import (
ModelType,
QuantizationType,
QuantizationPackage,
get_quantization_package,
)
def get_model(model_name: str, use_device_map: bool = True):
def get_torch_dtype():
dev = get_torch_device()
if should_use_bf16(device=dev):
@@ -34,13 +36,74 @@ def get_model(model_name: str, use_device_map: bool = True):
else:
req_torch_dtype = torch_float32
return req_torch_dtype
def get_quanto_config(type: QuantizationType):
from transformers import QuantoConfig
quanto_config = None
if type == QuantizationType.EightBit:
quanto_config = QuantoConfig(weights="int8")
elif type == QuantizationType.FourBit:
quanto_config = QuantoConfig(weight="int4")
elif type == QuantizationType.EightFloat:
quanto_config = QuantoConfig(weights="float8")
return quanto_config
def get_bitsandbytes_config(type: QuantizationType):
from transformers import BitsAndBytesConfig
bnb_config = None
if type == QuantizationType.EightBit:
bnb_config = BitsAndBytesConfig(load_in_8bit=True)
elif type == QuantizationType.FourBit:
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=get_torch_dtype(),
)
return bnb_config
def get_quantization_config(type: QuantizationType):
quant_config = None
quant_package = get_quantization_package()
if quant_package == QuantizationPackage.QUANTO:
quant_config = get_quanto_config(type)
elif quant_package == QuantizationPackage.BITSANDBYTES:
quant_config = get_bitsandbytes_config(type)
return quant_config
def get_model(model_name: str, type: QuantizationType, use_device_map: bool = True):
req_torch_dtype = get_torch_dtype()
quant_config = get_quantization_config(type)
if use_device_map:
model = AutoModelForCausalLM.from_pretrained(
model_name, device_map="auto", torch_dtype=req_torch_dtype
model_name,
device_map="auto",
torch_dtype=req_torch_dtype,
quantization_config=quant_config,
)
else:
dev = get_torch_device()
model = AutoModelForCausalLM.from_pretrained(
model_name, device=dev, torch_dtype=req_torch_dtype
model_name,
device=dev,
torch_dtype=req_torch_dtype,
quantization_config=quant_config,
)
# torch.compile is supported only in Linux
@@ -58,59 +121,34 @@ def get_tokenizer(model_name: str):
except:
tokenizer = AutoTokenizer.from_pretrained(model_name, use_fast=False)
tokenizer.pad_token_id = tokenizer.eos_token_id
tokenizer.padding_side = "left"
return tokenizer
def get_default_pipeline(model_name: str) -> Pipeline:
model = get_model(model_name)
tokenizer = get_tokenizer(model_name)
def get_model_tokenizer(
model_path: str,
type: ModelType,
quant_type: QuantizationType,
is_native: bool = True,
):
if type == ModelType.ONNX:
if is_native:
model = ORTModelForCausalLM.from_pretrained(model_path)
else:
model = ORTModelForCausalLM.from_pretrained(model_path, export=True)
elif type == ModelType.BETTERTRANSFORMER or ModelType.DEFAULT:
# is_mps = is_device_mps(get_torch_device())
pipe = tf_pipe(
task="text-generation",
model=model,
tokenizer=tokenizer,
framework="pt",
)
model = get_model(model_path, quant_type, use_device_map=True)
return pipe
if type == ModelType.BETTERTRANSFORMER:
try:
model = BetterTransformer.transform(model)
except:
pass
tokenizer = get_tokenizer(model_path)
def get_onnx_pipeline(model_name: str, is_native: bool = True) -> Pipeline:
if is_native:
model = ORTModelForCausalLM.from_pretrained(model_name)
else:
model = ORTModelForCausalLM.from_pretrained(model_name, export=True)
tokenizer = get_tokenizer(model_name)
pipe = opt_pipe(
task="text-generation",
model=model,
tokenizer=tokenizer,
accelerator="ort",
framework="pt",
)
return pipe
def get_bettertransformer_pipeline(model_name: str) -> Pipeline:
is_mps = is_device_mps(get_torch_device())
if is_mps:
model = get_model(model_name, use_device_map=False)
else:
model = get_model(model_name, use_device_map=True)
tokenizer = get_tokenizer(model_name)
pipe = opt_pipe(
task="text-generation",
model=model,
tokenizer=tokenizer,
accelerator="bettertransformer",
framework="pt",
device=get_torch_device() if is_mps else None,
)
return pipe
return (model, tokenizer)
+72 -4
View File
@@ -1,4 +1,72 @@
from os import listdir
from enum import Enum
from platform import system
from torch import __version__ as torch_version
class ModelType(Enum):
DEFAULT = 1
ONNX = 2
BETTERTRANSFORMER = 3
NONE = 4
def check_torch_version_is_enough(min_major: int, min_minor: int) -> bool:
torch_version_splitted = torch_version.split(".")
torch_version_major = int(torch_version_splitted[0])
torch_version_minor = int(torch_version_splitted[1])
if torch_version_major >= min_major and torch_version_minor >= min_minor:
return True
else:
return False
class QuantizationPackage(Enum):
NONE = 1
QUANTO = 2
BITSANDBYTES = 3
class QuantizationType(Enum):
NONE = 1
EightBit = 2
FourBit = 3
EightFloat = 4
def get_quantization_package() -> QuantizationPackage:
if system() == "Linux":
return QuantizationPackage.BITSANDBYTES
elif check_torch_version_is_enough(2, 2):
return QuantizationPackage.QUANTO
else:
return QuantizationPackage.NONE
def get_usable_quantize_sizes() -> list[str]:
quant_package = get_quantization_package()
quant_sizes = ["none"]
if quant_package == QuantizationPackage.BITSANDBYTES:
quant_sizes = quant_sizes + ["int8", "int4"]
elif quant_package == QuantizationPackage.QUANTO:
quant_sizes = quant_sizes + ["int8", "float8", "int4"]
return quant_sizes
def str_to_quant_type(type_str: str) -> QuantizationType:
quantize_type = QuantizationType.NONE
if type_str == "int8":
quantize_type = QuantizationType.EightBit
elif type_str == "int4":
quantize_type = QuantizationType.FourBit
elif type_str == "float8":
quantize_type = QuantizationType.EightFloat
return quantize_type
def get_variable_dictionary(given_class) -> dict:
@@ -9,16 +77,16 @@ def get_variable_dictionary(given_class) -> dict:
}
def get_accelerator_type(path: str) -> str:
def get_accelerator_type(path: str) -> ModelType:
files = listdir(path)
accelerator_type = "none"
accelerator_type = ModelType.NONE
for file in files:
if file.endswith(".onnx"):
accelerator_type = "onnx"
accelerator_type = ModelType.ONNX
break
if file.endswith(".bin") or file.endswith(".safetensors"):
accelerator_type = "bettertransformer"
accelerator_type = ModelType.BETTERTRANSFORMER
break
return accelerator_type
+12 -8
View File
@@ -8,6 +8,7 @@ from random import randint
from datetime import date
from generator.generate import GenerateArgs, Generator, get_generated_texts
from generator.utility import get_usable_quantize_sizes
from comfy.sd import CLIP
from folder_paths import models_dir, base_path
@@ -21,17 +22,19 @@ class PromptGenerator:
@classmethod
def INPUT_TYPES(s):
quantize_sizes = get_usable_quantize_sizes()
model_names = [
file
for file in listdir(join(models_dir, "prompt_generators"))
if isdir(join(models_dir, "prompt_generators", file))
]
return {
"required": {
"clip": ("CLIP",),
"model_name": (
[
file
for file in listdir(join(models_dir, "prompt_generators"))
if isdir(join(models_dir, "prompt_generators", file))
],
),
"model_name": (model_names,),
"accelerate": (["enable", "disable"],),
"quantize": (quantize_sizes,),
"prompt": (
"STRING",
{
@@ -182,6 +185,7 @@ class PromptGenerator:
clip: CLIP,
model_name: str,
accelerate: str,
quantize: str,
prompt: str,
seed: int,
lock: str,
@@ -262,7 +266,7 @@ class PromptGenerator:
file = open(prompt_log_filename, "w")
file.close()
generator = Generator(model_path, is_accelerate)
generator = Generator(model_path, is_accelerate, quantize)
self._gen_settings = GenerateArgs(
guidance_scale=cfg,