levzalt's picture
Release step-1000 Nanosaur2 inpainting adapter and ComfyUI support
943b415 verified
Raw History Blame Contribute Delete
7.16 kB
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_