Add Gemma3 GGUF support and mmap/memory fixes
Multiple improvements and fixes across GGUF loading and model patching: - dequant.py: Fix index dtype for IQ4 dequant gathers by casting indices to int64 to avoid dtype issues. - loader.py: - Extend TXT_ARCH_LIST with gemma3. - Fix get_field to extract scalar values via .item(). - Add get_gguf_metadata to collect simple GGUF metadata (string/int/float/bool). - Change gguf_sd_loader to return (state_dict, extra) where extra includes arch_str and metadata. - Dequantize 1D BF16 tensors to float32 to avoid incorrect quantization for bias/1D params. - Add GEMMA3_SD_MAP and gemma3_norm_corrections to reverse a Gemma3-specific norm offset (apply -1.0 correction and dequantize if needed). - Add gguf_gemma3_tokenizer_loader to reconstruct a SentencePiece tokenizer from GGUF metadata. - Update gguf_clip_loader to handle gemma3: map keys, apply norm corrections, and produce tokenizer data. - Ensure gguf_mmproj_loader and other callers unpack the new gguf_sd_loader return value. - nodes.py: - Improve GGUFModelPatcher mmap handling: track named modules to unmap, add pin_weight_to_device to safely move modules when releasing mmap, clear tracking after release. - Pass GGUF metadata into comfy.sd.load_diffusion_model_state_dict when supported, and add error checks when loading fails. These changes add Gemma3 model/tokenizer support, fix dtype and BF16 edge cases, and improve low-memory mmap/unmap handling for safer weight pinning and loading.
This commit is contained in:
+2
-2
@@ -247,7 +247,7 @@ def dequantize_blocks_IQ4_NL(blocks, block_size, type_size, dtype=None):
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
|
||||
qs = qs.reshape((n_blocks, -1, 1, block_size//2)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape((1, 1, 2, 1))
|
||||
qs = (qs & 0x0F).reshape((n_blocks, -1, 1)).to(torch.int32)
|
||||
qs = (qs & 0x0F).reshape((n_blocks, -1, 1)).to(torch.int64)
|
||||
|
||||
kvalues = KVALUES.to(qs.device).expand(*qs.shape[:-1], 16)
|
||||
qs = torch.gather(kvalues, dim=-1, index=qs).reshape((n_blocks, -1))
|
||||
@@ -277,7 +277,7 @@ def dequantize_blocks_IQ4_XS(blocks, block_size, type_size, dtype=None):
|
||||
qs = qs.reshape((n_blocks, -1, 32, 1)) & 0x0F
|
||||
|
||||
kvalues = KVALUES.to(qs.device).expand(*qs.shape[:-1], 16)
|
||||
qs = torch.gather(kvalues, dim=-1, index=qs.to(torch.int32)).reshape((n_blocks, -1, 32))
|
||||
qs = torch.gather(kvalues, dim=-1, index=qs.to(torch.int64)).reshape((n_blocks, -1, 32))
|
||||
del kvalues # see IQ4_NL
|
||||
del shift_a
|
||||
del shift_b
|
||||
|
||||
@@ -10,7 +10,7 @@ from .ops import GGMLTensor
|
||||
from .dequant import is_quantized, dequantize_tensor
|
||||
|
||||
IMG_ARCH_LIST = {"flux", "sd1", "sdxl", "sd3", "aura", "hidream", "cosmos", "ltxv", "hyvid", "wan", "lumina2", "qwen_image"}
|
||||
TXT_ARCH_LIST = {"t5", "t5encoder", "llama", "qwen2vl", "qwen3", "qwen3vl"}
|
||||
TXT_ARCH_LIST = {"t5", "t5encoder", "llama", "qwen2vl", "qwen3", "qwen3vl", "gemma3"}
|
||||
VIS_TYPE_LIST = {"clip-vision", "mmproj"}
|
||||
|
||||
def get_orig_shape(reader, tensor_name):
|
||||
@@ -33,7 +33,7 @@ def get_field(reader, field_name, field_type):
|
||||
raise TypeError(f"Bad type for GGUF {field_name} key: expected string, got {field.types!r}")
|
||||
return str(field.parts[field.data[-1]], encoding="utf-8")
|
||||
elif field_type in [int, float, bool]:
|
||||
return field_type(field.parts[field.data[-1]])
|
||||
return field_type(field.parts[field.data[-1]].item())
|
||||
else:
|
||||
raise TypeError(f"Unknown field type {field_type}")
|
||||
|
||||
@@ -48,7 +48,26 @@ def get_list_field(reader, field_name, field_type):
|
||||
else:
|
||||
raise TypeError(f"Unknown field type {field_type}")
|
||||
|
||||
def gguf_sd_loader(path, handle_prefix="model.diffusion_model.", return_arch=False, is_text_model=False):
|
||||
def get_gguf_metadata(reader):
|
||||
"""Extract all simple metadata fields like safetensors"""
|
||||
metadata = {}
|
||||
for field_name in reader.fields:
|
||||
try:
|
||||
field = reader.get_field(field_name)
|
||||
if len(field.types) == 1: # Simple scalar fields only
|
||||
if field.types[0] == gguf.GGUFValueType.STRING:
|
||||
metadata[field_name] = str(field.parts[field.data[-1]], "utf-8")
|
||||
elif field.types[0] == gguf.GGUFValueType.INT32:
|
||||
metadata[field_name] = int(field.parts[field.data[-1]])
|
||||
elif field.types[0] == gguf.GGUFValueType.F32:
|
||||
metadata[field_name] = float(field.parts[field.data[-1]])
|
||||
elif field.types[0] == gguf.GGUFValueType.BOOL:
|
||||
metadata[field_name] = bool(field.parts[field.data[-1]])
|
||||
except:
|
||||
continue
|
||||
return metadata
|
||||
|
||||
def gguf_sd_loader(path, handle_prefix="model.diffusion_model.", is_text_model=False):
|
||||
"""
|
||||
Read state dict as fake tensors
|
||||
"""
|
||||
@@ -119,6 +138,10 @@ def gguf_sd_loader(path, handle_prefix="model.diffusion_model.", return_arch=Fal
|
||||
torch_tensor = torch_tensor.view(*shape)
|
||||
state_dict[sd_key] = GGMLTensor(torch_tensor, tensor_type=tensor.tensor_type, tensor_shape=shape)
|
||||
|
||||
# 1D tensors shouldn't be quantized, this is a fix for BF16
|
||||
if len(shape) <= 1 and tensor.tensor_type == gguf.GGMLQuantizationType.BF16:
|
||||
state_dict[sd_key] = dequantize_tensor(state_dict[sd_key], dtype=torch.float32)
|
||||
|
||||
# keep track of loaded tensor types
|
||||
tensor_type_str = getattr(tensor.tensor_type, "name", repr(tensor.tensor_type))
|
||||
qtype_dict[tensor_type_str] = qtype_dict.get(tensor_type_str, 0) + 1
|
||||
@@ -132,9 +155,12 @@ def gguf_sd_loader(path, handle_prefix="model.diffusion_model.", return_arch=Fal
|
||||
max_key = max(qsd.keys(), key=lambda k: qsd[k].numel())
|
||||
state_dict[max_key].is_largest_weight = True
|
||||
|
||||
if return_arch:
|
||||
return (state_dict, arch_str)
|
||||
return state_dict
|
||||
# extra info to return
|
||||
extra = {
|
||||
"arch_str": arch_str,
|
||||
"metadata": get_gguf_metadata(reader)
|
||||
}
|
||||
return (state_dict, extra)
|
||||
|
||||
# for remapping llama.cpp -> original key names
|
||||
T5_SD_MAP = {
|
||||
@@ -173,6 +199,13 @@ LLAMA_SD_MAP = {
|
||||
"output.weight": "lm_head.weight",
|
||||
}
|
||||
|
||||
GEMMA3_SD_MAP = LLAMA_SD_MAP.copy()
|
||||
GEMMA3_SD_MAP.update({
|
||||
"ffn_norm": "pre_feedforward_layernorm",
|
||||
"post_ffw_norm": "post_feedforward_layernorm",
|
||||
"post_attention_norm": "post_attention_layernorm",
|
||||
})
|
||||
|
||||
CLIP_VISION_SD_MAP = {
|
||||
"mm.": "visual.merger.mlp.",
|
||||
"v.post_ln.": "visual.merger.ln_q.",
|
||||
@@ -206,6 +239,28 @@ def llama_permute(raw_sd, n_head, n_head_kv):
|
||||
sd[k] = v
|
||||
return sd
|
||||
|
||||
def gemma3_norm_corrections(sd):
|
||||
# Reverse change from Gemma3Model modify_tensors in llama.cpp convert script
|
||||
norm_patterns = [
|
||||
"input_layernorm.weight",
|
||||
"post_attention_layernorm.weight",
|
||||
"pre_feedforward_layernorm.weight",
|
||||
"post_feedforward_layernorm.weight",
|
||||
"self_attn.q_norm.weight",
|
||||
"self_attn.k_norm.weight",
|
||||
"model.norm.weight"
|
||||
]
|
||||
corrected = 0
|
||||
for key in list(sd.keys()):
|
||||
if any(p in key for p in norm_patterns):
|
||||
if is_quantized(sd[key]):
|
||||
sd[key] = dequantize_tensor(sd[key], dtype=torch.float32) - 1.0
|
||||
else:
|
||||
sd[key] = sd[key].float() - 1.0
|
||||
corrected += 1
|
||||
#logging.info(f"Gemma3: Applied -1 norm correction to {corrected} tensors")
|
||||
return sd
|
||||
|
||||
def strip_quant_suffix(name):
|
||||
pattern = r"[-_]?(?:ud-)?i?q[0-9]_[a-z0-9_\-]{1,8}$"
|
||||
match = re.search(pattern, name, re.IGNORECASE)
|
||||
@@ -242,7 +297,7 @@ def gguf_mmproj_loader(path):
|
||||
|
||||
logging.info(f"Using mmproj '{target[0]}' for text encoder '{tenc_fname}'.")
|
||||
target = os.path.join(root, target[0])
|
||||
vsd = gguf_sd_loader(target, is_text_model=True)
|
||||
vsd, _ = gguf_sd_loader(target, is_text_model=True)
|
||||
|
||||
# concat 4D to 5D
|
||||
if "v.patch_embd.weight.1" in vsd:
|
||||
@@ -370,8 +425,51 @@ def gguf_tekken_tokenizer_loader(path, temb_shape):
|
||||
del reader
|
||||
return torch.ByteTensor(list(json.dumps(data).encode('utf-8')))
|
||||
|
||||
def gguf_gemma3_tokenizer_loader(path):
|
||||
#TODO: merge into gguf_tokenizer_loader
|
||||
logging.info("Attempting to recreate sentencepiece tokenizer from GGUF file metadata...")
|
||||
try:
|
||||
from sentencepiece import sentencepiece_model_pb2 as model
|
||||
except ImportError:
|
||||
raise ImportError("Please install sentencepiece and protobuf.\npip install sentencepiece protobuf")
|
||||
spm = model.ModelProto()
|
||||
reader = gguf.GGUFReader(path)
|
||||
|
||||
spm.normalizer_spec.name = "identity"
|
||||
spm.normalizer_spec.add_dummy_prefix = False
|
||||
spm.trainer_spec.model_type = 2
|
||||
spm.trainer_spec.input_format = "tsv"
|
||||
spm.trainer_spec.byte_fallback = True
|
||||
spm.trainer_spec.max_sentence_length = 4192
|
||||
spm.trainer_spec.bos_piece = "<bos>"
|
||||
|
||||
tokens = get_list_field(reader, "tokenizer.ggml.tokens", str)
|
||||
scores = get_list_field(reader, "tokenizer.ggml.scores", float)
|
||||
toktype = get_list_field(reader, "tokenizer.ggml.token_type", int)
|
||||
|
||||
if not tokens or not scores or not toktype:
|
||||
raise ValueError("Missing tokenizer metadata")
|
||||
|
||||
for idx in range(len(tokens)):
|
||||
piece = spm.SentencePiece()
|
||||
piece.piece = tokens[idx]
|
||||
if idx == 3: # UNK position
|
||||
piece.type = 2 # UNK Token
|
||||
piece.score = 0.0 # UNK Score
|
||||
else:
|
||||
piece.type = toktype[idx]
|
||||
piece.score = scores[idx]
|
||||
spm.pieces.append(piece)
|
||||
|
||||
spm.trainer_spec.vocab_size = len(spm.pieces)
|
||||
logging.info(f"Created tokenizer with vocab size of {len(spm.pieces)}")
|
||||
|
||||
del reader
|
||||
return torch.ByteTensor(list(spm.SerializeToString()))
|
||||
|
||||
def gguf_clip_loader(path):
|
||||
sd, arch = gguf_sd_loader(path, return_arch=True, is_text_model=True)
|
||||
sd, extra = gguf_sd_loader(path, is_text_model=True)
|
||||
arch = extra.get("arch_str", None)
|
||||
if arch in {"t5", "t5encoder"}:
|
||||
temb_key = "token_embd.weight"
|
||||
if temb_key in sd and sd[temb_key].shape == (256384, 4096):
|
||||
@@ -381,17 +479,23 @@ def gguf_clip_loader(path):
|
||||
logging.warning(f"Dequantizing {temb_key} to prevent runtime OOM.")
|
||||
sd[temb_key] = dequantize_tensor(sd[temb_key], dtype=torch.float16)
|
||||
sd = sd_map_replace(sd, T5_SD_MAP)
|
||||
elif arch in {"llama", "qwen2vl", "qwen3", "qwen3vl"}:
|
||||
elif arch in {"llama", "qwen2vl", "qwen3", "qwen3vl", "gemma3"}:
|
||||
# TODO: pass model_options["vocab_size"] to loader somehow
|
||||
temb_key = "token_embd.weight"
|
||||
if temb_key in sd and sd[temb_key].shape[0] >= (64 * 1024):
|
||||
if arch == "llama" and sd[temb_key].shape == (131072, 5120):
|
||||
# non-standard Comfy-Org tokenizer
|
||||
sd["tekken_model"] = gguf_tekken_tokenizer_loader(path, sd[temb_key].shape)
|
||||
elif arch == "gemma3":
|
||||
sd["spiece_model"] = gguf_gemma3_tokenizer_loader(path)
|
||||
# See note above for T5.
|
||||
logging.warning(f"Dequantizing {temb_key} to prevent runtime OOM.")
|
||||
sd[temb_key] = dequantize_tensor(sd[temb_key], dtype=torch.float16)
|
||||
sd = sd_map_replace(sd, LLAMA_SD_MAP)
|
||||
if arch == "gemma3":
|
||||
sd = sd_map_replace(sd, GEMMA3_SD_MAP)
|
||||
sd = gemma3_norm_corrections(sd)
|
||||
else:
|
||||
sd = sd_map_replace(sd, LLAMA_SD_MAP)
|
||||
if arch == "llama":
|
||||
sd = llama_permute(sd, 32, 8) # L3 / Mistral
|
||||
if arch == "qwen2vl":
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
# Modified by Maxed-Out-99
|
||||
import logging
|
||||
import collections
|
||||
import inspect
|
||||
|
||||
import comfy.sd
|
||||
import comfy.lora
|
||||
@@ -78,8 +79,21 @@ class GGUFModelPatcher(comfy.model_patcher.ModelPatcher):
|
||||
# TODO: Find another way to not unload after patches
|
||||
return super().unpatch_model(device_to=device_to, unpatch_weights=unpatch_weights)
|
||||
|
||||
def pin_weight_to_device(self, key):
|
||||
op_key = key.rsplit('.', 1)[0]
|
||||
if not self.mmap_released and op_key in self.named_modules_to_munmap:
|
||||
# TODO: possible to OOM, find better way to detach
|
||||
self.named_modules_to_munmap[op_key].to(self.load_device).to(self.offload_device)
|
||||
del self.named_modules_to_munmap[op_key]
|
||||
super().pin_weight_to_device(key)
|
||||
|
||||
mmap_released = False
|
||||
named_modules_to_munmap = {}
|
||||
|
||||
def load(self, *args, force_patch_weights=False, **kwargs):
|
||||
if not self.mmap_released:
|
||||
self.named_modules_to_munmap = dict(self.model.named_modules())
|
||||
|
||||
# always call `patch_weight_to_device` even for lowvram
|
||||
super().load(*args, force_patch_weights=True, **kwargs)
|
||||
|
||||
@@ -87,7 +101,7 @@ class GGUFModelPatcher(comfy.model_patcher.ModelPatcher):
|
||||
if not self.mmap_released:
|
||||
linked = []
|
||||
if kwargs.get("lowvram_model_memory", 0) > 0:
|
||||
for n, m in self.model.named_modules():
|
||||
for n, m in self.named_modules_to_munmap.items():
|
||||
if hasattr(m, "weight"):
|
||||
device = getattr(m.weight, "device", None)
|
||||
if device == self.offload_device:
|
||||
@@ -104,6 +118,7 @@ class GGUFModelPatcher(comfy.model_patcher.ModelPatcher):
|
||||
# TODO: possible to OOM, find better way to detach
|
||||
m.to(self.load_device).to(self.offload_device)
|
||||
self.mmap_released = True
|
||||
self.named_modules_to_munmap = {}
|
||||
|
||||
def clone(self, *args, **kwargs):
|
||||
src_cls = self.__class__
|
||||
@@ -113,6 +128,7 @@ class GGUFModelPatcher(comfy.model_patcher.ModelPatcher):
|
||||
self.__class__ = src_cls
|
||||
# GGUF specific clone values below
|
||||
n.patch_on_device = getattr(self, "patch_on_device", False)
|
||||
n.mmap_released = getattr(self, "mmap_released", False)
|
||||
if src_cls != GGUFModelPatcher:
|
||||
n.size = 0 # force recalc
|
||||
return n
|
||||
@@ -149,13 +165,25 @@ class UNETLoaderUnified:
|
||||
ops = GGMLOps()
|
||||
|
||||
# Load state dict from GGUF
|
||||
sd = gguf_sd_loader(unet_path)
|
||||
if sd is None:
|
||||
loaded = gguf_sd_loader(unet_path)
|
||||
if loaded is None:
|
||||
raise RuntimeError(f"Failed to load GGUF model: {unet_path}")
|
||||
|
||||
if isinstance(loaded, tuple) and len(loaded) == 2:
|
||||
sd, extra = loaded
|
||||
else:
|
||||
sd, extra = loaded, {}
|
||||
|
||||
kwargs = {}
|
||||
valid_params = inspect.signature(comfy.sd.load_diffusion_model_state_dict).parameters
|
||||
if "metadata" in valid_params:
|
||||
kwargs["metadata"] = extra.get("metadata", {})
|
||||
|
||||
model = comfy.sd.load_diffusion_model_state_dict(
|
||||
sd, model_options={"custom_operations": ops}
|
||||
sd, model_options={"custom_operations": ops}, **kwargs
|
||||
)
|
||||
if model is None:
|
||||
raise RuntimeError(f"Could not detect GGUF model type for: {unet_path}")
|
||||
model = GGUFModelPatcher.clone(model)
|
||||
return (model,)
|
||||
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "smartmodelloaders-mxd"
|
||||
description = "Smart, unified model loaders for ComfyUI that support both standard .safetensors and quantized .gguf formats — no switching nodes required. Includes flexible UNET and CLIP loaders that work across models like SDXL, SD3, Flux, and more."
|
||||
version = "1.0.4"
|
||||
version = "1.0.5"
|
||||
license = {file = "LICENSE"}
|
||||
# classifiers = [
|
||||
# # For OS-independent nodes (works on all operating systems)
|
||||
|
||||
Reference in New Issue
Block a user