1 Commits
Author SHA1 Message Date
City bf8cb71be6 Initial attempt at gemma2 tokenizer 2025-11-08 19:34:56 +01:00
3 changed files with 84 additions and 244 deletions
+1 -54
View File
@@ -23,7 +23,7 @@ def dequantize_tensor(tensor, dtype=None, dequant_dtype=None):
return dequantize(tensor.data, qtype, oshape, dtype=dequant_dtype).to(dtype)
else:
# this is incredibly slow
tqdm.write(f"Falling back to numpy dequant for qtype: {getattr(qtype, 'name', repr(qtype))}")
tqdm.write(f"Falling back to numpy dequant for qtype: {qtype}")
new = gguf.quants.dequantize(tensor.cpu().numpy(), qtype)
return torch.from_numpy(new).to(tensor.device, dtype=dtype)
@@ -48,10 +48,6 @@ def to_uint32(x):
x = x.view(torch.uint8).to(torch.int32)
return (x[:, 0] | x[:, 1] << 8 | x[:, 2] << 16 | x[:, 3] << 24).unsqueeze(1)
def to_uint16(x):
x = x.view(torch.uint8).to(torch.int32)
return (x[:, 0] | x[:, 1] << 8).unsqueeze(1)
def split_block_dims(blocks, *args):
n_max = blocks.shape[1]
dims = list(args) + [n_max - sum(args)]
@@ -237,53 +233,6 @@ def dequantize_blocks_Q2_K(blocks, block_size, type_size, dtype=None):
return qs.reshape((n_blocks, -1))
# IQ quants
KVALUES = torch.tensor([-127, -104, -83, -65, -49, -35, -22, -10, 1, 13, 25, 38, 53, 69, 89, 113], dtype=torch.int8)
def dequantize_blocks_IQ4_NL(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, qs = split_block_dims(blocks, 2)
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.int64)
kvalues = KVALUES.to(qs.device).expand(*qs.shape[:-1], 16)
qs = torch.gather(kvalues, dim=-1, index=qs).reshape((n_blocks, -1))
del kvalues # should still be view, but just to be safe
return (d * qs)
def dequantize_blocks_IQ4_XS(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, scales_h, scales_l, qs = split_block_dims(blocks, 2, 2, QK_K // 64)
d = d.view(torch.float16).to(dtype)
scales_h = to_uint16(scales_h)
shift_a = torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape((1, 1, 2))
shift_b = torch.tensor([2 * i for i in range(QK_K // 32)], device=d.device, dtype=torch.uint8).reshape((1, -1, 1))
scales_l = scales_l.reshape((n_blocks, -1, 1)) >> shift_a.reshape((1, 1, 2))
scales_h = scales_h.reshape((n_blocks, -1, 1)) >> shift_b.reshape((1, -1, 1))
scales_l = scales_l.reshape((n_blocks, -1)) & 0x0F
scales_h = scales_h.reshape((n_blocks, -1)).to(torch.uint8) & 0x03
scales = (scales_l | (scales_h << 4)).to(torch.int8) - 32
dl = (d * scales.to(dtype)).reshape((n_blocks, -1, 1))
qs = qs.reshape((n_blocks, -1, 1, 16)) >> shift_a.reshape((1, 1, 2, 1))
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.int64)).reshape((n_blocks, -1, 32))
del kvalues # see IQ4_NL
del shift_a
del shift_b
return (dl * qs).reshape((n_blocks, -1))
dequantize_functions = {
gguf.GGMLQuantizationType.BF16: dequantize_blocks_BF16,
gguf.GGMLQuantizationType.Q8_0: dequantize_blocks_Q8_0,
@@ -296,6 +245,4 @@ dequantize_functions = {
gguf.GGMLQuantizationType.Q4_K: dequantize_blocks_Q4_K,
gguf.GGMLQuantizationType.Q3_K: dequantize_blocks_Q3_K,
gguf.GGMLQuantizationType.Q2_K: dequantize_blocks_Q2_K,
gguf.GGMLQuantizationType.IQ4_NL: dequantize_blocks_IQ4_NL,
gguf.GGMLQuantizationType.IQ4_XS: dequantize_blocks_IQ4_XS,
}
+81 -180
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", "gemma3"}
TXT_ARCH_LIST = {"t5", "t5encoder", "llama", "qwen2vl", "gemma2"}
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]].item())
return field_type(field.parts[field.data[-1]])
else:
raise TypeError(f"Unknown field type {field_type}")
@@ -48,26 +48,7 @@ def get_list_field(reader, field_name, field_type):
else:
raise TypeError(f"Unknown field type {field_type}")
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):
def gguf_sd_loader(path, handle_prefix="model.diffusion_model.", return_arch=False, is_text_model=False):
"""
Read state dict as fake tensors
"""
@@ -93,9 +74,9 @@ def gguf_sd_loader(path, handle_prefix="model.diffusion_model.", is_text_model=F
compat = None
arch_str = get_field(reader, "general.architecture", str)
type_str = get_field(reader, "general.type", str)
if arch_str in [None, "pig", "cow"]:
if arch_str in [None, "pig"]:
if is_text_model:
raise ValueError(f"This gguf file is incompatible with llama.cpp!\nConsider using safetensors or a compatible gguf file\n({path})")
raise ValueError(f"This text model is incompatible with llama.cpp!\nConsider using the safetensors version\n({path})")
compat = "sd.cpp" if arch_str is None else arch_str
# import here to avoid changes to convert.py breaking regular models
from .tools.convert import detect_arch
@@ -138,10 +119,6 @@ def gguf_sd_loader(path, handle_prefix="model.diffusion_model.", is_text_model=F
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
@@ -155,12 +132,9 @@ def gguf_sd_loader(path, handle_prefix="model.diffusion_model.", is_text_model=F
max_key = max(qsd.keys(), key=lambda k: qsd[k].numel())
state_dict[max_key].is_largest_weight = True
# extra info to return
extra = {
"arch_str": arch_str,
"metadata": get_gguf_metadata(reader)
}
return (state_dict, extra)
if return_arch:
return (state_dict, arch_str)
return state_dict
# for remapping llama.cpp -> original key names
T5_SD_MAP = {
@@ -183,9 +157,6 @@ T5_SD_MAP = {
LLAMA_SD_MAP = {
"blk.": "model.layers.",
"attn_norm": "input_layernorm",
"attn_q_norm.": "self_attn.q_norm.",
"attn_k_norm.": "self_attn.k_norm.",
"attn_v_norm.": "self_attn.v_norm.",
"attn_q": "self_attn.q_proj",
"attn_k": "self_attn.k_proj",
"attn_v": "self_attn.v_proj",
@@ -199,8 +170,8 @@ LLAMA_SD_MAP = {
"output.weight": "lm_head.weight",
}
GEMMA3_SD_MAP = LLAMA_SD_MAP.copy()
GEMMA3_SD_MAP.update({
GEMMA_SD_MAP = LLAMA_SD_MAP.copy()
GEMMA_SD_MAP.update({
"ffn_norm": "pre_feedforward_layernorm",
"post_ffw_norm": "post_feedforward_layernorm",
"post_attention_norm": "post_attention_layernorm",
@@ -239,28 +210,6 @@ 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)
@@ -297,7 +246,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:
@@ -345,26 +294,32 @@ def gguf_tokenizer_loader(path, temb_shape):
reader = gguf.GGUFReader(path)
if get_field(reader, "tokenizer.ggml.model", str) == "t5":
model_str = get_field(reader, "tokenizer.ggml.model", str)
if model_str == "t5":
if temb_shape == (256384, 4096): # probably UMT5
spm.trainer_spec.model_type == 1 # Unigram (do we have a T5 w/ BPE?)
spm.trainer_spec.max_sentence_length = 4096
else:
raise NotImplementedError("Unknown model, can't set tokenizer!")
elif model_str == "llama":
if temb_shape == (256000, 2304): # probably gemma
spm.trainer_spec.model_type == 2 # BPE
# TODO: something is missing, can't match 1:1
spm.trainer_spec.max_sentence_length = 0
spm.trainer_spec.max_sentencepiece_length = 16
spm.trainer_spec.split_digits = True
spm.trainer_spec.allow_whitespace_only_pieces = True
else:
raise NotImplementedError("Unknown model, can't set tokenizer!")
spm.normalizer_spec.add_dummy_prefix = get_field(reader, "tokenizer.ggml.add_space_prefix", bool)
spm.normalizer_spec.remove_extra_whitespaces = get_field(reader, "tokenizer.ggml.remove_extra_whitespaces", bool)
spm.normalizer_spec.add_dummy_prefix = get_field(reader, "tokenizer.ggml.add_space_prefix", bool) or False
spm.normalizer_spec.remove_extra_whitespaces = get_field(reader, "tokenizer.ggml.remove_extra_whitespaces", bool) or False
tokens = get_list_field(reader, "tokenizer.ggml.tokens", str)
scores = get_list_field(reader, "tokenizer.ggml.scores", float)
toktypes = get_list_field(reader, "tokenizer.ggml.token_type", int)
for idx, (token, score, toktype) in enumerate(zip(tokens, scores, toktypes)):
# # These aren't present in the original?
# if toktype == 5 and idx >= temb_shape[0]%1000):
# continue
for token, score, toktype in zip(tokens, scores, toktypes):
piece = spm.SentencePiece()
piece.piece = token
piece.score = score
@@ -373,134 +328,80 @@ def gguf_tokenizer_loader(path, temb_shape):
# unsure if any of these are correct
spm.trainer_spec.byte_fallback = True
spm.trainer_spec.vocab_size = len(tokens) # split off unused?
spm.trainer_spec.max_sentence_length = 4096
spm.trainer_spec.eos_id = get_field(reader, "tokenizer.ggml.eos_token_id", int)
spm.trainer_spec.pad_id = get_field(reader, "tokenizer.ggml.padding_token_id", int)
spm.trainer_spec.vocab_size = len(tokens)
# map special token IDs
tok_map = {
"bos_id": "tokenizer.ggml.bos_token_id",
"eos_id": "tokenizer.ggml.eos_token_id",
"pad_id": "tokenizer.ggml.padding_token_id",
"unk_id": "tokenizer.ggml.unknown_token_id",
}
for sp, gg in tok_map.items():
val = get_field(reader, gg, int)
if val is not None:
logging.debug(f"setting sp:{sp} to {val}")
setattr(spm.trainer_spec, sp, val)
# fix special token
if model_str == "llama" and hasattr(spm.trainer_spec, "unk_id"):
spm.pieces[spm.trainer_spec.unk_id].type = 2
for p in tok_map.keys():
if hasattr(spm.trainer_spec, p):
val = spm.pieces[getattr(spm.trainer_spec, p)].piece
setattr(spm.trainer_spec, p.replace("_id", "_piece"), val)
if temb_shape == (256000, 2304):
# for some reason the ggml tokenizer has these set to -1000 instead of 0...?
for p in spm.pieces:
if p.score == -1000 and p.type > 1:
p.score = 0.0
logging.info(f"Created tokenizer with vocab size of {len(spm.pieces)}")
del reader
return torch.ByteTensor(list(spm.SerializeToString()))
def gguf_tekken_tokenizer_loader(path, temb_shape):
# convert ggml (hf) tokenizer metadata to tekken/comfy data
logging.info("Attempting to recreate tekken tokenizer from GGUF file metadata...")
import json
import base64
from transformers.convert_slow_tokenizer import bytes_to_unicode
reader = gguf.GGUFReader(path)
model_str = get_field(reader, "tokenizer.ggml.model", str)
if model_str == "gpt2":
if temb_shape == (131072, 5120): # probably Mistral
data = {
"config": {"num_vocab_tokens": 150000, "default_vocab_size": 131072},
"vocab": [],
"special_tokens": [],
}
else:
raise NotImplementedError("Unknown model, can't set tokenizer!")
else:
raise NotImplementedError("Unknown model, can't set tokenizer!")
tokens = get_list_field(reader, "tokenizer.ggml.tokens", str)
toktypes = get_list_field(reader, "tokenizer.ggml.token_type", int)
decoder = {v: k for k, v in bytes_to_unicode().items()}
for idx, (token, toktype) in enumerate(zip(tokens, toktypes)):
if toktype == 3:
data["special_tokens"].append(
{'rank': idx, 'token_str': token, 'is_control': True}
)
else:
tok = bytes([decoder[char] for char in token])
data["vocab"].append({
"rank": len(data["vocab"]),
"token_bytes": base64.b64encode(tok).decode("ascii"),
"token_str": tok.decode("utf-8", errors="replace") # ?
})
logging.info(f"Created tekken tokenizer with vocab size of {len(data['vocab'])} (+{len(data['special_tokens'])})")
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 dequantize_temb(sd, temb_key="token_embd.weight"):
# TODO: dequantizing token embed here is janky but otherwise we OOM due to tensor being massive.
if temb_key in sd and is_quantized(sd[temb_key]):
logging.warning(f"Dequantizing {temb_key} to prevent runtime OOM.")
sd[temb_key] = dequantize_tensor(sd[temb_key], dtype=torch.float16)
return sd
def gguf_clip_loader(path):
sd, extra = gguf_sd_loader(path, is_text_model=True)
arch = extra.get("arch_str", None)
sd, arch = gguf_sd_loader(path, return_arch=True, is_text_model=True)
temb_key = "token_embd.weight"
if arch in {"t5", "t5encoder"}:
temb_key = "token_embd.weight"
if temb_key in sd and sd[temb_key].shape == (256384, 4096):
# non-standard Comfy-Org tokenizer
sd["spiece_model"] = gguf_tokenizer_loader(path, sd[temb_key].shape)
# TODO: dequantizing token embed here is janky but otherwise we OOM due to tensor being massive.
logging.warning(f"Dequantizing {temb_key} to prevent runtime OOM.")
sd[temb_key] = dequantize_tensor(sd[temb_key], dtype=torch.float16)
sd = dequantize_temb(sd, temb_key)
sd = sd_map_replace(sd, T5_SD_MAP)
elif arch in {"llama", "qwen2vl", "qwen3", "qwen3vl", "gemma3"}:
elif arch in {"llama", "qwen2vl"}:
# 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)
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)
sd = dequantize_temb(sd, temb_key)
sd = sd_map_replace(sd, LLAMA_SD_MAP)
if arch == "llama":
sd = llama_permute(sd, 32, 8) # L3 / Mistral
sd = llama_permute(sd, 32, 8) # L3
if arch == "qwen2vl":
vsd = gguf_mmproj_loader(path)
sd.update(vsd)
elif arch in {"gemma2"}:
if temb_key in sd:
# non-standard Comfy-Org tokenizer
sd["spiece_model"] = gguf_tokenizer_loader(path, sd[temb_key].shape)
# # TODO: for verifying tokenizer accuracy, remove this
# from safetensors.torch import load_file
# sd["spiece_model"] = load_file(r"models\clip\gemma_2_2b_fp16.safetensors")["spiece_model"]
sd = dequantize_temb(sd, temb_key)
sd = sd_map_replace(sd, GEMMA_SD_MAP)
# Reverse change from Gemma2Model.modify_tensors in convert_hf_to_gguf.py
for k,v in sd.items():
if k.endswith("norm.weight"):
if is_quantized(v):
v = dequantize_tensor(v, torch.float16)
sd[k] = v - 1.0
else:
pass
return sd
+2 -10
View File
@@ -1,7 +1,6 @@
# (c) City96 || Apache-2.0 (apache.org/licenses/LICENSE-2.0)
import torch
import logging
import inspect
import collections
import nodes
@@ -166,15 +165,9 @@ class UnetLoaderGGUF:
# init model
unet_path = folder_paths.get_full_path("unet", unet_name)
sd, extra = gguf_sd_loader(unet_path)
kwargs = {}
valid_params = inspect.signature(comfy.sd.load_diffusion_model_state_dict).parameters
if "metadata" in valid_params:
kwargs["metadata"] = extra.get("metadata", {})
sd = gguf_sd_loader(unet_path)
model = comfy.sd.load_diffusion_model_state_dict(
sd, model_options={"custom_operations": ops}, **kwargs,
sd, model_options={"custom_operations": ops}
)
if model is None:
logging.error("ERROR UNSUPPORTED UNET {}".format(unet_path))
@@ -326,4 +319,3 @@ NODE_CLASS_MAPPINGS = {
"QuadrupleCLIPLoaderGGUF": QuadrupleCLIPLoaderGGUF,
"UnetLoaderGGUFAdvanced": UnetLoaderGGUFAdvanced,
}