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_