Compare commits
1
Commits
main
..
gemma_tenc
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bf8cb71be6 |
+1
-54
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user