File size: 7,163 Bytes
943b415 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 | import re
from dataclasses import dataclass
import torch
import comfy.model_management
import comfy.text_encoders.llama
from comfy import sd1_clip
from comfy.text_encoders.spiece_tokenizer import SPieceTokenizer
MAX_LENGTH = 256
@dataclass
class Gemma3_270M_Config:
vocab_size: int = 262144
hidden_size: int = 640
intermediate_size: int = 2048
num_hidden_layers: int = 18
num_attention_heads: int = 4
num_key_value_heads: int = 1
max_position_embeddings: int = 32768
rms_norm_eps: float = 1e-6
rope_theta = [1000000.0, 10000.0]
transformer_type: str = "gemma3"
head_dim = 256
rms_norm_add = True
mlp_activation = "gelu_pytorch_tanh"
qkv_bias = False
rope_dims = None
q_norm = "gemma3"
k_norm = "gemma3"
sliding_attention = [512, 512, 512, 512, 512, False]
rope_scale = None
final_norm: bool = True
lm_head: bool = False
stop_tokens = [1, 106]
class Gemma3_270M(comfy.text_encoders.llama.BaseLlama, torch.nn.Module):
def __init__(self, config_dict, dtype, device, operations):
super().__init__()
config = Gemma3_270M_Config(**config_dict)
self.num_layers = config.num_hidden_layers
self.model = comfy.text_encoders.llama.Llama2_(config, device=device, dtype=dtype, ops=operations)
self.dtype = dtype
def parse_prompt_emphasis(caption):
"""Strip "(text:weight)" groups; return the plain text and (start, end, weight) character spans into it."""
weight_pattern = re.compile(r"[+-]?(?:\d+(?:\.\d+)?|\.\d+)$")
spans = []
parts = []
cursor = 0
output_len = 0
idx = 0
while idx < len(caption):
if caption[idx] != "(":
idx += 1
continue
depth = 1
end = idx + 1
while end < len(caption) and depth > 0:
if caption[end] == "(":
depth += 1
elif caption[end] == ")":
depth -= 1
end += 1
if depth != 0:
idx += 1
continue
inner = caption[idx + 1:end - 1]
inner_depth = 0
colon_idx = -1
for inner_idx, char in enumerate(inner):
if char == "(":
inner_depth += 1
elif char == ")":
inner_depth -= 1
elif char == ":" and inner_depth == 0:
colon_idx = inner_idx
emphasized_text = inner[:colon_idx]
weight_text = inner[colon_idx + 1:].strip()
if colon_idx == -1 or not emphasized_text or not weight_pattern.fullmatch(weight_text):
idx += 1
continue
parts.append(caption[cursor:idx])
output_len += idx - cursor
parts.append(emphasized_text)
spans.append((output_len, output_len + len(emphasized_text), float(weight_text)))
output_len += len(emphasized_text)
cursor = end
idx = end
parts.append(caption[cursor:])
return "".join(parts), spans
class Gemma3_270MTokenizer(sd1_clip.SDTokenizer):
def __init__(self, embedding_directory=None, tokenizer_data={}):
tokenizer = tokenizer_data.get("spiece_model", None)
super().__init__(tokenizer, pad_with_end=False, embedding_size=640, embedding_key="gemma3_270m", tokenizer_class=SPieceTokenizer, has_end_token=False, pad_to_max_length=False, max_length=MAX_LENGTH, min_length=1, disable_weights=True, tokenizer_args={"add_bos": True, "add_eos": False}, tokenizer_data=tokenizer_data)
def tokenize_with_weights(self, text, return_word_ids=False, **kwargs):
# A token's weight is the product of the emphasis groups its characters overlap.
text, spans = parse_prompt_emphasis(text)
spm = self.tokenizer.tokenizer
if hasattr(spm, "EncodeAsOffsetMapping"):
encoded = spm.encode(text, add_bos=False, return_type="offset_mapping")
tokens = zip(encoded["ids"], encoded["offsets"])
else: # sentencepiece < 0.2.2
encoded = spm.encode(text, add_bos=False, out_type="immutable_proto")
tokens = [(piece.id, (piece.begin, piece.end)) for piece in encoded.pieces]
batch = [(self.start_token, 1.0, 0)]
for word_id, (token, (token_begin, token_end)) in enumerate(tokens, start=1):
weight = 1.0
for begin, end, span_weight in spans:
if token_begin < end and token_end > begin:
weight *= span_weight
batch.append((token, weight, word_id))
batch = batch[:self.max_length]
if not return_word_ids:
batch = [(token, weight) for token, weight, _ in batch]
return [batch]
def state_dict(self):
return {"spiece_model": self.tokenizer.serialize_model()}
class Nanosaur2Tokenizer(sd1_clip.SD1Tokenizer):
def __init__(self, embedding_directory=None, tokenizer_data={}):
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma3_270m", tokenizer=Gemma3_270MTokenizer)
class Gemma3_270MModel(sd1_clip.SDClipModel):
def __init__(self, device="cpu", layer="hidden", layer_idx=-2, dtype=None, model_options={}):
super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, dtype=dtype, special_tokens={"start": 2, "pad": 0}, layer_norm_hidden_state=True, model_class=Gemma3_270M, model_options=model_options)
def load_sd(self, sd):
# Gemma adds 1 to norm weights before casting, so preserve their checkpoint dtype.
for name, module in self.transformer.named_modules():
if isinstance(module, comfy.text_encoders.llama.RMSNorm):
module.to(dtype=sd[f"{name}.weight"].dtype)
comfy.model_management.archive_model_dtypes(module)
return super().load_sd(sd)
def encode_token_weights(self, token_weight_pairs):
# Emphasis weights go to the diffusion model, which scales attention to each text token by them.
out, pooled = self.encode([[token for token, _ in section] for section in token_weight_pairs])
device = comfy.model_management.intermediate_device()
token_weights = torch.tensor([[weight for _, weight in section] for section in token_weight_pairs], device=device)
return out.to(device), pooled, {"token_weights": token_weights}
class Nanosaur2TEModel(sd1_clip.SD1ClipModel):
def __init__(self, device="cpu", dtype=None, model_options={}):
super().__init__(device=device, dtype=dtype, name="gemma3_270m", clip_model=Gemma3_270MModel, model_options=model_options)
def te(dtype_llama=None, llama_quantization_metadata=None):
class Nanosaur2TEModel_(Nanosaur2TEModel):
def __init__(self, device="cpu", dtype=None, model_options={}):
if dtype_llama is not None:
dtype = dtype_llama
if llama_quantization_metadata is not None:
model_options = model_options.copy()
model_options["quantization_metadata"] = llama_quantization_metadata
super().__init__(device=device, dtype=dtype, model_options=model_options)
return Nanosaur2TEModel_
|