mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
@@ -1,11 +1,35 @@
|
||||
from pathlib import Path
|
||||
import os
|
||||
import time
|
||||
import torch
|
||||
import gguf
|
||||
from .gguf_utils import TORCH_COMPATIBLE_QTYPES
|
||||
from .gguf_tensor import GGMLTensor
|
||||
import diffusers
|
||||
import transformers
|
||||
|
||||
|
||||
def load_gguf_state_dict(path: str, compute_dtype: torch.dtype) -> dict[str, GGMLTensor]:
|
||||
def install_gguf():
|
||||
# pip install git+https://github.com/junejae/transformers@feature/t5-gguf
|
||||
# https://github.com/ggerganov/llama.cpp/issues/9566
|
||||
from installer import install
|
||||
install('gguf', quiet=True)
|
||||
import importlib
|
||||
import gguf
|
||||
from modules import shared
|
||||
scripts_dir = os.path.join(os.path.dirname(gguf.__file__), '..', 'scripts')
|
||||
if os.path.exists(scripts_dir):
|
||||
os.rename(scripts_dir, scripts_dir + str(time.time()))
|
||||
# monkey patch transformers/diffusers so they detect newly installed gguf pacakge correctly
|
||||
ver = importlib.metadata.version('gguf')
|
||||
transformers.utils.import_utils._is_gguf_available = True # pylint: disable=protected-access
|
||||
transformers.utils.import_utils._gguf_version = ver # pylint: disable=protected-access
|
||||
diffusers.utils.import_utils._is_gguf_available = True # pylint: disable=protected-access
|
||||
diffusers.utils.import_utils._gguf_version = ver # pylint: disable=protected-access
|
||||
shared.log.debug(f'Load GGUF: version={ver}')
|
||||
return gguf
|
||||
|
||||
|
||||
def load_gguf_state_dict(path: str, compute_dtype: torch.dtype) -> dict:
|
||||
gguf = install_gguf()
|
||||
from .gguf_utils import TORCH_COMPATIBLE_QTYPES
|
||||
from .gguf_tensor import GGMLTensor
|
||||
sd: dict[str, GGMLTensor] = {}
|
||||
stats = {}
|
||||
reader = gguf.GGUFReader(path)
|
||||
@@ -19,3 +43,14 @@ def load_gguf_state_dict(path: str, compute_dtype: torch.dtype) -> dict[str, GGM
|
||||
stats[tensor.tensor_type.name] = 0
|
||||
stats[tensor.tensor_type.name] += 1
|
||||
return sd, stats
|
||||
|
||||
|
||||
def load_gguf(path, cls, compute_dtype: torch.dtype):
|
||||
_gguf = install_gguf()
|
||||
module = cls.from_single_file(
|
||||
path,
|
||||
quantization_config = diffusers.GGUFQuantizationConfig(compute_dtype=compute_dtype),
|
||||
torch_dtype=compute_dtype,
|
||||
)
|
||||
module.gguf = 'gguf'
|
||||
return module
|
||||
|
||||
+14
-4
@@ -160,9 +160,10 @@ def load_quants(kwargs, repo_id, cache_dir):
|
||||
return kwargs
|
||||
|
||||
|
||||
"""
|
||||
def load_flux_gguf(file_path):
|
||||
transformer = None
|
||||
model_te.install_gguf()
|
||||
ggml.install_gguf()
|
||||
from accelerate import init_empty_weights
|
||||
from diffusers.loaders.single_file_utils import convert_flux_transformer_checkpoint_to_diffusers
|
||||
from modules import ggml, sd_hijack_accelerate
|
||||
@@ -180,9 +181,11 @@ def load_flux_gguf(file_path):
|
||||
continue
|
||||
applied += 1
|
||||
sd_hijack_accelerate.hijack_set_module_tensor_simple(transformer, tensor_name=param_name, value=param, device=0)
|
||||
transformer.gguf = 'gguf'
|
||||
state_dict[param_name] = None
|
||||
shared.log.debug(f'Load model: type=Unet/Transformer applied={applied} skipped={skipped} stats={stats}')
|
||||
return transformer, None
|
||||
"""
|
||||
|
||||
|
||||
def load_transformer(file_path): # triggered by opts.sd_unet change
|
||||
@@ -197,7 +200,9 @@ def load_transformer(file_path): # triggered by opts.sd_unet change
|
||||
}
|
||||
shared.log.info(f'Load module: type=UNet/Transformer file="{file_path}" offload={shared.opts.diffusers_offload_mode} quant={quant} dtype={devices.dtype}')
|
||||
if 'gguf' in file_path.lower():
|
||||
_transformer, _text_encoder_2 = load_flux_gguf(file_path)
|
||||
# _transformer, _text_encoder_2 = load_flux_gguf(file_path)
|
||||
from modules import ggml
|
||||
_transformer = ggml.load_gguf(file_path, cls=diffusers.FluxTransformer2DModel, compute_dtype=devices.dtype)
|
||||
if _transformer is not None:
|
||||
transformer = _transformer
|
||||
elif quant == 'qint8' or quant == 'qint4':
|
||||
@@ -336,9 +341,14 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
|
||||
cls = diffusers.FluxPipeline
|
||||
shared.log.debug(f'Load model: type=FLUX cls={cls.__name__} preloaded={list(kwargs)} revision={diffusers_load_config.get("revision", None)}')
|
||||
for c in kwargs:
|
||||
if getattr(kwargs[c], 'quantization_method', None) is not None or getattr(kwargs[c], 'gguf', None) is not None:
|
||||
shared.log.debug(f'Load model: type=FLUX component={c} dtype={kwargs[c].dtype} quant={getattr(kwargs[c], 'quantization_method', None) or getattr(kwargs[c], 'gguf', None)}')
|
||||
if kwargs[c].dtype == torch.float32 and devices.dtype != torch.float32:
|
||||
shared.log.warning(f'Load model: type=FLUX component={c} dtype={kwargs[c].dtype} cast dtype={devices.dtype} recast')
|
||||
kwargs[c] = kwargs[c].to(dtype=devices.dtype)
|
||||
try:
|
||||
kwargs[c] = kwargs[c].to(dtype=devices.dtype)
|
||||
shared.log.warning(f'Load model: type=FLUX component={c} dtype={kwargs[c].dtype} cast dtype={devices.dtype} recast')
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
allow_quant = 'gguf' not in (sd_unet.loaded_unet or '')
|
||||
fn = checkpoint_info.path
|
||||
|
||||
+10
-3
@@ -13,7 +13,9 @@ def load_overrides(kwargs, cache_dir):
|
||||
sd_unet.loaded_unet = shared.opts.sd_unet
|
||||
shared.log.debug(f'Load model: type=SD3 unet="{shared.opts.sd_unet}" fmt=safetensors')
|
||||
elif fn.endswith('.gguf'):
|
||||
kwargs = load_gguf(kwargs, fn)
|
||||
from modules import ggml
|
||||
# kwargs = load_gguf(kwargs, fn)
|
||||
kwargs['transformer'] = ggml.load_gguf(fn, cls=diffusers.SD3Transformer2DModel, compute_dtype=devices.dtype)
|
||||
sd_unet.loaded_unet = shared.opts.sd_unet
|
||||
shared.log.debug(f'Load model: type=SD3 unet="{shared.opts.sd_unet}" fmt=gguf')
|
||||
except Exception as e:
|
||||
@@ -90,8 +92,9 @@ def load_missing(kwargs, fn, cache_dir):
|
||||
return kwargs
|
||||
|
||||
|
||||
"""
|
||||
def load_gguf(kwargs, fn):
|
||||
model_te.install_gguf()
|
||||
ggml.install_gguf()
|
||||
from accelerate import init_empty_weights
|
||||
from diffusers.loaders.single_file_utils import convert_sd3_transformer_checkpoint_to_diffusers
|
||||
from modules import ggml, sd_hijack_accelerate
|
||||
@@ -108,10 +111,12 @@ def load_gguf(kwargs, fn):
|
||||
continue
|
||||
applied += 1
|
||||
sd_hijack_accelerate.hijack_set_module_tensor_simple(transformer, tensor_name=param_name, value=param, device=0)
|
||||
transformer.gguf = 'gguf'
|
||||
state_dict[param_name] = None
|
||||
shared.log.debug(f'Load model: type=Unet/Transformer applied={applied} skipped={skipped} stats={stats} compute={devices.dtype}')
|
||||
kwargs['transformer'] = transformer
|
||||
return kwargs
|
||||
"""
|
||||
|
||||
|
||||
def load_sd3(checkpoint_info, cache_dir=None, config=None):
|
||||
@@ -139,7 +144,9 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None):
|
||||
# kwargs = load_missing(kwargs, fn, cache_dir)
|
||||
repo_id = fn
|
||||
elif fn.endswith('.gguf'):
|
||||
kwargs = load_gguf(kwargs, fn)
|
||||
from modules import ggml
|
||||
kwargs['transformer'] = ggml.load_gguf(fn, cls=diffusers.SD3Transformer2DModel, compute_dtype=devices.dtype)
|
||||
# kwargs = load_gguf(kwargs, fn)
|
||||
kwargs = load_missing(kwargs, fn, cache_dir)
|
||||
kwargs['variant'] = 'fp16'
|
||||
else:
|
||||
|
||||
+3
-16
@@ -12,20 +12,6 @@ debug = os.environ.get('SD_LOAD_DEBUG', None) is not None
|
||||
loaded_te = None
|
||||
|
||||
|
||||
def install_gguf():
|
||||
# pip install git+https://github.com/junejae/transformers@feature/t5-gguf
|
||||
install('gguf', quiet=True)
|
||||
# https://github.com/ggerganov/llama.cpp/issues/9566
|
||||
import gguf
|
||||
scripts_dir = os.path.join(os.path.dirname(gguf.__file__), '..', 'scripts')
|
||||
if os.path.exists(scripts_dir):
|
||||
os.rename(scripts_dir, scripts_dir + '_gguf')
|
||||
# monkey patch transformers so they detect gguf pacakge correctly
|
||||
import importlib
|
||||
transformers.utils.import_utils._is_gguf_available = True # pylint: disable=protected-access
|
||||
transformers.utils.import_utils._gguf_version = importlib.metadata.version('gguf') # pylint: disable=protected-access
|
||||
|
||||
|
||||
def load_t5(name=None, cache_dir=None):
|
||||
global loaded_te # pylint: disable=global-statement
|
||||
if name is None:
|
||||
@@ -34,8 +20,9 @@ def load_t5(name=None, cache_dir=None):
|
||||
modelloader.hf_login()
|
||||
repo_id = 'stabilityai/stable-diffusion-3-medium-diffusers'
|
||||
fn = te_dict.get(name) if name in te_dict else None
|
||||
if fn is not None and 'gguf' in name.lower():
|
||||
install_gguf()
|
||||
if fn is not None and name.lower().endswith('gguf'):
|
||||
from modules import ggml
|
||||
ggml.install_gguf()
|
||||
with open(os.path.join('configs', 'flux', 'text_encoder_2', 'config.json'), encoding='utf8') as f:
|
||||
t5_config = transformers.T5Config(**json.load(f))
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(None, gguf_file=fn, config=t5_config, device_map="auto", cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
|
||||
@@ -96,6 +96,17 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
|
||||
guess = 'FLUX'
|
||||
if size > 11000 and size < 16000:
|
||||
warn(f'Model detected as FLUX UNET model, but attempting to load a base model: {op}={f} size={size} MB')
|
||||
# guess for diffusers
|
||||
index = os.path.join(f, 'model_index.json')
|
||||
if os.path.exists(index) and os.path.isfile(index):
|
||||
index = shared.readfile(index, silent=True)
|
||||
cls = index.get('_class_name', None)
|
||||
if cls is not None:
|
||||
pipeline = getattr(diffusers, cls)
|
||||
if 'Flux' in pipeline.__name__:
|
||||
guess = 'FLUX'
|
||||
if 'StableDiffusion3' in pipeline.__name__:
|
||||
guess = 'Stable Diffusion 3'
|
||||
# switch for specific variant
|
||||
if guess == 'Stable Diffusion' and 'inpaint' in f.lower():
|
||||
guess = 'Stable Diffusion Inpaint'
|
||||
|
||||
Reference in New Issue
Block a user