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:
Maxed-Out-99
2026-02-22 08:21:31 -08:00
parent b7f561c389
commit 7fd86b156a
4 changed files with 149 additions and 17 deletions
+2 -2
View File
@@ -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
+114 -10
View File
@@ -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":
+32 -4
View File
@@ -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
View File
@@ -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)