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_