quantization is added
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user