Download nanosaur2_support/text_encoder.py from levzalt/Nanosaur2-Inpaint-ControlNet: direct link, hf CLI and curl.
- Browser
- Download file 7.16 kB
-
https://huggingface.co/levzalt/Nanosaur2-Inpaint-ControlNet/resolve/main/nanosaur2_support/text_encoder.py
- Command line
-
hf download hf://levzalt/Nanosaur2-Inpaint-ControlNet/nanosaur2_support/text_encoder.py
-
curl -L -o text_encoder.py https://huggingface.co/levzalt/Nanosaur2-Inpaint-ControlNet/resolve/main/nanosaur2_support/text_encoder.py
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 | |
| 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_ | |