From 7fd86b156afbe3bb453958900d5b6bd8c074ab8a Mon Sep 17 00:00:00 2001 From: Maxed-Out-99 Date: Sun, 22 Feb 2026 08:21:31 -0800 Subject: [PATCH] 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. --- dequant.py | 4 +- loader.py | 124 +++++++++++++++++++++++++++++++++++++++++++++---- nodes.py | 36 ++++++++++++-- pyproject.toml | 2 +- 4 files changed, 149 insertions(+), 17 deletions(-) diff --git a/dequant.py b/dequant.py index 9f2100c..9689bc2 100644 --- a/dequant.py +++ b/dequant.py @@ -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 diff --git a/loader.py b/loader.py index 85008a5..ef9fbc5 100644 --- a/loader.py +++ b/loader.py @@ -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 = "" + + 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": diff --git a/nodes.py b/nodes.py index e34111c..1e8e283 100644 --- a/nodes.py +++ b/nodes.py @@ -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,) diff --git a/pyproject.toml b/pyproject.toml index 6326744..f1760e9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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)