Spaces:
Running on Zero
Running on Zero
Vendor ComfyUI + custom nodes, add Gradio app with ZeroGPU + model auto-download (part 3)
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +1 -0
- comfy/text_encoders/sa_t5.py +22 -0
- comfy/text_encoders/sam3_clip.py +97 -0
- comfy/text_encoders/sd2_clip.py +23 -0
- comfy/text_encoders/sd2_clip_config.json +23 -0
- comfy/text_encoders/sd3_clip.py +167 -0
- comfy/text_encoders/sensenova.py +149 -0
- comfy/text_encoders/spiece_tokenizer.py +59 -0
- comfy/text_encoders/t5.py +249 -0
- comfy/text_encoders/t5_config_base.json +22 -0
- comfy/text_encoders/t5_config_xxl.json +22 -0
- comfy/text_encoders/t5_old_config_xxl.json +22 -0
- comfy/text_encoders/t5_pile_config_xl.json +22 -0
- comfy/text_encoders/t5_pile_tokenizer/tokenizer.model +3 -0
- comfy/text_encoders/t5_tokenizer/special_tokens_map.json +125 -0
- comfy/text_encoders/t5_tokenizer/tokenizer.json +0 -0
- comfy/text_encoders/t5_tokenizer/tokenizer_config.json +939 -0
- comfy/text_encoders/umt5_config_base.json +22 -0
- comfy/text_encoders/umt5_config_xxl.json +22 -0
- comfy/text_encoders/wan.py +37 -0
- comfy/text_encoders/z_image.py +46 -0
- comfy/utils.py +1535 -0
- comfy/weight_adapter/__init__.py +42 -0
- comfy/weight_adapter/base.py +396 -0
- comfy/weight_adapter/boft.py +218 -0
- comfy/weight_adapter/bypass.py +441 -0
- comfy/weight_adapter/glora.py +290 -0
- comfy/weight_adapter/loha.py +378 -0
- comfy/weight_adapter/lokr.py +481 -0
- comfy/weight_adapter/lora.py +368 -0
- comfy/weight_adapter/oft.py +327 -0
- comfy_api/feature_flags.py +166 -0
- comfy_api/generate_api_stubs.py +86 -0
- comfy_api/input/__init__.py +26 -0
- comfy_api/input/basic_types.py +14 -0
- comfy_api/input/video_types.py +6 -0
- comfy_api/input_impl/__init__.py +7 -0
- comfy_api/input_impl/video_types.py +2 -0
- comfy_api/internal/__init__.py +150 -0
- comfy_api/internal/api_registry.py +39 -0
- comfy_api/internal/async_to_sync.py +1002 -0
- comfy_api/internal/singleton.py +33 -0
- comfy_api/latest/__init__.py +177 -0
- comfy_api/latest/_caching.py +42 -0
- comfy_api/latest/_input/__init__.py +17 -0
- comfy_api/latest/_input/basic_types.py +42 -0
- comfy_api/latest/_input/curve_types.py +219 -0
- comfy_api/latest/_input/range_types.py +70 -0
- comfy_api/latest/_input/video_types.py +198 -0
- comfy_api/latest/_input_impl/__init__.py +7 -0
.gitattributes
CHANGED
|
@@ -1,3 +1,4 @@
|
|
| 1 |
/web/assets/** linguist-generated
|
| 2 |
/web/** linguist-vendored
|
| 3 |
comfy_api_nodes/apis/__init__.py linguist-generated
|
|
|
|
|
|
| 1 |
/web/assets/** linguist-generated
|
| 2 |
/web/** linguist-vendored
|
| 3 |
comfy_api_nodes/apis/__init__.py linguist-generated
|
| 4 |
+
comfy/text_encoders/t5_pile_tokenizer/tokenizer.model filter=lfs diff=lfs merge=lfs -text
|
comfy/text_encoders/sa_t5.py
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from comfy import sd1_clip
|
| 2 |
+
from transformers import T5TokenizerFast
|
| 3 |
+
import comfy.text_encoders.t5
|
| 4 |
+
import os
|
| 5 |
+
|
| 6 |
+
class T5BaseModel(sd1_clip.SDClipModel):
|
| 7 |
+
def __init__(self, device="cpu", layer="last", layer_idx=None, dtype=None, model_options={}):
|
| 8 |
+
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "t5_config_base.json")
|
| 9 |
+
super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, model_options=model_options, special_tokens={"end": 1, "pad": 0}, model_class=comfy.text_encoders.t5.T5, enable_attention_masks=True, zero_out_masked=True)
|
| 10 |
+
|
| 11 |
+
class T5BaseTokenizer(sd1_clip.SDTokenizer):
|
| 12 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 13 |
+
tokenizer_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "t5_tokenizer")
|
| 14 |
+
super().__init__(tokenizer_path, pad_with_end=False, embedding_size=768, embedding_key='t5base', tokenizer_class=T5TokenizerFast, has_start_token=False, pad_to_max_length=False, max_length=99999999, min_length=128, tokenizer_data=tokenizer_data)
|
| 15 |
+
|
| 16 |
+
class SAT5Tokenizer(sd1_clip.SD1Tokenizer):
|
| 17 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 18 |
+
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, clip_name="t5base", tokenizer=T5BaseTokenizer)
|
| 19 |
+
|
| 20 |
+
class SAT5Model(sd1_clip.SD1ClipModel):
|
| 21 |
+
def __init__(self, device="cpu", dtype=None, model_options={}, **kwargs):
|
| 22 |
+
super().__init__(device=device, dtype=dtype, model_options=model_options, name="t5base", clip_model=T5BaseModel, **kwargs)
|
comfy/text_encoders/sam3_clip.py
ADDED
|
@@ -0,0 +1,97 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import re
|
| 2 |
+
from comfy import sd1_clip
|
| 3 |
+
|
| 4 |
+
SAM3_CLIP_CONFIG = {
|
| 5 |
+
"architectures": ["CLIPTextModel"],
|
| 6 |
+
"hidden_act": "quick_gelu",
|
| 7 |
+
"hidden_size": 1024,
|
| 8 |
+
"intermediate_size": 4096,
|
| 9 |
+
"num_attention_heads": 16,
|
| 10 |
+
"num_hidden_layers": 24,
|
| 11 |
+
"max_position_embeddings": 32,
|
| 12 |
+
"projection_dim": 512,
|
| 13 |
+
"vocab_size": 49408,
|
| 14 |
+
"layer_norm_eps": 1e-5,
|
| 15 |
+
"eos_token_id": 49407,
|
| 16 |
+
}
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class SAM3ClipModel(sd1_clip.SDClipModel):
|
| 20 |
+
def __init__(self, device="cpu", dtype=None, model_options={}):
|
| 21 |
+
super().__init__(device=device, dtype=dtype, max_length=32, layer="last", textmodel_json_config=SAM3_CLIP_CONFIG, special_tokens={"start": 49406, "end": 49407, "pad": 0}, return_projected_pooled=False, return_attention_masks=True, enable_attention_masks=True, model_options=model_options)
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class SAM3Tokenizer(sd1_clip.SDTokenizer):
|
| 25 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 26 |
+
super().__init__(max_length=32, pad_with_end=False, pad_token=0, embedding_directory=embedding_directory, embedding_size=1024, embedding_key="sam3_clip", tokenizer_data=tokenizer_data)
|
| 27 |
+
self.disable_weights = True
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def _parse_prompts(text):
|
| 31 |
+
"""Split comma-separated prompts with optional :N max detections per category"""
|
| 32 |
+
text = text.replace("(", "").replace(")", "")
|
| 33 |
+
parts = [p.strip() for p in text.split(",") if p.strip()]
|
| 34 |
+
result = []
|
| 35 |
+
for part in parts:
|
| 36 |
+
m = re.match(r'^(.+?)\s*:\s*([\d.]+)\s*$', part)
|
| 37 |
+
if m:
|
| 38 |
+
text_part = m.group(1).strip()
|
| 39 |
+
val = m.group(2)
|
| 40 |
+
max_det = max(1, round(float(val)))
|
| 41 |
+
result.append((text_part, max_det))
|
| 42 |
+
else:
|
| 43 |
+
result.append((part, 1))
|
| 44 |
+
return result
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
class SAM3TokenizerWrapper(sd1_clip.SD1Tokenizer):
|
| 48 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 49 |
+
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, clip_name="l", tokenizer=SAM3Tokenizer, name="sam3_clip")
|
| 50 |
+
|
| 51 |
+
def tokenize_with_weights(self, text: str, return_word_ids=False, **kwargs):
|
| 52 |
+
parsed = _parse_prompts(text)
|
| 53 |
+
if len(parsed) <= 1 and (not parsed or parsed[0][1] == 1):
|
| 54 |
+
return super().tokenize_with_weights(text, return_word_ids, **kwargs)
|
| 55 |
+
# Tokenize each prompt part separately, store per-part batches and metadata
|
| 56 |
+
inner = getattr(self, self.clip)
|
| 57 |
+
per_prompt = []
|
| 58 |
+
for prompt_text, max_det in parsed:
|
| 59 |
+
batches = inner.tokenize_with_weights(prompt_text, return_word_ids, **kwargs)
|
| 60 |
+
per_prompt.append((batches, max_det))
|
| 61 |
+
# Main output uses first prompt's tokens (for compatibility)
|
| 62 |
+
out = {self.clip_name: per_prompt[0][0], "sam3_per_prompt": per_prompt}
|
| 63 |
+
return out
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
class SAM3ClipModelWrapper(sd1_clip.SD1ClipModel):
|
| 67 |
+
def __init__(self, device="cpu", dtype=None, model_options={}, **kwargs):
|
| 68 |
+
super().__init__(device=device, dtype=dtype, model_options=model_options, clip_name="l", clip_model=SAM3ClipModel, name="sam3_clip")
|
| 69 |
+
|
| 70 |
+
def encode_token_weights(self, token_weight_pairs):
|
| 71 |
+
per_prompt = token_weight_pairs.pop("sam3_per_prompt", None)
|
| 72 |
+
if per_prompt is None:
|
| 73 |
+
return super().encode_token_weights(token_weight_pairs)
|
| 74 |
+
|
| 75 |
+
# Encode each prompt separately, pack into extra dict
|
| 76 |
+
inner = getattr(self, self.clip)
|
| 77 |
+
multi_cond = []
|
| 78 |
+
first_pooled = None
|
| 79 |
+
for batches, max_det in per_prompt:
|
| 80 |
+
out = inner.encode_token_weights(batches)
|
| 81 |
+
cond, pooled = out[0], out[1]
|
| 82 |
+
extra = out[2] if len(out) > 2 else {}
|
| 83 |
+
if first_pooled is None:
|
| 84 |
+
first_pooled = pooled
|
| 85 |
+
multi_cond.append({
|
| 86 |
+
"cond": cond,
|
| 87 |
+
"attention_mask": extra.get("attention_mask"),
|
| 88 |
+
"max_detections": max_det,
|
| 89 |
+
})
|
| 90 |
+
|
| 91 |
+
# Return first prompt as main (for non-SAM3 consumers), all prompts in metadata
|
| 92 |
+
main = multi_cond[0]
|
| 93 |
+
main_extra = {}
|
| 94 |
+
if main["attention_mask"] is not None:
|
| 95 |
+
main_extra["attention_mask"] = main["attention_mask"]
|
| 96 |
+
main_extra["sam3_multi_cond"] = multi_cond
|
| 97 |
+
return (main["cond"], first_pooled, main_extra)
|
comfy/text_encoders/sd2_clip.py
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from comfy import sd1_clip
|
| 2 |
+
import os
|
| 3 |
+
|
| 4 |
+
class SD2ClipHModel(sd1_clip.SDClipModel):
|
| 5 |
+
def __init__(self, arch="ViT-H-14", device="cpu", max_length=77, freeze=True, layer="penultimate", layer_idx=None, dtype=None, model_options={}):
|
| 6 |
+
if layer == "penultimate":
|
| 7 |
+
layer="hidden"
|
| 8 |
+
layer_idx=-2
|
| 9 |
+
|
| 10 |
+
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "sd2_clip_config.json")
|
| 11 |
+
super().__init__(device=device, freeze=freeze, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, special_tokens={"start": 49406, "end": 49407, "pad": 0}, return_projected_pooled=True, model_options=model_options)
|
| 12 |
+
|
| 13 |
+
class SD2ClipHTokenizer(sd1_clip.SDTokenizer):
|
| 14 |
+
def __init__(self, tokenizer_path=None, embedding_directory=None, tokenizer_data={}):
|
| 15 |
+
super().__init__(tokenizer_path, pad_with_end=False, embedding_directory=embedding_directory, embedding_size=1024, embedding_key='clip_h', tokenizer_data=tokenizer_data)
|
| 16 |
+
|
| 17 |
+
class SD2Tokenizer(sd1_clip.SD1Tokenizer):
|
| 18 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 19 |
+
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, clip_name="h", tokenizer=SD2ClipHTokenizer)
|
| 20 |
+
|
| 21 |
+
class SD2ClipModel(sd1_clip.SD1ClipModel):
|
| 22 |
+
def __init__(self, device="cpu", dtype=None, model_options={}, **kwargs):
|
| 23 |
+
super().__init__(device=device, dtype=dtype, model_options=model_options, clip_name="h", clip_model=SD2ClipHModel, **kwargs)
|
comfy/text_encoders/sd2_clip_config.json
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"CLIPTextModel"
|
| 4 |
+
],
|
| 5 |
+
"attention_dropout": 0.0,
|
| 6 |
+
"bos_token_id": 0,
|
| 7 |
+
"dropout": 0.0,
|
| 8 |
+
"eos_token_id": 49407,
|
| 9 |
+
"hidden_act": "gelu",
|
| 10 |
+
"hidden_size": 1024,
|
| 11 |
+
"initializer_factor": 1.0,
|
| 12 |
+
"initializer_range": 0.02,
|
| 13 |
+
"intermediate_size": 4096,
|
| 14 |
+
"layer_norm_eps": 1e-05,
|
| 15 |
+
"max_position_embeddings": 77,
|
| 16 |
+
"model_type": "clip_text_model",
|
| 17 |
+
"num_attention_heads": 16,
|
| 18 |
+
"num_hidden_layers": 24,
|
| 19 |
+
"pad_token_id": 1,
|
| 20 |
+
"projection_dim": 1024,
|
| 21 |
+
"torch_dtype": "float32",
|
| 22 |
+
"vocab_size": 49408
|
| 23 |
+
}
|
comfy/text_encoders/sd3_clip.py
ADDED
|
@@ -0,0 +1,167 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from comfy import sd1_clip
|
| 2 |
+
from comfy import sdxl_clip
|
| 3 |
+
from transformers import T5TokenizerFast
|
| 4 |
+
import comfy.text_encoders.t5
|
| 5 |
+
import torch
|
| 6 |
+
import os
|
| 7 |
+
import comfy.model_management
|
| 8 |
+
import logging
|
| 9 |
+
import comfy.utils
|
| 10 |
+
|
| 11 |
+
class T5XXLModel(sd1_clip.SDClipModel):
|
| 12 |
+
def __init__(self, device="cpu", layer="last", layer_idx=None, dtype=None, attention_mask=False, model_options={}):
|
| 13 |
+
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "t5_config_xxl.json")
|
| 14 |
+
t5xxl_quantization_metadata = model_options.get("t5xxl_quantization_metadata", None)
|
| 15 |
+
if t5xxl_quantization_metadata is not None:
|
| 16 |
+
model_options = model_options.copy()
|
| 17 |
+
model_options["quantization_metadata"] = t5xxl_quantization_metadata
|
| 18 |
+
|
| 19 |
+
model_options = {**model_options, "model_name": "t5xxl"}
|
| 20 |
+
super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, special_tokens={"end": 1, "pad": 0}, model_class=comfy.text_encoders.t5.T5, enable_attention_masks=attention_mask, return_attention_masks=attention_mask, model_options=model_options)
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def t5_xxl_detect(state_dict, prefix=""):
|
| 24 |
+
out = {}
|
| 25 |
+
t5_key = "{}encoder.final_layer_norm.weight".format(prefix)
|
| 26 |
+
if t5_key in state_dict:
|
| 27 |
+
out["dtype_t5"] = state_dict[t5_key].dtype
|
| 28 |
+
|
| 29 |
+
quant = comfy.utils.detect_layer_quantization(state_dict, prefix)
|
| 30 |
+
if quant is not None:
|
| 31 |
+
out["t5_quantization_metadata"] = quant
|
| 32 |
+
|
| 33 |
+
return out
|
| 34 |
+
|
| 35 |
+
class T5XXLTokenizer(sd1_clip.SDTokenizer):
|
| 36 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}, min_length=77, max_length=99999999):
|
| 37 |
+
tokenizer_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "t5_tokenizer")
|
| 38 |
+
super().__init__(tokenizer_path, embedding_directory=embedding_directory, pad_with_end=False, embedding_size=4096, embedding_key='t5xxl', tokenizer_class=T5TokenizerFast, has_start_token=False, pad_to_max_length=False, max_length=max_length, min_length=min_length, tokenizer_data=tokenizer_data)
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class SD3Tokenizer:
|
| 42 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 43 |
+
self.clip_l = sd1_clip.SDTokenizer(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data)
|
| 44 |
+
self.clip_g = sdxl_clip.SDXLClipGTokenizer(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data)
|
| 45 |
+
self.t5xxl = T5XXLTokenizer(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data)
|
| 46 |
+
|
| 47 |
+
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs):
|
| 48 |
+
out = {}
|
| 49 |
+
out["g"] = self.clip_g.tokenize_with_weights(text, return_word_ids, **kwargs)
|
| 50 |
+
out["l"] = self.clip_l.tokenize_with_weights(text, return_word_ids, **kwargs)
|
| 51 |
+
out["t5xxl"] = self.t5xxl.tokenize_with_weights(text, return_word_ids, **kwargs)
|
| 52 |
+
return out
|
| 53 |
+
|
| 54 |
+
def untokenize(self, token_weight_pair):
|
| 55 |
+
return self.clip_g.untokenize(token_weight_pair)
|
| 56 |
+
|
| 57 |
+
def state_dict(self):
|
| 58 |
+
return {}
|
| 59 |
+
|
| 60 |
+
class SD3ClipModel(torch.nn.Module):
|
| 61 |
+
def __init__(self, clip_l=True, clip_g=True, t5=True, dtype_t5=None, t5_attention_mask=False, device="cpu", dtype=None, model_options={}):
|
| 62 |
+
super().__init__()
|
| 63 |
+
self.dtypes = set()
|
| 64 |
+
if clip_l:
|
| 65 |
+
self.clip_l = sd1_clip.SDClipModel(layer="hidden", layer_idx=-2, device=device, dtype=dtype, layer_norm_hidden_state=False, return_projected_pooled=False, model_options=model_options)
|
| 66 |
+
self.dtypes.add(dtype)
|
| 67 |
+
else:
|
| 68 |
+
self.clip_l = None
|
| 69 |
+
|
| 70 |
+
if clip_g:
|
| 71 |
+
self.clip_g = sdxl_clip.SDXLClipG(device=device, dtype=dtype, model_options=model_options)
|
| 72 |
+
self.dtypes.add(dtype)
|
| 73 |
+
else:
|
| 74 |
+
self.clip_g = None
|
| 75 |
+
|
| 76 |
+
if t5:
|
| 77 |
+
dtype_t5 = comfy.model_management.pick_weight_dtype(dtype_t5, dtype, device)
|
| 78 |
+
self.t5_attention_mask = t5_attention_mask
|
| 79 |
+
self.t5xxl = T5XXLModel(device=device, dtype=dtype_t5, model_options=model_options, attention_mask=self.t5_attention_mask)
|
| 80 |
+
self.dtypes.add(dtype_t5)
|
| 81 |
+
else:
|
| 82 |
+
self.t5xxl = None
|
| 83 |
+
|
| 84 |
+
logging.debug("Created SD3 text encoder with: clip_l {}, clip_g {}, t5xxl {}:{}".format(clip_l, clip_g, t5, dtype_t5))
|
| 85 |
+
|
| 86 |
+
def set_clip_options(self, options):
|
| 87 |
+
if self.clip_l is not None:
|
| 88 |
+
self.clip_l.set_clip_options(options)
|
| 89 |
+
if self.clip_g is not None:
|
| 90 |
+
self.clip_g.set_clip_options(options)
|
| 91 |
+
if self.t5xxl is not None:
|
| 92 |
+
self.t5xxl.set_clip_options(options)
|
| 93 |
+
|
| 94 |
+
def reset_clip_options(self):
|
| 95 |
+
if self.clip_l is not None:
|
| 96 |
+
self.clip_l.reset_clip_options()
|
| 97 |
+
if self.clip_g is not None:
|
| 98 |
+
self.clip_g.reset_clip_options()
|
| 99 |
+
if self.t5xxl is not None:
|
| 100 |
+
self.t5xxl.reset_clip_options()
|
| 101 |
+
|
| 102 |
+
def encode_token_weights(self, token_weight_pairs):
|
| 103 |
+
token_weight_pairs_l = token_weight_pairs["l"]
|
| 104 |
+
token_weight_pairs_g = token_weight_pairs["g"]
|
| 105 |
+
token_weight_pairs_t5 = token_weight_pairs["t5xxl"]
|
| 106 |
+
lg_out = None
|
| 107 |
+
pooled = None
|
| 108 |
+
out = None
|
| 109 |
+
extra = {}
|
| 110 |
+
|
| 111 |
+
if len(token_weight_pairs_g) > 0 or len(token_weight_pairs_l) > 0:
|
| 112 |
+
if self.clip_l is not None:
|
| 113 |
+
lg_out, l_pooled = self.clip_l.encode_token_weights(token_weight_pairs_l)
|
| 114 |
+
else:
|
| 115 |
+
l_pooled = torch.zeros((1, 768), device=comfy.model_management.intermediate_device())
|
| 116 |
+
|
| 117 |
+
if self.clip_g is not None:
|
| 118 |
+
g_out, g_pooled = self.clip_g.encode_token_weights(token_weight_pairs_g)
|
| 119 |
+
if lg_out is not None:
|
| 120 |
+
cut_to = min(lg_out.shape[1], g_out.shape[1])
|
| 121 |
+
lg_out = torch.cat([lg_out[:,:cut_to], g_out[:,:cut_to]], dim=-1)
|
| 122 |
+
else:
|
| 123 |
+
lg_out = torch.nn.functional.pad(g_out, (768, 0))
|
| 124 |
+
else:
|
| 125 |
+
g_out = None
|
| 126 |
+
g_pooled = torch.zeros((1, 1280), device=comfy.model_management.intermediate_device())
|
| 127 |
+
|
| 128 |
+
if lg_out is not None:
|
| 129 |
+
lg_out = torch.nn.functional.pad(lg_out, (0, 4096 - lg_out.shape[-1]))
|
| 130 |
+
out = lg_out
|
| 131 |
+
pooled = torch.cat((l_pooled, g_pooled), dim=-1)
|
| 132 |
+
|
| 133 |
+
if self.t5xxl is not None:
|
| 134 |
+
t5_output = self.t5xxl.encode_token_weights(token_weight_pairs_t5)
|
| 135 |
+
t5_out, t5_pooled = t5_output[:2]
|
| 136 |
+
if self.t5_attention_mask:
|
| 137 |
+
extra["attention_mask"] = t5_output[2]["attention_mask"]
|
| 138 |
+
|
| 139 |
+
if lg_out is not None:
|
| 140 |
+
out = torch.cat([lg_out, t5_out], dim=-2)
|
| 141 |
+
else:
|
| 142 |
+
out = t5_out
|
| 143 |
+
|
| 144 |
+
if out is None:
|
| 145 |
+
out = torch.zeros((1, 77, 4096), device=comfy.model_management.intermediate_device())
|
| 146 |
+
|
| 147 |
+
if pooled is None:
|
| 148 |
+
pooled = torch.zeros((1, 768 + 1280), device=comfy.model_management.intermediate_device())
|
| 149 |
+
|
| 150 |
+
return out, pooled, extra
|
| 151 |
+
|
| 152 |
+
def load_sd(self, sd):
|
| 153 |
+
if "text_model.encoder.layers.30.mlp.fc1.weight" in sd:
|
| 154 |
+
return self.clip_g.load_sd(sd)
|
| 155 |
+
elif "text_model.encoder.layers.1.mlp.fc1.weight" in sd:
|
| 156 |
+
return self.clip_l.load_sd(sd)
|
| 157 |
+
else:
|
| 158 |
+
return self.t5xxl.load_sd(sd)
|
| 159 |
+
|
| 160 |
+
def sd3_clip(clip_l=True, clip_g=True, t5=True, dtype_t5=None, t5_quantization_metadata=None, t5_attention_mask=False):
|
| 161 |
+
class SD3ClipModel_(SD3ClipModel):
|
| 162 |
+
def __init__(self, device="cpu", dtype=None, model_options={}):
|
| 163 |
+
if t5_quantization_metadata is not None:
|
| 164 |
+
model_options = model_options.copy()
|
| 165 |
+
model_options["t5xxl_quantization_metadata"] = t5_quantization_metadata
|
| 166 |
+
super().__init__(clip_l=clip_l, clip_g=clip_g, t5=t5, dtype_t5=dtype_t5, t5_attention_mask=t5_attention_mask, device=device, dtype=dtype, model_options=model_options)
|
| 167 |
+
return SD3ClipModel_
|
comfy/text_encoders/sensenova.py
ADDED
|
@@ -0,0 +1,149 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Tokenizer-only conditioning for SenseNova U1.5.
|
| 2 |
+
|
| 3 |
+
The language model is part of the diffusion checkpoint, so CLIP only needs to
|
| 4 |
+
produce token ids. SenseNova extends the Qwen vocabulary with image-control
|
| 5 |
+
tokens; their order is significant because the checkpoint embeds them by id.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
import os
|
| 9 |
+
|
| 10 |
+
import torch
|
| 11 |
+
from transformers import Qwen2Tokenizer
|
| 12 |
+
|
| 13 |
+
from comfy import sd1_clip
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
SYSTEM_MESSAGE = (
|
| 17 |
+
"You are an image generation and editing assistant that accurately understands and executes "
|
| 18 |
+
"user intent.\n\nYou support two modes:\n\n1. Think Mode:\nIf the task requires reasoning, you "
|
| 19 |
+
"MUST start with a <think></think> block. Put all reasoning inside the block using plain text. "
|
| 20 |
+
"DO NOT include any image tags. Keep it reasonable and directly useful for producing the final "
|
| 21 |
+
"image.\n\n2. Non-Think Mode:\nIf no reasoning is needed, directly produce the final image.\n\n"
|
| 22 |
+
"Task Types:\n\nA. Text-to-Image Generation:\n"
|
| 23 |
+
"- Generate a high-quality image based on the user's description.\n"
|
| 24 |
+
"- Ensure visual clarity, semantic consistency, and completeness.\n"
|
| 25 |
+
"- DO NOT introduce elements that contradict or override the user's intent.\n\n"
|
| 26 |
+
"B. Image Editing:\n"
|
| 27 |
+
"- Use the provided image(s) as input or reference for modification or transformation.\n"
|
| 28 |
+
"- The result can be an edited image or a new image based on the reference(s).\n"
|
| 29 |
+
"- Preserve all unspecified attributes unless explicitly changed.\n\n"
|
| 30 |
+
"General Rules:\n"
|
| 31 |
+
"- For any visible text in the image, follow the language specified for the rendered text in "
|
| 32 |
+
"the user's description, not the language of the prompt. If no language is specified, use the "
|
| 33 |
+
"user's input language."
|
| 34 |
+
)
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def build_generation_prompt(text):
|
| 38 |
+
return (
|
| 39 |
+
f"<|im_start|>system\n{SYSTEM_MESSAGE}<|im_end|>\n"
|
| 40 |
+
f"<|im_start|>user\n{text}<|im_end|>\n"
|
| 41 |
+
"<|im_start|>assistant\n<think>\n\n</think>\n\n<img>"
|
| 42 |
+
)
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def build_unconditional_prompt():
|
| 46 |
+
return "<|im_start|>user\n<|im_end|>\n<|im_start|>assistant\n<img>"
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
class SenseNovaQwen2Tokenizer:
|
| 50 |
+
@classmethod
|
| 51 |
+
def from_pretrained(cls, *args, **kwargs):
|
| 52 |
+
tokenizer = Qwen2Tokenizer.from_pretrained(*args, **kwargs)
|
| 53 |
+
existing_special_tokens = [
|
| 54 |
+
token
|
| 55 |
+
for _, token in sorted(tokenizer.added_tokens_decoder.items())
|
| 56 |
+
if token.special
|
| 57 |
+
]
|
| 58 |
+
extra_tokens = [
|
| 59 |
+
"<IMG_CONTEXT>",
|
| 60 |
+
"<img>",
|
| 61 |
+
"</img>",
|
| 62 |
+
"<quad>",
|
| 63 |
+
"</quad>",
|
| 64 |
+
"<ref>",
|
| 65 |
+
"</ref>",
|
| 66 |
+
"<box>",
|
| 67 |
+
"</box>",
|
| 68 |
+
"<|action_start|>",
|
| 69 |
+
"<|action_end|>",
|
| 70 |
+
"<|plugin|>",
|
| 71 |
+
"<|interpreter|>",
|
| 72 |
+
]
|
| 73 |
+
extra_tokens.extend(f"<FAKE_PAD_{index}>" for index in range(254))
|
| 74 |
+
tokenizer.add_special_tokens(
|
| 75 |
+
{"additional_special_tokens": existing_special_tokens + extra_tokens}
|
| 76 |
+
)
|
| 77 |
+
return tokenizer
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
class SenseNovaQwenTokenizer(sd1_clip.SDTokenizer):
|
| 81 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 82 |
+
tokenizer_path = os.path.join(
|
| 83 |
+
os.path.dirname(os.path.realpath(__file__)), "qwen25_tokenizer"
|
| 84 |
+
)
|
| 85 |
+
super().__init__(
|
| 86 |
+
tokenizer_path,
|
| 87 |
+
pad_with_end=False,
|
| 88 |
+
embedding_size=4096,
|
| 89 |
+
embedding_key="sensenova_u15",
|
| 90 |
+
tokenizer_class=SenseNovaQwen2Tokenizer,
|
| 91 |
+
has_start_token=False,
|
| 92 |
+
has_end_token=False,
|
| 93 |
+
pad_to_max_length=False,
|
| 94 |
+
max_length=99999999,
|
| 95 |
+
min_length=1,
|
| 96 |
+
pad_token=151643,
|
| 97 |
+
tokenizer_data=tokenizer_data,
|
| 98 |
+
)
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
class SenseNovaTokenizer(sd1_clip.SD1Tokenizer):
|
| 102 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 103 |
+
super().__init__(
|
| 104 |
+
embedding_directory=embedding_directory,
|
| 105 |
+
tokenizer_data=tokenizer_data,
|
| 106 |
+
name="sensenova_u15",
|
| 107 |
+
tokenizer=SenseNovaQwenTokenizer,
|
| 108 |
+
)
|
| 109 |
+
|
| 110 |
+
def tokenize_with_weights(self, text, return_word_ids=False, **kwargs):
|
| 111 |
+
prompt = build_generation_prompt(text) if text else build_unconditional_prompt()
|
| 112 |
+
tokens = super().tokenize_with_weights(
|
| 113 |
+
prompt,
|
| 114 |
+
return_word_ids=return_word_ids,
|
| 115 |
+
disable_weights=True,
|
| 116 |
+
**kwargs,
|
| 117 |
+
)
|
| 118 |
+
values = tokens["sensenova_u15"][0]
|
| 119 |
+
values = [value for value in values if int(value[0]) != 151643]
|
| 120 |
+
return {"sensenova_u15": [values]}
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
class SenseNovaTextEncoder(torch.nn.Module):
|
| 124 |
+
def __init__(self, device="cpu", dtype=None, model_options={}):
|
| 125 |
+
super().__init__()
|
| 126 |
+
self.dtypes = {torch.float32}
|
| 127 |
+
self.disable_offload = True
|
| 128 |
+
self.device = torch.device("cpu") if device is None else torch.device(device)
|
| 129 |
+
|
| 130 |
+
def encode_token_weights(self, token_weight_pairs):
|
| 131 |
+
pairs = token_weight_pairs["sensenova_u15"][0]
|
| 132 |
+
input_ids = torch.tensor([[int(value[0]) for value in pairs]], dtype=torch.long)
|
| 133 |
+
return (
|
| 134 |
+
input_ids.unsqueeze(-1).to(torch.float32),
|
| 135 |
+
None,
|
| 136 |
+
{"text_input_ids": input_ids},
|
| 137 |
+
)
|
| 138 |
+
|
| 139 |
+
def load_sd(self, sd):
|
| 140 |
+
return []
|
| 141 |
+
|
| 142 |
+
def get_sd(self):
|
| 143 |
+
return {}
|
| 144 |
+
|
| 145 |
+
def reset_clip_options(self):
|
| 146 |
+
pass
|
| 147 |
+
|
| 148 |
+
def set_clip_options(self, options):
|
| 149 |
+
pass
|
comfy/text_encoders/spiece_tokenizer.py
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import os
|
| 3 |
+
|
| 4 |
+
class SPieceTokenizer:
|
| 5 |
+
@staticmethod
|
| 6 |
+
def from_pretrained(path, **kwargs):
|
| 7 |
+
return SPieceTokenizer(path, **kwargs)
|
| 8 |
+
|
| 9 |
+
def __init__(self, tokenizer_path, add_bos=False, add_eos=True, special_tokens=None):
|
| 10 |
+
self.add_bos = add_bos
|
| 11 |
+
self.add_eos = add_eos
|
| 12 |
+
self.special_tokens = special_tokens
|
| 13 |
+
import sentencepiece
|
| 14 |
+
if torch.is_tensor(tokenizer_path):
|
| 15 |
+
tokenizer_path = tokenizer_path.numpy().tobytes()
|
| 16 |
+
|
| 17 |
+
if isinstance(tokenizer_path, bytes):
|
| 18 |
+
self.tokenizer = sentencepiece.SentencePieceProcessor(model_proto=tokenizer_path, add_bos=self.add_bos, add_eos=self.add_eos)
|
| 19 |
+
else:
|
| 20 |
+
if not os.path.isfile(tokenizer_path):
|
| 21 |
+
raise ValueError("invalid tokenizer")
|
| 22 |
+
self.tokenizer = sentencepiece.SentencePieceProcessor(model_file=tokenizer_path, add_bos=self.add_bos, add_eos=self.add_eos)
|
| 23 |
+
|
| 24 |
+
def get_vocab(self):
|
| 25 |
+
out = {}
|
| 26 |
+
for i in range(self.tokenizer.get_piece_size()):
|
| 27 |
+
out[self.tokenizer.id_to_piece(i)] = i
|
| 28 |
+
return out
|
| 29 |
+
|
| 30 |
+
def __call__(self, string):
|
| 31 |
+
if self.special_tokens is not None:
|
| 32 |
+
import re
|
| 33 |
+
special_tokens_pattern = '|'.join(re.escape(token) for token in self.special_tokens.keys())
|
| 34 |
+
if special_tokens_pattern and re.search(special_tokens_pattern, string):
|
| 35 |
+
parts = re.split(f'({special_tokens_pattern})', string)
|
| 36 |
+
result = []
|
| 37 |
+
for part in parts:
|
| 38 |
+
if not part:
|
| 39 |
+
continue
|
| 40 |
+
if part in self.special_tokens:
|
| 41 |
+
result.append(self.special_tokens[part])
|
| 42 |
+
else:
|
| 43 |
+
encoded = self.tokenizer.encode(part, add_bos=False, add_eos=False)
|
| 44 |
+
result.extend(encoded)
|
| 45 |
+
return {"input_ids": result}
|
| 46 |
+
|
| 47 |
+
out = self.tokenizer.encode(string)
|
| 48 |
+
return {"input_ids": out}
|
| 49 |
+
|
| 50 |
+
def decode(self, token_ids, skip_special_tokens=False):
|
| 51 |
+
|
| 52 |
+
if skip_special_tokens and self.special_tokens:
|
| 53 |
+
special_token_ids = set(self.special_tokens.values())
|
| 54 |
+
token_ids = [tid for tid in token_ids if tid not in special_token_ids]
|
| 55 |
+
|
| 56 |
+
return self.tokenizer.decode(token_ids)
|
| 57 |
+
|
| 58 |
+
def serialize_model(self):
|
| 59 |
+
return torch.ByteTensor(list(self.tokenizer.serialized_model_proto()))
|
comfy/text_encoders/t5.py
ADDED
|
@@ -0,0 +1,249 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import math
|
| 3 |
+
from comfy.ldm.modules.attention import optimized_attention_for_device
|
| 4 |
+
import comfy.ops
|
| 5 |
+
|
| 6 |
+
class T5LayerNorm(torch.nn.Module):
|
| 7 |
+
def __init__(self, hidden_size, eps=1e-6, dtype=None, device=None, operations=None):
|
| 8 |
+
super().__init__()
|
| 9 |
+
self.weight = torch.nn.Parameter(torch.empty(hidden_size, dtype=dtype, device=device))
|
| 10 |
+
self.variance_epsilon = eps
|
| 11 |
+
|
| 12 |
+
def forward(self, x):
|
| 13 |
+
variance = x.pow(2).mean(-1, keepdim=True)
|
| 14 |
+
x = x * torch.rsqrt(variance + self.variance_epsilon)
|
| 15 |
+
return comfy.ops.cast_to_input(self.weight, x) * x
|
| 16 |
+
|
| 17 |
+
activations = {
|
| 18 |
+
"gelu_pytorch_tanh": lambda a: torch.nn.functional.gelu(a, approximate="tanh"),
|
| 19 |
+
"relu": torch.nn.functional.relu,
|
| 20 |
+
}
|
| 21 |
+
|
| 22 |
+
class T5DenseActDense(torch.nn.Module):
|
| 23 |
+
def __init__(self, model_dim, ff_dim, ff_activation, dtype, device, operations):
|
| 24 |
+
super().__init__()
|
| 25 |
+
self.wi = operations.Linear(model_dim, ff_dim, bias=False, dtype=dtype, device=device)
|
| 26 |
+
self.wo = operations.Linear(ff_dim, model_dim, bias=False, dtype=dtype, device=device)
|
| 27 |
+
# self.dropout = nn.Dropout(config.dropout_rate)
|
| 28 |
+
self.act = activations[ff_activation]
|
| 29 |
+
|
| 30 |
+
def forward(self, x):
|
| 31 |
+
x = self.act(self.wi(x))
|
| 32 |
+
# x = self.dropout(x)
|
| 33 |
+
x = self.wo(x)
|
| 34 |
+
return x
|
| 35 |
+
|
| 36 |
+
class T5DenseGatedActDense(torch.nn.Module):
|
| 37 |
+
def __init__(self, model_dim, ff_dim, ff_activation, dtype, device, operations):
|
| 38 |
+
super().__init__()
|
| 39 |
+
self.wi_0 = operations.Linear(model_dim, ff_dim, bias=False, dtype=dtype, device=device)
|
| 40 |
+
self.wi_1 = operations.Linear(model_dim, ff_dim, bias=False, dtype=dtype, device=device)
|
| 41 |
+
self.wo = operations.Linear(ff_dim, model_dim, bias=False, dtype=dtype, device=device)
|
| 42 |
+
# self.dropout = nn.Dropout(config.dropout_rate)
|
| 43 |
+
self.act = activations[ff_activation]
|
| 44 |
+
|
| 45 |
+
def forward(self, x):
|
| 46 |
+
hidden_gelu = self.act(self.wi_0(x))
|
| 47 |
+
hidden_linear = self.wi_1(x)
|
| 48 |
+
x = hidden_gelu * hidden_linear
|
| 49 |
+
# x = self.dropout(x)
|
| 50 |
+
x = self.wo(x)
|
| 51 |
+
return x
|
| 52 |
+
|
| 53 |
+
class T5LayerFF(torch.nn.Module):
|
| 54 |
+
def __init__(self, model_dim, ff_dim, ff_activation, gated_act, dtype, device, operations):
|
| 55 |
+
super().__init__()
|
| 56 |
+
if gated_act:
|
| 57 |
+
self.DenseReluDense = T5DenseGatedActDense(model_dim, ff_dim, ff_activation, dtype, device, operations)
|
| 58 |
+
else:
|
| 59 |
+
self.DenseReluDense = T5DenseActDense(model_dim, ff_dim, ff_activation, dtype, device, operations)
|
| 60 |
+
|
| 61 |
+
self.layer_norm = T5LayerNorm(model_dim, dtype=dtype, device=device, operations=operations)
|
| 62 |
+
# self.dropout = nn.Dropout(config.dropout_rate)
|
| 63 |
+
|
| 64 |
+
def forward(self, x):
|
| 65 |
+
forwarded_states = self.layer_norm(x)
|
| 66 |
+
forwarded_states = self.DenseReluDense(forwarded_states)
|
| 67 |
+
# x = x + self.dropout(forwarded_states)
|
| 68 |
+
x += forwarded_states
|
| 69 |
+
return x
|
| 70 |
+
|
| 71 |
+
class T5Attention(torch.nn.Module):
|
| 72 |
+
def __init__(self, model_dim, inner_dim, num_heads, relative_attention_bias, dtype, device, operations):
|
| 73 |
+
super().__init__()
|
| 74 |
+
|
| 75 |
+
# Mesh TensorFlow initialization to avoid scaling before softmax
|
| 76 |
+
self.q = operations.Linear(model_dim, inner_dim, bias=False, dtype=dtype, device=device)
|
| 77 |
+
self.k = operations.Linear(model_dim, inner_dim, bias=False, dtype=dtype, device=device)
|
| 78 |
+
self.v = operations.Linear(model_dim, inner_dim, bias=False, dtype=dtype, device=device)
|
| 79 |
+
self.o = operations.Linear(inner_dim, model_dim, bias=False, dtype=dtype, device=device)
|
| 80 |
+
self.num_heads = num_heads
|
| 81 |
+
|
| 82 |
+
self.relative_attention_bias = None
|
| 83 |
+
if relative_attention_bias:
|
| 84 |
+
self.relative_attention_num_buckets = 32
|
| 85 |
+
self.relative_attention_max_distance = 128
|
| 86 |
+
self.relative_attention_bias = operations.Embedding(self.relative_attention_num_buckets, self.num_heads, device=device, dtype=dtype)
|
| 87 |
+
|
| 88 |
+
@staticmethod
|
| 89 |
+
def _relative_position_bucket(relative_position, bidirectional=True, num_buckets=32, max_distance=128):
|
| 90 |
+
"""
|
| 91 |
+
Adapted from Mesh Tensorflow:
|
| 92 |
+
https://github.com/tensorflow/mesh/blob/0cb87fe07da627bf0b7e60475d59f95ed6b5be3d/mesh_tensorflow/transformer/transformer_layers.py#L593
|
| 93 |
+
|
| 94 |
+
Translate relative position to a bucket number for relative attention. The relative position is defined as
|
| 95 |
+
memory_position - query_position, i.e. the distance in tokens from the attending position to the attended-to
|
| 96 |
+
position. If bidirectional=False, then positive relative positions are invalid. We use smaller buckets for
|
| 97 |
+
small absolute relative_position and larger buckets for larger absolute relative_positions. All relative
|
| 98 |
+
positions >=max_distance map to the same bucket. All relative positions <=-max_distance map to the same bucket.
|
| 99 |
+
This should allow for more graceful generalization to longer sequences than the model has been trained on
|
| 100 |
+
|
| 101 |
+
Args:
|
| 102 |
+
relative_position: an int32 Tensor
|
| 103 |
+
bidirectional: a boolean - whether the attention is bidirectional
|
| 104 |
+
num_buckets: an integer
|
| 105 |
+
max_distance: an integer
|
| 106 |
+
|
| 107 |
+
Returns:
|
| 108 |
+
a Tensor with the same shape as relative_position, containing int32 values in the range [0, num_buckets)
|
| 109 |
+
"""
|
| 110 |
+
relative_buckets = 0
|
| 111 |
+
if bidirectional:
|
| 112 |
+
num_buckets //= 2
|
| 113 |
+
relative_buckets += (relative_position > 0).to(torch.long) * num_buckets
|
| 114 |
+
relative_position = torch.abs(relative_position)
|
| 115 |
+
else:
|
| 116 |
+
relative_position = -torch.min(relative_position, torch.zeros_like(relative_position))
|
| 117 |
+
# now relative_position is in the range [0, inf)
|
| 118 |
+
|
| 119 |
+
# half of the buckets are for exact increments in positions
|
| 120 |
+
max_exact = num_buckets // 2
|
| 121 |
+
is_small = relative_position < max_exact
|
| 122 |
+
|
| 123 |
+
# The other half of the buckets are for logarithmically bigger bins in positions up to max_distance
|
| 124 |
+
relative_position_if_large = max_exact + (
|
| 125 |
+
torch.log(relative_position.float() / max_exact)
|
| 126 |
+
/ math.log(max_distance / max_exact)
|
| 127 |
+
* (num_buckets - max_exact)
|
| 128 |
+
).to(torch.long)
|
| 129 |
+
relative_position_if_large = torch.min(
|
| 130 |
+
relative_position_if_large, torch.full_like(relative_position_if_large, num_buckets - 1)
|
| 131 |
+
)
|
| 132 |
+
|
| 133 |
+
relative_buckets += torch.where(is_small, relative_position, relative_position_if_large)
|
| 134 |
+
return relative_buckets
|
| 135 |
+
|
| 136 |
+
def compute_bias(self, query_length, key_length, device, dtype):
|
| 137 |
+
"""Compute binned relative position bias"""
|
| 138 |
+
context_position = torch.arange(query_length, dtype=torch.long, device=device)[:, None]
|
| 139 |
+
memory_position = torch.arange(key_length, dtype=torch.long, device=device)[None, :]
|
| 140 |
+
relative_position = memory_position - context_position # shape (query_length, key_length)
|
| 141 |
+
relative_position_bucket = self._relative_position_bucket(
|
| 142 |
+
relative_position, # shape (query_length, key_length)
|
| 143 |
+
bidirectional=True,
|
| 144 |
+
num_buckets=self.relative_attention_num_buckets,
|
| 145 |
+
max_distance=self.relative_attention_max_distance,
|
| 146 |
+
)
|
| 147 |
+
values = self.relative_attention_bias(relative_position_bucket, out_dtype=dtype) # shape (query_length, key_length, num_heads)
|
| 148 |
+
values = values.permute([2, 0, 1]).unsqueeze(0) # shape (1, num_heads, query_length, key_length)
|
| 149 |
+
return values.contiguous()
|
| 150 |
+
|
| 151 |
+
def forward(self, x, mask=None, past_bias=None, optimized_attention=None):
|
| 152 |
+
q = self.q(x)
|
| 153 |
+
k = self.k(x)
|
| 154 |
+
v = self.v(x)
|
| 155 |
+
if self.relative_attention_bias is not None:
|
| 156 |
+
past_bias = self.compute_bias(x.shape[1], x.shape[1], x.device, x.dtype)
|
| 157 |
+
|
| 158 |
+
if past_bias is not None:
|
| 159 |
+
if mask is not None:
|
| 160 |
+
mask = mask + past_bias
|
| 161 |
+
else:
|
| 162 |
+
mask = past_bias
|
| 163 |
+
|
| 164 |
+
out = optimized_attention(q, k * ((k.shape[-1] / self.num_heads) ** 0.5), v, self.num_heads, mask)
|
| 165 |
+
return self.o(out), past_bias
|
| 166 |
+
|
| 167 |
+
class T5LayerSelfAttention(torch.nn.Module):
|
| 168 |
+
def __init__(self, model_dim, inner_dim, ff_dim, num_heads, relative_attention_bias, dtype, device, operations):
|
| 169 |
+
super().__init__()
|
| 170 |
+
self.SelfAttention = T5Attention(model_dim, inner_dim, num_heads, relative_attention_bias, dtype, device, operations)
|
| 171 |
+
self.layer_norm = T5LayerNorm(model_dim, dtype=dtype, device=device, operations=operations)
|
| 172 |
+
# self.dropout = nn.Dropout(config.dropout_rate)
|
| 173 |
+
|
| 174 |
+
def forward(self, x, mask=None, past_bias=None, optimized_attention=None):
|
| 175 |
+
output, past_bias = self.SelfAttention(self.layer_norm(x), mask=mask, past_bias=past_bias, optimized_attention=optimized_attention)
|
| 176 |
+
# x = x + self.dropout(attention_output)
|
| 177 |
+
x += output
|
| 178 |
+
return x, past_bias
|
| 179 |
+
|
| 180 |
+
class T5Block(torch.nn.Module):
|
| 181 |
+
def __init__(self, model_dim, inner_dim, ff_dim, ff_activation, gated_act, num_heads, relative_attention_bias, dtype, device, operations):
|
| 182 |
+
super().__init__()
|
| 183 |
+
self.layer = torch.nn.ModuleList()
|
| 184 |
+
self.layer.append(T5LayerSelfAttention(model_dim, inner_dim, ff_dim, num_heads, relative_attention_bias, dtype, device, operations))
|
| 185 |
+
self.layer.append(T5LayerFF(model_dim, ff_dim, ff_activation, gated_act, dtype, device, operations))
|
| 186 |
+
|
| 187 |
+
def forward(self, x, mask=None, past_bias=None, optimized_attention=None):
|
| 188 |
+
x, past_bias = self.layer[0](x, mask, past_bias, optimized_attention)
|
| 189 |
+
x = self.layer[-1](x)
|
| 190 |
+
return x, past_bias
|
| 191 |
+
|
| 192 |
+
class T5Stack(torch.nn.Module):
|
| 193 |
+
def __init__(self, num_layers, model_dim, inner_dim, ff_dim, ff_activation, gated_act, num_heads, relative_attention, dtype, device, operations):
|
| 194 |
+
super().__init__()
|
| 195 |
+
|
| 196 |
+
self.block = torch.nn.ModuleList(
|
| 197 |
+
[T5Block(model_dim, inner_dim, ff_dim, ff_activation, gated_act, num_heads, relative_attention_bias=((not relative_attention) or (i == 0)), dtype=dtype, device=device, operations=operations) for i in range(num_layers)]
|
| 198 |
+
)
|
| 199 |
+
self.final_layer_norm = T5LayerNorm(model_dim, dtype=dtype, device=device, operations=operations)
|
| 200 |
+
# self.dropout = nn.Dropout(config.dropout_rate)
|
| 201 |
+
|
| 202 |
+
def forward(self, x, attention_mask=None, intermediate_output=None, final_layer_norm_intermediate=True, dtype=None, embeds_info=[]):
|
| 203 |
+
mask = None
|
| 204 |
+
if attention_mask is not None:
|
| 205 |
+
mask = 1.0 - attention_mask.to(x.dtype).reshape((attention_mask.shape[0], 1, -1, attention_mask.shape[-1])).expand(attention_mask.shape[0], 1, attention_mask.shape[-1], attention_mask.shape[-1])
|
| 206 |
+
mask = mask.masked_fill(mask.to(torch.bool), -torch.finfo(x.dtype).max)
|
| 207 |
+
|
| 208 |
+
intermediate = None
|
| 209 |
+
optimized_attention = optimized_attention_for_device(x.device, mask=attention_mask is not None, small_input=True)
|
| 210 |
+
past_bias = None
|
| 211 |
+
|
| 212 |
+
if intermediate_output is not None:
|
| 213 |
+
if intermediate_output < 0:
|
| 214 |
+
intermediate_output = len(self.block) + intermediate_output
|
| 215 |
+
|
| 216 |
+
for i, l in enumerate(self.block):
|
| 217 |
+
x, past_bias = l(x, mask, past_bias, optimized_attention)
|
| 218 |
+
if i == intermediate_output:
|
| 219 |
+
intermediate = x.clone()
|
| 220 |
+
x = self.final_layer_norm(x)
|
| 221 |
+
if intermediate is not None and final_layer_norm_intermediate:
|
| 222 |
+
intermediate = self.final_layer_norm(intermediate)
|
| 223 |
+
return x, intermediate
|
| 224 |
+
|
| 225 |
+
class T5(torch.nn.Module):
|
| 226 |
+
def __init__(self, config_dict, dtype, device, operations):
|
| 227 |
+
super().__init__()
|
| 228 |
+
self.num_layers = config_dict["num_layers"]
|
| 229 |
+
model_dim = config_dict["d_model"]
|
| 230 |
+
inner_dim = config_dict["d_kv"] * config_dict["num_heads"]
|
| 231 |
+
|
| 232 |
+
self.encoder = T5Stack(self.num_layers, model_dim, inner_dim, config_dict["d_ff"], config_dict["dense_act_fn"], config_dict["is_gated_act"], config_dict["num_heads"], config_dict["model_type"] != "umt5", dtype, device, operations)
|
| 233 |
+
self.dtype = dtype
|
| 234 |
+
self.shared = operations.Embedding(config_dict["vocab_size"], model_dim, device=device, dtype=dtype)
|
| 235 |
+
|
| 236 |
+
def get_input_embeddings(self):
|
| 237 |
+
return self.shared
|
| 238 |
+
|
| 239 |
+
def set_input_embeddings(self, embeddings):
|
| 240 |
+
self.shared = embeddings
|
| 241 |
+
|
| 242 |
+
def forward(self, input_ids, attention_mask, embeds=None, num_tokens=None, **kwargs):
|
| 243 |
+
if input_ids is None:
|
| 244 |
+
x = embeds
|
| 245 |
+
else:
|
| 246 |
+
x = self.shared(input_ids, out_dtype=kwargs.get("dtype", torch.float32))
|
| 247 |
+
if self.dtype not in [torch.float32, torch.float16, torch.bfloat16]:
|
| 248 |
+
x = torch.nan_to_num(x) #Fix for fp8 T5 base
|
| 249 |
+
return self.encoder(x, attention_mask=attention_mask, **kwargs)
|
comfy/text_encoders/t5_config_base.json
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"d_ff": 3072,
|
| 3 |
+
"d_kv": 64,
|
| 4 |
+
"d_model": 768,
|
| 5 |
+
"decoder_start_token_id": 0,
|
| 6 |
+
"dropout_rate": 0.1,
|
| 7 |
+
"eos_token_id": 1,
|
| 8 |
+
"dense_act_fn": "relu",
|
| 9 |
+
"initializer_factor": 1.0,
|
| 10 |
+
"is_encoder_decoder": true,
|
| 11 |
+
"is_gated_act": false,
|
| 12 |
+
"layer_norm_epsilon": 1e-06,
|
| 13 |
+
"model_type": "t5",
|
| 14 |
+
"num_decoder_layers": 12,
|
| 15 |
+
"num_heads": 12,
|
| 16 |
+
"num_layers": 12,
|
| 17 |
+
"output_past": true,
|
| 18 |
+
"pad_token_id": 0,
|
| 19 |
+
"relative_attention_num_buckets": 32,
|
| 20 |
+
"tie_word_embeddings": false,
|
| 21 |
+
"vocab_size": 32128
|
| 22 |
+
}
|
comfy/text_encoders/t5_config_xxl.json
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"d_ff": 10240,
|
| 3 |
+
"d_kv": 64,
|
| 4 |
+
"d_model": 4096,
|
| 5 |
+
"decoder_start_token_id": 0,
|
| 6 |
+
"dropout_rate": 0.1,
|
| 7 |
+
"eos_token_id": 1,
|
| 8 |
+
"dense_act_fn": "gelu_pytorch_tanh",
|
| 9 |
+
"initializer_factor": 1.0,
|
| 10 |
+
"is_encoder_decoder": true,
|
| 11 |
+
"is_gated_act": true,
|
| 12 |
+
"layer_norm_epsilon": 1e-06,
|
| 13 |
+
"model_type": "t5",
|
| 14 |
+
"num_decoder_layers": 24,
|
| 15 |
+
"num_heads": 64,
|
| 16 |
+
"num_layers": 24,
|
| 17 |
+
"output_past": true,
|
| 18 |
+
"pad_token_id": 0,
|
| 19 |
+
"relative_attention_num_buckets": 32,
|
| 20 |
+
"tie_word_embeddings": false,
|
| 21 |
+
"vocab_size": 32128
|
| 22 |
+
}
|
comfy/text_encoders/t5_old_config_xxl.json
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"d_ff": 65536,
|
| 3 |
+
"d_kv": 128,
|
| 4 |
+
"d_model": 1024,
|
| 5 |
+
"decoder_start_token_id": 0,
|
| 6 |
+
"dropout_rate": 0.1,
|
| 7 |
+
"eos_token_id": 1,
|
| 8 |
+
"dense_act_fn": "relu",
|
| 9 |
+
"initializer_factor": 1.0,
|
| 10 |
+
"is_encoder_decoder": true,
|
| 11 |
+
"is_gated_act": false,
|
| 12 |
+
"layer_norm_epsilon": 1e-06,
|
| 13 |
+
"model_type": "t5",
|
| 14 |
+
"num_decoder_layers": 24,
|
| 15 |
+
"num_heads": 128,
|
| 16 |
+
"num_layers": 24,
|
| 17 |
+
"output_past": true,
|
| 18 |
+
"pad_token_id": 0,
|
| 19 |
+
"relative_attention_num_buckets": 32,
|
| 20 |
+
"tie_word_embeddings": false,
|
| 21 |
+
"vocab_size": 32128
|
| 22 |
+
}
|
comfy/text_encoders/t5_pile_config_xl.json
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"d_ff": 5120,
|
| 3 |
+
"d_kv": 64,
|
| 4 |
+
"d_model": 2048,
|
| 5 |
+
"decoder_start_token_id": 0,
|
| 6 |
+
"dropout_rate": 0.1,
|
| 7 |
+
"eos_token_id": 2,
|
| 8 |
+
"dense_act_fn": "gelu_pytorch_tanh",
|
| 9 |
+
"initializer_factor": 1.0,
|
| 10 |
+
"is_encoder_decoder": true,
|
| 11 |
+
"is_gated_act": true,
|
| 12 |
+
"layer_norm_epsilon": 1e-06,
|
| 13 |
+
"model_type": "umt5",
|
| 14 |
+
"num_decoder_layers": 24,
|
| 15 |
+
"num_heads": 32,
|
| 16 |
+
"num_layers": 24,
|
| 17 |
+
"output_past": true,
|
| 18 |
+
"pad_token_id": 1,
|
| 19 |
+
"relative_attention_num_buckets": 32,
|
| 20 |
+
"tie_word_embeddings": false,
|
| 21 |
+
"vocab_size": 32128
|
| 22 |
+
}
|
comfy/text_encoders/t5_pile_tokenizer/tokenizer.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:9e556afd44213b6bd1be2b850ebbbd98f5481437a8021afaf58ee7fb1818d347
|
| 3 |
+
size 499723
|
comfy/text_encoders/t5_tokenizer/special_tokens_map.json
ADDED
|
@@ -0,0 +1,125 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"additional_special_tokens": [
|
| 3 |
+
"<extra_id_0>",
|
| 4 |
+
"<extra_id_1>",
|
| 5 |
+
"<extra_id_2>",
|
| 6 |
+
"<extra_id_3>",
|
| 7 |
+
"<extra_id_4>",
|
| 8 |
+
"<extra_id_5>",
|
| 9 |
+
"<extra_id_6>",
|
| 10 |
+
"<extra_id_7>",
|
| 11 |
+
"<extra_id_8>",
|
| 12 |
+
"<extra_id_9>",
|
| 13 |
+
"<extra_id_10>",
|
| 14 |
+
"<extra_id_11>",
|
| 15 |
+
"<extra_id_12>",
|
| 16 |
+
"<extra_id_13>",
|
| 17 |
+
"<extra_id_14>",
|
| 18 |
+
"<extra_id_15>",
|
| 19 |
+
"<extra_id_16>",
|
| 20 |
+
"<extra_id_17>",
|
| 21 |
+
"<extra_id_18>",
|
| 22 |
+
"<extra_id_19>",
|
| 23 |
+
"<extra_id_20>",
|
| 24 |
+
"<extra_id_21>",
|
| 25 |
+
"<extra_id_22>",
|
| 26 |
+
"<extra_id_23>",
|
| 27 |
+
"<extra_id_24>",
|
| 28 |
+
"<extra_id_25>",
|
| 29 |
+
"<extra_id_26>",
|
| 30 |
+
"<extra_id_27>",
|
| 31 |
+
"<extra_id_28>",
|
| 32 |
+
"<extra_id_29>",
|
| 33 |
+
"<extra_id_30>",
|
| 34 |
+
"<extra_id_31>",
|
| 35 |
+
"<extra_id_32>",
|
| 36 |
+
"<extra_id_33>",
|
| 37 |
+
"<extra_id_34>",
|
| 38 |
+
"<extra_id_35>",
|
| 39 |
+
"<extra_id_36>",
|
| 40 |
+
"<extra_id_37>",
|
| 41 |
+
"<extra_id_38>",
|
| 42 |
+
"<extra_id_39>",
|
| 43 |
+
"<extra_id_40>",
|
| 44 |
+
"<extra_id_41>",
|
| 45 |
+
"<extra_id_42>",
|
| 46 |
+
"<extra_id_43>",
|
| 47 |
+
"<extra_id_44>",
|
| 48 |
+
"<extra_id_45>",
|
| 49 |
+
"<extra_id_46>",
|
| 50 |
+
"<extra_id_47>",
|
| 51 |
+
"<extra_id_48>",
|
| 52 |
+
"<extra_id_49>",
|
| 53 |
+
"<extra_id_50>",
|
| 54 |
+
"<extra_id_51>",
|
| 55 |
+
"<extra_id_52>",
|
| 56 |
+
"<extra_id_53>",
|
| 57 |
+
"<extra_id_54>",
|
| 58 |
+
"<extra_id_55>",
|
| 59 |
+
"<extra_id_56>",
|
| 60 |
+
"<extra_id_57>",
|
| 61 |
+
"<extra_id_58>",
|
| 62 |
+
"<extra_id_59>",
|
| 63 |
+
"<extra_id_60>",
|
| 64 |
+
"<extra_id_61>",
|
| 65 |
+
"<extra_id_62>",
|
| 66 |
+
"<extra_id_63>",
|
| 67 |
+
"<extra_id_64>",
|
| 68 |
+
"<extra_id_65>",
|
| 69 |
+
"<extra_id_66>",
|
| 70 |
+
"<extra_id_67>",
|
| 71 |
+
"<extra_id_68>",
|
| 72 |
+
"<extra_id_69>",
|
| 73 |
+
"<extra_id_70>",
|
| 74 |
+
"<extra_id_71>",
|
| 75 |
+
"<extra_id_72>",
|
| 76 |
+
"<extra_id_73>",
|
| 77 |
+
"<extra_id_74>",
|
| 78 |
+
"<extra_id_75>",
|
| 79 |
+
"<extra_id_76>",
|
| 80 |
+
"<extra_id_77>",
|
| 81 |
+
"<extra_id_78>",
|
| 82 |
+
"<extra_id_79>",
|
| 83 |
+
"<extra_id_80>",
|
| 84 |
+
"<extra_id_81>",
|
| 85 |
+
"<extra_id_82>",
|
| 86 |
+
"<extra_id_83>",
|
| 87 |
+
"<extra_id_84>",
|
| 88 |
+
"<extra_id_85>",
|
| 89 |
+
"<extra_id_86>",
|
| 90 |
+
"<extra_id_87>",
|
| 91 |
+
"<extra_id_88>",
|
| 92 |
+
"<extra_id_89>",
|
| 93 |
+
"<extra_id_90>",
|
| 94 |
+
"<extra_id_91>",
|
| 95 |
+
"<extra_id_92>",
|
| 96 |
+
"<extra_id_93>",
|
| 97 |
+
"<extra_id_94>",
|
| 98 |
+
"<extra_id_95>",
|
| 99 |
+
"<extra_id_96>",
|
| 100 |
+
"<extra_id_97>",
|
| 101 |
+
"<extra_id_98>",
|
| 102 |
+
"<extra_id_99>"
|
| 103 |
+
],
|
| 104 |
+
"eos_token": {
|
| 105 |
+
"content": "</s>",
|
| 106 |
+
"lstrip": false,
|
| 107 |
+
"normalized": false,
|
| 108 |
+
"rstrip": false,
|
| 109 |
+
"single_word": false
|
| 110 |
+
},
|
| 111 |
+
"pad_token": {
|
| 112 |
+
"content": "<pad>",
|
| 113 |
+
"lstrip": false,
|
| 114 |
+
"normalized": false,
|
| 115 |
+
"rstrip": false,
|
| 116 |
+
"single_word": false
|
| 117 |
+
},
|
| 118 |
+
"unk_token": {
|
| 119 |
+
"content": "<unk>",
|
| 120 |
+
"lstrip": false,
|
| 121 |
+
"normalized": false,
|
| 122 |
+
"rstrip": false,
|
| 123 |
+
"single_word": false
|
| 124 |
+
}
|
| 125 |
+
}
|
comfy/text_encoders/t5_tokenizer/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
comfy/text_encoders/t5_tokenizer/tokenizer_config.json
ADDED
|
@@ -0,0 +1,939 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"added_tokens_decoder": {
|
| 3 |
+
"0": {
|
| 4 |
+
"content": "<pad>",
|
| 5 |
+
"lstrip": false,
|
| 6 |
+
"normalized": false,
|
| 7 |
+
"rstrip": false,
|
| 8 |
+
"single_word": false,
|
| 9 |
+
"special": true
|
| 10 |
+
},
|
| 11 |
+
"1": {
|
| 12 |
+
"content": "</s>",
|
| 13 |
+
"lstrip": false,
|
| 14 |
+
"normalized": false,
|
| 15 |
+
"rstrip": false,
|
| 16 |
+
"single_word": false,
|
| 17 |
+
"special": true
|
| 18 |
+
},
|
| 19 |
+
"2": {
|
| 20 |
+
"content": "<unk>",
|
| 21 |
+
"lstrip": false,
|
| 22 |
+
"normalized": false,
|
| 23 |
+
"rstrip": false,
|
| 24 |
+
"single_word": false,
|
| 25 |
+
"special": true
|
| 26 |
+
},
|
| 27 |
+
"32000": {
|
| 28 |
+
"content": "<extra_id_99>",
|
| 29 |
+
"lstrip": false,
|
| 30 |
+
"normalized": false,
|
| 31 |
+
"rstrip": false,
|
| 32 |
+
"single_word": false,
|
| 33 |
+
"special": true
|
| 34 |
+
},
|
| 35 |
+
"32001": {
|
| 36 |
+
"content": "<extra_id_98>",
|
| 37 |
+
"lstrip": false,
|
| 38 |
+
"normalized": false,
|
| 39 |
+
"rstrip": false,
|
| 40 |
+
"single_word": false,
|
| 41 |
+
"special": true
|
| 42 |
+
},
|
| 43 |
+
"32002": {
|
| 44 |
+
"content": "<extra_id_97>",
|
| 45 |
+
"lstrip": false,
|
| 46 |
+
"normalized": false,
|
| 47 |
+
"rstrip": false,
|
| 48 |
+
"single_word": false,
|
| 49 |
+
"special": true
|
| 50 |
+
},
|
| 51 |
+
"32003": {
|
| 52 |
+
"content": "<extra_id_96>",
|
| 53 |
+
"lstrip": false,
|
| 54 |
+
"normalized": false,
|
| 55 |
+
"rstrip": false,
|
| 56 |
+
"single_word": false,
|
| 57 |
+
"special": true
|
| 58 |
+
},
|
| 59 |
+
"32004": {
|
| 60 |
+
"content": "<extra_id_95>",
|
| 61 |
+
"lstrip": false,
|
| 62 |
+
"normalized": false,
|
| 63 |
+
"rstrip": false,
|
| 64 |
+
"single_word": false,
|
| 65 |
+
"special": true
|
| 66 |
+
},
|
| 67 |
+
"32005": {
|
| 68 |
+
"content": "<extra_id_94>",
|
| 69 |
+
"lstrip": false,
|
| 70 |
+
"normalized": false,
|
| 71 |
+
"rstrip": false,
|
| 72 |
+
"single_word": false,
|
| 73 |
+
"special": true
|
| 74 |
+
},
|
| 75 |
+
"32006": {
|
| 76 |
+
"content": "<extra_id_93>",
|
| 77 |
+
"lstrip": false,
|
| 78 |
+
"normalized": false,
|
| 79 |
+
"rstrip": false,
|
| 80 |
+
"single_word": false,
|
| 81 |
+
"special": true
|
| 82 |
+
},
|
| 83 |
+
"32007": {
|
| 84 |
+
"content": "<extra_id_92>",
|
| 85 |
+
"lstrip": false,
|
| 86 |
+
"normalized": false,
|
| 87 |
+
"rstrip": false,
|
| 88 |
+
"single_word": false,
|
| 89 |
+
"special": true
|
| 90 |
+
},
|
| 91 |
+
"32008": {
|
| 92 |
+
"content": "<extra_id_91>",
|
| 93 |
+
"lstrip": false,
|
| 94 |
+
"normalized": false,
|
| 95 |
+
"rstrip": false,
|
| 96 |
+
"single_word": false,
|
| 97 |
+
"special": true
|
| 98 |
+
},
|
| 99 |
+
"32009": {
|
| 100 |
+
"content": "<extra_id_90>",
|
| 101 |
+
"lstrip": false,
|
| 102 |
+
"normalized": false,
|
| 103 |
+
"rstrip": false,
|
| 104 |
+
"single_word": false,
|
| 105 |
+
"special": true
|
| 106 |
+
},
|
| 107 |
+
"32010": {
|
| 108 |
+
"content": "<extra_id_89>",
|
| 109 |
+
"lstrip": false,
|
| 110 |
+
"normalized": false,
|
| 111 |
+
"rstrip": false,
|
| 112 |
+
"single_word": false,
|
| 113 |
+
"special": true
|
| 114 |
+
},
|
| 115 |
+
"32011": {
|
| 116 |
+
"content": "<extra_id_88>",
|
| 117 |
+
"lstrip": false,
|
| 118 |
+
"normalized": false,
|
| 119 |
+
"rstrip": false,
|
| 120 |
+
"single_word": false,
|
| 121 |
+
"special": true
|
| 122 |
+
},
|
| 123 |
+
"32012": {
|
| 124 |
+
"content": "<extra_id_87>",
|
| 125 |
+
"lstrip": false,
|
| 126 |
+
"normalized": false,
|
| 127 |
+
"rstrip": false,
|
| 128 |
+
"single_word": false,
|
| 129 |
+
"special": true
|
| 130 |
+
},
|
| 131 |
+
"32013": {
|
| 132 |
+
"content": "<extra_id_86>",
|
| 133 |
+
"lstrip": false,
|
| 134 |
+
"normalized": false,
|
| 135 |
+
"rstrip": false,
|
| 136 |
+
"single_word": false,
|
| 137 |
+
"special": true
|
| 138 |
+
},
|
| 139 |
+
"32014": {
|
| 140 |
+
"content": "<extra_id_85>",
|
| 141 |
+
"lstrip": false,
|
| 142 |
+
"normalized": false,
|
| 143 |
+
"rstrip": false,
|
| 144 |
+
"single_word": false,
|
| 145 |
+
"special": true
|
| 146 |
+
},
|
| 147 |
+
"32015": {
|
| 148 |
+
"content": "<extra_id_84>",
|
| 149 |
+
"lstrip": false,
|
| 150 |
+
"normalized": false,
|
| 151 |
+
"rstrip": false,
|
| 152 |
+
"single_word": false,
|
| 153 |
+
"special": true
|
| 154 |
+
},
|
| 155 |
+
"32016": {
|
| 156 |
+
"content": "<extra_id_83>",
|
| 157 |
+
"lstrip": false,
|
| 158 |
+
"normalized": false,
|
| 159 |
+
"rstrip": false,
|
| 160 |
+
"single_word": false,
|
| 161 |
+
"special": true
|
| 162 |
+
},
|
| 163 |
+
"32017": {
|
| 164 |
+
"content": "<extra_id_82>",
|
| 165 |
+
"lstrip": false,
|
| 166 |
+
"normalized": false,
|
| 167 |
+
"rstrip": false,
|
| 168 |
+
"single_word": false,
|
| 169 |
+
"special": true
|
| 170 |
+
},
|
| 171 |
+
"32018": {
|
| 172 |
+
"content": "<extra_id_81>",
|
| 173 |
+
"lstrip": false,
|
| 174 |
+
"normalized": false,
|
| 175 |
+
"rstrip": false,
|
| 176 |
+
"single_word": false,
|
| 177 |
+
"special": true
|
| 178 |
+
},
|
| 179 |
+
"32019": {
|
| 180 |
+
"content": "<extra_id_80>",
|
| 181 |
+
"lstrip": false,
|
| 182 |
+
"normalized": false,
|
| 183 |
+
"rstrip": false,
|
| 184 |
+
"single_word": false,
|
| 185 |
+
"special": true
|
| 186 |
+
},
|
| 187 |
+
"32020": {
|
| 188 |
+
"content": "<extra_id_79>",
|
| 189 |
+
"lstrip": false,
|
| 190 |
+
"normalized": false,
|
| 191 |
+
"rstrip": false,
|
| 192 |
+
"single_word": false,
|
| 193 |
+
"special": true
|
| 194 |
+
},
|
| 195 |
+
"32021": {
|
| 196 |
+
"content": "<extra_id_78>",
|
| 197 |
+
"lstrip": false,
|
| 198 |
+
"normalized": false,
|
| 199 |
+
"rstrip": false,
|
| 200 |
+
"single_word": false,
|
| 201 |
+
"special": true
|
| 202 |
+
},
|
| 203 |
+
"32022": {
|
| 204 |
+
"content": "<extra_id_77>",
|
| 205 |
+
"lstrip": false,
|
| 206 |
+
"normalized": false,
|
| 207 |
+
"rstrip": false,
|
| 208 |
+
"single_word": false,
|
| 209 |
+
"special": true
|
| 210 |
+
},
|
| 211 |
+
"32023": {
|
| 212 |
+
"content": "<extra_id_76>",
|
| 213 |
+
"lstrip": false,
|
| 214 |
+
"normalized": false,
|
| 215 |
+
"rstrip": false,
|
| 216 |
+
"single_word": false,
|
| 217 |
+
"special": true
|
| 218 |
+
},
|
| 219 |
+
"32024": {
|
| 220 |
+
"content": "<extra_id_75>",
|
| 221 |
+
"lstrip": false,
|
| 222 |
+
"normalized": false,
|
| 223 |
+
"rstrip": false,
|
| 224 |
+
"single_word": false,
|
| 225 |
+
"special": true
|
| 226 |
+
},
|
| 227 |
+
"32025": {
|
| 228 |
+
"content": "<extra_id_74>",
|
| 229 |
+
"lstrip": false,
|
| 230 |
+
"normalized": false,
|
| 231 |
+
"rstrip": false,
|
| 232 |
+
"single_word": false,
|
| 233 |
+
"special": true
|
| 234 |
+
},
|
| 235 |
+
"32026": {
|
| 236 |
+
"content": "<extra_id_73>",
|
| 237 |
+
"lstrip": false,
|
| 238 |
+
"normalized": false,
|
| 239 |
+
"rstrip": false,
|
| 240 |
+
"single_word": false,
|
| 241 |
+
"special": true
|
| 242 |
+
},
|
| 243 |
+
"32027": {
|
| 244 |
+
"content": "<extra_id_72>",
|
| 245 |
+
"lstrip": false,
|
| 246 |
+
"normalized": false,
|
| 247 |
+
"rstrip": false,
|
| 248 |
+
"single_word": false,
|
| 249 |
+
"special": true
|
| 250 |
+
},
|
| 251 |
+
"32028": {
|
| 252 |
+
"content": "<extra_id_71>",
|
| 253 |
+
"lstrip": false,
|
| 254 |
+
"normalized": false,
|
| 255 |
+
"rstrip": false,
|
| 256 |
+
"single_word": false,
|
| 257 |
+
"special": true
|
| 258 |
+
},
|
| 259 |
+
"32029": {
|
| 260 |
+
"content": "<extra_id_70>",
|
| 261 |
+
"lstrip": false,
|
| 262 |
+
"normalized": false,
|
| 263 |
+
"rstrip": false,
|
| 264 |
+
"single_word": false,
|
| 265 |
+
"special": true
|
| 266 |
+
},
|
| 267 |
+
"32030": {
|
| 268 |
+
"content": "<extra_id_69>",
|
| 269 |
+
"lstrip": false,
|
| 270 |
+
"normalized": false,
|
| 271 |
+
"rstrip": false,
|
| 272 |
+
"single_word": false,
|
| 273 |
+
"special": true
|
| 274 |
+
},
|
| 275 |
+
"32031": {
|
| 276 |
+
"content": "<extra_id_68>",
|
| 277 |
+
"lstrip": false,
|
| 278 |
+
"normalized": false,
|
| 279 |
+
"rstrip": false,
|
| 280 |
+
"single_word": false,
|
| 281 |
+
"special": true
|
| 282 |
+
},
|
| 283 |
+
"32032": {
|
| 284 |
+
"content": "<extra_id_67>",
|
| 285 |
+
"lstrip": false,
|
| 286 |
+
"normalized": false,
|
| 287 |
+
"rstrip": false,
|
| 288 |
+
"single_word": false,
|
| 289 |
+
"special": true
|
| 290 |
+
},
|
| 291 |
+
"32033": {
|
| 292 |
+
"content": "<extra_id_66>",
|
| 293 |
+
"lstrip": false,
|
| 294 |
+
"normalized": false,
|
| 295 |
+
"rstrip": false,
|
| 296 |
+
"single_word": false,
|
| 297 |
+
"special": true
|
| 298 |
+
},
|
| 299 |
+
"32034": {
|
| 300 |
+
"content": "<extra_id_65>",
|
| 301 |
+
"lstrip": false,
|
| 302 |
+
"normalized": false,
|
| 303 |
+
"rstrip": false,
|
| 304 |
+
"single_word": false,
|
| 305 |
+
"special": true
|
| 306 |
+
},
|
| 307 |
+
"32035": {
|
| 308 |
+
"content": "<extra_id_64>",
|
| 309 |
+
"lstrip": false,
|
| 310 |
+
"normalized": false,
|
| 311 |
+
"rstrip": false,
|
| 312 |
+
"single_word": false,
|
| 313 |
+
"special": true
|
| 314 |
+
},
|
| 315 |
+
"32036": {
|
| 316 |
+
"content": "<extra_id_63>",
|
| 317 |
+
"lstrip": false,
|
| 318 |
+
"normalized": false,
|
| 319 |
+
"rstrip": false,
|
| 320 |
+
"single_word": false,
|
| 321 |
+
"special": true
|
| 322 |
+
},
|
| 323 |
+
"32037": {
|
| 324 |
+
"content": "<extra_id_62>",
|
| 325 |
+
"lstrip": false,
|
| 326 |
+
"normalized": false,
|
| 327 |
+
"rstrip": false,
|
| 328 |
+
"single_word": false,
|
| 329 |
+
"special": true
|
| 330 |
+
},
|
| 331 |
+
"32038": {
|
| 332 |
+
"content": "<extra_id_61>",
|
| 333 |
+
"lstrip": false,
|
| 334 |
+
"normalized": false,
|
| 335 |
+
"rstrip": false,
|
| 336 |
+
"single_word": false,
|
| 337 |
+
"special": true
|
| 338 |
+
},
|
| 339 |
+
"32039": {
|
| 340 |
+
"content": "<extra_id_60>",
|
| 341 |
+
"lstrip": false,
|
| 342 |
+
"normalized": false,
|
| 343 |
+
"rstrip": false,
|
| 344 |
+
"single_word": false,
|
| 345 |
+
"special": true
|
| 346 |
+
},
|
| 347 |
+
"32040": {
|
| 348 |
+
"content": "<extra_id_59>",
|
| 349 |
+
"lstrip": false,
|
| 350 |
+
"normalized": false,
|
| 351 |
+
"rstrip": false,
|
| 352 |
+
"single_word": false,
|
| 353 |
+
"special": true
|
| 354 |
+
},
|
| 355 |
+
"32041": {
|
| 356 |
+
"content": "<extra_id_58>",
|
| 357 |
+
"lstrip": false,
|
| 358 |
+
"normalized": false,
|
| 359 |
+
"rstrip": false,
|
| 360 |
+
"single_word": false,
|
| 361 |
+
"special": true
|
| 362 |
+
},
|
| 363 |
+
"32042": {
|
| 364 |
+
"content": "<extra_id_57>",
|
| 365 |
+
"lstrip": false,
|
| 366 |
+
"normalized": false,
|
| 367 |
+
"rstrip": false,
|
| 368 |
+
"single_word": false,
|
| 369 |
+
"special": true
|
| 370 |
+
},
|
| 371 |
+
"32043": {
|
| 372 |
+
"content": "<extra_id_56>",
|
| 373 |
+
"lstrip": false,
|
| 374 |
+
"normalized": false,
|
| 375 |
+
"rstrip": false,
|
| 376 |
+
"single_word": false,
|
| 377 |
+
"special": true
|
| 378 |
+
},
|
| 379 |
+
"32044": {
|
| 380 |
+
"content": "<extra_id_55>",
|
| 381 |
+
"lstrip": false,
|
| 382 |
+
"normalized": false,
|
| 383 |
+
"rstrip": false,
|
| 384 |
+
"single_word": false,
|
| 385 |
+
"special": true
|
| 386 |
+
},
|
| 387 |
+
"32045": {
|
| 388 |
+
"content": "<extra_id_54>",
|
| 389 |
+
"lstrip": false,
|
| 390 |
+
"normalized": false,
|
| 391 |
+
"rstrip": false,
|
| 392 |
+
"single_word": false,
|
| 393 |
+
"special": true
|
| 394 |
+
},
|
| 395 |
+
"32046": {
|
| 396 |
+
"content": "<extra_id_53>",
|
| 397 |
+
"lstrip": false,
|
| 398 |
+
"normalized": false,
|
| 399 |
+
"rstrip": false,
|
| 400 |
+
"single_word": false,
|
| 401 |
+
"special": true
|
| 402 |
+
},
|
| 403 |
+
"32047": {
|
| 404 |
+
"content": "<extra_id_52>",
|
| 405 |
+
"lstrip": false,
|
| 406 |
+
"normalized": false,
|
| 407 |
+
"rstrip": false,
|
| 408 |
+
"single_word": false,
|
| 409 |
+
"special": true
|
| 410 |
+
},
|
| 411 |
+
"32048": {
|
| 412 |
+
"content": "<extra_id_51>",
|
| 413 |
+
"lstrip": false,
|
| 414 |
+
"normalized": false,
|
| 415 |
+
"rstrip": false,
|
| 416 |
+
"single_word": false,
|
| 417 |
+
"special": true
|
| 418 |
+
},
|
| 419 |
+
"32049": {
|
| 420 |
+
"content": "<extra_id_50>",
|
| 421 |
+
"lstrip": false,
|
| 422 |
+
"normalized": false,
|
| 423 |
+
"rstrip": false,
|
| 424 |
+
"single_word": false,
|
| 425 |
+
"special": true
|
| 426 |
+
},
|
| 427 |
+
"32050": {
|
| 428 |
+
"content": "<extra_id_49>",
|
| 429 |
+
"lstrip": false,
|
| 430 |
+
"normalized": false,
|
| 431 |
+
"rstrip": false,
|
| 432 |
+
"single_word": false,
|
| 433 |
+
"special": true
|
| 434 |
+
},
|
| 435 |
+
"32051": {
|
| 436 |
+
"content": "<extra_id_48>",
|
| 437 |
+
"lstrip": false,
|
| 438 |
+
"normalized": false,
|
| 439 |
+
"rstrip": false,
|
| 440 |
+
"single_word": false,
|
| 441 |
+
"special": true
|
| 442 |
+
},
|
| 443 |
+
"32052": {
|
| 444 |
+
"content": "<extra_id_47>",
|
| 445 |
+
"lstrip": false,
|
| 446 |
+
"normalized": false,
|
| 447 |
+
"rstrip": false,
|
| 448 |
+
"single_word": false,
|
| 449 |
+
"special": true
|
| 450 |
+
},
|
| 451 |
+
"32053": {
|
| 452 |
+
"content": "<extra_id_46>",
|
| 453 |
+
"lstrip": false,
|
| 454 |
+
"normalized": false,
|
| 455 |
+
"rstrip": false,
|
| 456 |
+
"single_word": false,
|
| 457 |
+
"special": true
|
| 458 |
+
},
|
| 459 |
+
"32054": {
|
| 460 |
+
"content": "<extra_id_45>",
|
| 461 |
+
"lstrip": false,
|
| 462 |
+
"normalized": false,
|
| 463 |
+
"rstrip": false,
|
| 464 |
+
"single_word": false,
|
| 465 |
+
"special": true
|
| 466 |
+
},
|
| 467 |
+
"32055": {
|
| 468 |
+
"content": "<extra_id_44>",
|
| 469 |
+
"lstrip": false,
|
| 470 |
+
"normalized": false,
|
| 471 |
+
"rstrip": false,
|
| 472 |
+
"single_word": false,
|
| 473 |
+
"special": true
|
| 474 |
+
},
|
| 475 |
+
"32056": {
|
| 476 |
+
"content": "<extra_id_43>",
|
| 477 |
+
"lstrip": false,
|
| 478 |
+
"normalized": false,
|
| 479 |
+
"rstrip": false,
|
| 480 |
+
"single_word": false,
|
| 481 |
+
"special": true
|
| 482 |
+
},
|
| 483 |
+
"32057": {
|
| 484 |
+
"content": "<extra_id_42>",
|
| 485 |
+
"lstrip": false,
|
| 486 |
+
"normalized": false,
|
| 487 |
+
"rstrip": false,
|
| 488 |
+
"single_word": false,
|
| 489 |
+
"special": true
|
| 490 |
+
},
|
| 491 |
+
"32058": {
|
| 492 |
+
"content": "<extra_id_41>",
|
| 493 |
+
"lstrip": false,
|
| 494 |
+
"normalized": false,
|
| 495 |
+
"rstrip": false,
|
| 496 |
+
"single_word": false,
|
| 497 |
+
"special": true
|
| 498 |
+
},
|
| 499 |
+
"32059": {
|
| 500 |
+
"content": "<extra_id_40>",
|
| 501 |
+
"lstrip": false,
|
| 502 |
+
"normalized": false,
|
| 503 |
+
"rstrip": false,
|
| 504 |
+
"single_word": false,
|
| 505 |
+
"special": true
|
| 506 |
+
},
|
| 507 |
+
"32060": {
|
| 508 |
+
"content": "<extra_id_39>",
|
| 509 |
+
"lstrip": false,
|
| 510 |
+
"normalized": false,
|
| 511 |
+
"rstrip": false,
|
| 512 |
+
"single_word": false,
|
| 513 |
+
"special": true
|
| 514 |
+
},
|
| 515 |
+
"32061": {
|
| 516 |
+
"content": "<extra_id_38>",
|
| 517 |
+
"lstrip": false,
|
| 518 |
+
"normalized": false,
|
| 519 |
+
"rstrip": false,
|
| 520 |
+
"single_word": false,
|
| 521 |
+
"special": true
|
| 522 |
+
},
|
| 523 |
+
"32062": {
|
| 524 |
+
"content": "<extra_id_37>",
|
| 525 |
+
"lstrip": false,
|
| 526 |
+
"normalized": false,
|
| 527 |
+
"rstrip": false,
|
| 528 |
+
"single_word": false,
|
| 529 |
+
"special": true
|
| 530 |
+
},
|
| 531 |
+
"32063": {
|
| 532 |
+
"content": "<extra_id_36>",
|
| 533 |
+
"lstrip": false,
|
| 534 |
+
"normalized": false,
|
| 535 |
+
"rstrip": false,
|
| 536 |
+
"single_word": false,
|
| 537 |
+
"special": true
|
| 538 |
+
},
|
| 539 |
+
"32064": {
|
| 540 |
+
"content": "<extra_id_35>",
|
| 541 |
+
"lstrip": false,
|
| 542 |
+
"normalized": false,
|
| 543 |
+
"rstrip": false,
|
| 544 |
+
"single_word": false,
|
| 545 |
+
"special": true
|
| 546 |
+
},
|
| 547 |
+
"32065": {
|
| 548 |
+
"content": "<extra_id_34>",
|
| 549 |
+
"lstrip": false,
|
| 550 |
+
"normalized": false,
|
| 551 |
+
"rstrip": false,
|
| 552 |
+
"single_word": false,
|
| 553 |
+
"special": true
|
| 554 |
+
},
|
| 555 |
+
"32066": {
|
| 556 |
+
"content": "<extra_id_33>",
|
| 557 |
+
"lstrip": false,
|
| 558 |
+
"normalized": false,
|
| 559 |
+
"rstrip": false,
|
| 560 |
+
"single_word": false,
|
| 561 |
+
"special": true
|
| 562 |
+
},
|
| 563 |
+
"32067": {
|
| 564 |
+
"content": "<extra_id_32>",
|
| 565 |
+
"lstrip": false,
|
| 566 |
+
"normalized": false,
|
| 567 |
+
"rstrip": false,
|
| 568 |
+
"single_word": false,
|
| 569 |
+
"special": true
|
| 570 |
+
},
|
| 571 |
+
"32068": {
|
| 572 |
+
"content": "<extra_id_31>",
|
| 573 |
+
"lstrip": false,
|
| 574 |
+
"normalized": false,
|
| 575 |
+
"rstrip": false,
|
| 576 |
+
"single_word": false,
|
| 577 |
+
"special": true
|
| 578 |
+
},
|
| 579 |
+
"32069": {
|
| 580 |
+
"content": "<extra_id_30>",
|
| 581 |
+
"lstrip": false,
|
| 582 |
+
"normalized": false,
|
| 583 |
+
"rstrip": false,
|
| 584 |
+
"single_word": false,
|
| 585 |
+
"special": true
|
| 586 |
+
},
|
| 587 |
+
"32070": {
|
| 588 |
+
"content": "<extra_id_29>",
|
| 589 |
+
"lstrip": false,
|
| 590 |
+
"normalized": false,
|
| 591 |
+
"rstrip": false,
|
| 592 |
+
"single_word": false,
|
| 593 |
+
"special": true
|
| 594 |
+
},
|
| 595 |
+
"32071": {
|
| 596 |
+
"content": "<extra_id_28>",
|
| 597 |
+
"lstrip": false,
|
| 598 |
+
"normalized": false,
|
| 599 |
+
"rstrip": false,
|
| 600 |
+
"single_word": false,
|
| 601 |
+
"special": true
|
| 602 |
+
},
|
| 603 |
+
"32072": {
|
| 604 |
+
"content": "<extra_id_27>",
|
| 605 |
+
"lstrip": false,
|
| 606 |
+
"normalized": false,
|
| 607 |
+
"rstrip": false,
|
| 608 |
+
"single_word": false,
|
| 609 |
+
"special": true
|
| 610 |
+
},
|
| 611 |
+
"32073": {
|
| 612 |
+
"content": "<extra_id_26>",
|
| 613 |
+
"lstrip": false,
|
| 614 |
+
"normalized": false,
|
| 615 |
+
"rstrip": false,
|
| 616 |
+
"single_word": false,
|
| 617 |
+
"special": true
|
| 618 |
+
},
|
| 619 |
+
"32074": {
|
| 620 |
+
"content": "<extra_id_25>",
|
| 621 |
+
"lstrip": false,
|
| 622 |
+
"normalized": false,
|
| 623 |
+
"rstrip": false,
|
| 624 |
+
"single_word": false,
|
| 625 |
+
"special": true
|
| 626 |
+
},
|
| 627 |
+
"32075": {
|
| 628 |
+
"content": "<extra_id_24>",
|
| 629 |
+
"lstrip": false,
|
| 630 |
+
"normalized": false,
|
| 631 |
+
"rstrip": false,
|
| 632 |
+
"single_word": false,
|
| 633 |
+
"special": true
|
| 634 |
+
},
|
| 635 |
+
"32076": {
|
| 636 |
+
"content": "<extra_id_23>",
|
| 637 |
+
"lstrip": false,
|
| 638 |
+
"normalized": false,
|
| 639 |
+
"rstrip": false,
|
| 640 |
+
"single_word": false,
|
| 641 |
+
"special": true
|
| 642 |
+
},
|
| 643 |
+
"32077": {
|
| 644 |
+
"content": "<extra_id_22>",
|
| 645 |
+
"lstrip": false,
|
| 646 |
+
"normalized": false,
|
| 647 |
+
"rstrip": false,
|
| 648 |
+
"single_word": false,
|
| 649 |
+
"special": true
|
| 650 |
+
},
|
| 651 |
+
"32078": {
|
| 652 |
+
"content": "<extra_id_21>",
|
| 653 |
+
"lstrip": false,
|
| 654 |
+
"normalized": false,
|
| 655 |
+
"rstrip": false,
|
| 656 |
+
"single_word": false,
|
| 657 |
+
"special": true
|
| 658 |
+
},
|
| 659 |
+
"32079": {
|
| 660 |
+
"content": "<extra_id_20>",
|
| 661 |
+
"lstrip": false,
|
| 662 |
+
"normalized": false,
|
| 663 |
+
"rstrip": false,
|
| 664 |
+
"single_word": false,
|
| 665 |
+
"special": true
|
| 666 |
+
},
|
| 667 |
+
"32080": {
|
| 668 |
+
"content": "<extra_id_19>",
|
| 669 |
+
"lstrip": false,
|
| 670 |
+
"normalized": false,
|
| 671 |
+
"rstrip": false,
|
| 672 |
+
"single_word": false,
|
| 673 |
+
"special": true
|
| 674 |
+
},
|
| 675 |
+
"32081": {
|
| 676 |
+
"content": "<extra_id_18>",
|
| 677 |
+
"lstrip": false,
|
| 678 |
+
"normalized": false,
|
| 679 |
+
"rstrip": false,
|
| 680 |
+
"single_word": false,
|
| 681 |
+
"special": true
|
| 682 |
+
},
|
| 683 |
+
"32082": {
|
| 684 |
+
"content": "<extra_id_17>",
|
| 685 |
+
"lstrip": false,
|
| 686 |
+
"normalized": false,
|
| 687 |
+
"rstrip": false,
|
| 688 |
+
"single_word": false,
|
| 689 |
+
"special": true
|
| 690 |
+
},
|
| 691 |
+
"32083": {
|
| 692 |
+
"content": "<extra_id_16>",
|
| 693 |
+
"lstrip": false,
|
| 694 |
+
"normalized": false,
|
| 695 |
+
"rstrip": false,
|
| 696 |
+
"single_word": false,
|
| 697 |
+
"special": true
|
| 698 |
+
},
|
| 699 |
+
"32084": {
|
| 700 |
+
"content": "<extra_id_15>",
|
| 701 |
+
"lstrip": false,
|
| 702 |
+
"normalized": false,
|
| 703 |
+
"rstrip": false,
|
| 704 |
+
"single_word": false,
|
| 705 |
+
"special": true
|
| 706 |
+
},
|
| 707 |
+
"32085": {
|
| 708 |
+
"content": "<extra_id_14>",
|
| 709 |
+
"lstrip": false,
|
| 710 |
+
"normalized": false,
|
| 711 |
+
"rstrip": false,
|
| 712 |
+
"single_word": false,
|
| 713 |
+
"special": true
|
| 714 |
+
},
|
| 715 |
+
"32086": {
|
| 716 |
+
"content": "<extra_id_13>",
|
| 717 |
+
"lstrip": false,
|
| 718 |
+
"normalized": false,
|
| 719 |
+
"rstrip": false,
|
| 720 |
+
"single_word": false,
|
| 721 |
+
"special": true
|
| 722 |
+
},
|
| 723 |
+
"32087": {
|
| 724 |
+
"content": "<extra_id_12>",
|
| 725 |
+
"lstrip": false,
|
| 726 |
+
"normalized": false,
|
| 727 |
+
"rstrip": false,
|
| 728 |
+
"single_word": false,
|
| 729 |
+
"special": true
|
| 730 |
+
},
|
| 731 |
+
"32088": {
|
| 732 |
+
"content": "<extra_id_11>",
|
| 733 |
+
"lstrip": false,
|
| 734 |
+
"normalized": false,
|
| 735 |
+
"rstrip": false,
|
| 736 |
+
"single_word": false,
|
| 737 |
+
"special": true
|
| 738 |
+
},
|
| 739 |
+
"32089": {
|
| 740 |
+
"content": "<extra_id_10>",
|
| 741 |
+
"lstrip": false,
|
| 742 |
+
"normalized": false,
|
| 743 |
+
"rstrip": false,
|
| 744 |
+
"single_word": false,
|
| 745 |
+
"special": true
|
| 746 |
+
},
|
| 747 |
+
"32090": {
|
| 748 |
+
"content": "<extra_id_9>",
|
| 749 |
+
"lstrip": false,
|
| 750 |
+
"normalized": false,
|
| 751 |
+
"rstrip": false,
|
| 752 |
+
"single_word": false,
|
| 753 |
+
"special": true
|
| 754 |
+
},
|
| 755 |
+
"32091": {
|
| 756 |
+
"content": "<extra_id_8>",
|
| 757 |
+
"lstrip": false,
|
| 758 |
+
"normalized": false,
|
| 759 |
+
"rstrip": false,
|
| 760 |
+
"single_word": false,
|
| 761 |
+
"special": true
|
| 762 |
+
},
|
| 763 |
+
"32092": {
|
| 764 |
+
"content": "<extra_id_7>",
|
| 765 |
+
"lstrip": false,
|
| 766 |
+
"normalized": false,
|
| 767 |
+
"rstrip": false,
|
| 768 |
+
"single_word": false,
|
| 769 |
+
"special": true
|
| 770 |
+
},
|
| 771 |
+
"32093": {
|
| 772 |
+
"content": "<extra_id_6>",
|
| 773 |
+
"lstrip": false,
|
| 774 |
+
"normalized": false,
|
| 775 |
+
"rstrip": false,
|
| 776 |
+
"single_word": false,
|
| 777 |
+
"special": true
|
| 778 |
+
},
|
| 779 |
+
"32094": {
|
| 780 |
+
"content": "<extra_id_5>",
|
| 781 |
+
"lstrip": false,
|
| 782 |
+
"normalized": false,
|
| 783 |
+
"rstrip": false,
|
| 784 |
+
"single_word": false,
|
| 785 |
+
"special": true
|
| 786 |
+
},
|
| 787 |
+
"32095": {
|
| 788 |
+
"content": "<extra_id_4>",
|
| 789 |
+
"lstrip": false,
|
| 790 |
+
"normalized": false,
|
| 791 |
+
"rstrip": false,
|
| 792 |
+
"single_word": false,
|
| 793 |
+
"special": true
|
| 794 |
+
},
|
| 795 |
+
"32096": {
|
| 796 |
+
"content": "<extra_id_3>",
|
| 797 |
+
"lstrip": false,
|
| 798 |
+
"normalized": false,
|
| 799 |
+
"rstrip": false,
|
| 800 |
+
"single_word": false,
|
| 801 |
+
"special": true
|
| 802 |
+
},
|
| 803 |
+
"32097": {
|
| 804 |
+
"content": "<extra_id_2>",
|
| 805 |
+
"lstrip": false,
|
| 806 |
+
"normalized": false,
|
| 807 |
+
"rstrip": false,
|
| 808 |
+
"single_word": false,
|
| 809 |
+
"special": true
|
| 810 |
+
},
|
| 811 |
+
"32098": {
|
| 812 |
+
"content": "<extra_id_1>",
|
| 813 |
+
"lstrip": false,
|
| 814 |
+
"normalized": false,
|
| 815 |
+
"rstrip": false,
|
| 816 |
+
"single_word": false,
|
| 817 |
+
"special": true
|
| 818 |
+
},
|
| 819 |
+
"32099": {
|
| 820 |
+
"content": "<extra_id_0>",
|
| 821 |
+
"lstrip": false,
|
| 822 |
+
"normalized": false,
|
| 823 |
+
"rstrip": false,
|
| 824 |
+
"single_word": false,
|
| 825 |
+
"special": true
|
| 826 |
+
}
|
| 827 |
+
},
|
| 828 |
+
"additional_special_tokens": [
|
| 829 |
+
"<extra_id_0>",
|
| 830 |
+
"<extra_id_1>",
|
| 831 |
+
"<extra_id_2>",
|
| 832 |
+
"<extra_id_3>",
|
| 833 |
+
"<extra_id_4>",
|
| 834 |
+
"<extra_id_5>",
|
| 835 |
+
"<extra_id_6>",
|
| 836 |
+
"<extra_id_7>",
|
| 837 |
+
"<extra_id_8>",
|
| 838 |
+
"<extra_id_9>",
|
| 839 |
+
"<extra_id_10>",
|
| 840 |
+
"<extra_id_11>",
|
| 841 |
+
"<extra_id_12>",
|
| 842 |
+
"<extra_id_13>",
|
| 843 |
+
"<extra_id_14>",
|
| 844 |
+
"<extra_id_15>",
|
| 845 |
+
"<extra_id_16>",
|
| 846 |
+
"<extra_id_17>",
|
| 847 |
+
"<extra_id_18>",
|
| 848 |
+
"<extra_id_19>",
|
| 849 |
+
"<extra_id_20>",
|
| 850 |
+
"<extra_id_21>",
|
| 851 |
+
"<extra_id_22>",
|
| 852 |
+
"<extra_id_23>",
|
| 853 |
+
"<extra_id_24>",
|
| 854 |
+
"<extra_id_25>",
|
| 855 |
+
"<extra_id_26>",
|
| 856 |
+
"<extra_id_27>",
|
| 857 |
+
"<extra_id_28>",
|
| 858 |
+
"<extra_id_29>",
|
| 859 |
+
"<extra_id_30>",
|
| 860 |
+
"<extra_id_31>",
|
| 861 |
+
"<extra_id_32>",
|
| 862 |
+
"<extra_id_33>",
|
| 863 |
+
"<extra_id_34>",
|
| 864 |
+
"<extra_id_35>",
|
| 865 |
+
"<extra_id_36>",
|
| 866 |
+
"<extra_id_37>",
|
| 867 |
+
"<extra_id_38>",
|
| 868 |
+
"<extra_id_39>",
|
| 869 |
+
"<extra_id_40>",
|
| 870 |
+
"<extra_id_41>",
|
| 871 |
+
"<extra_id_42>",
|
| 872 |
+
"<extra_id_43>",
|
| 873 |
+
"<extra_id_44>",
|
| 874 |
+
"<extra_id_45>",
|
| 875 |
+
"<extra_id_46>",
|
| 876 |
+
"<extra_id_47>",
|
| 877 |
+
"<extra_id_48>",
|
| 878 |
+
"<extra_id_49>",
|
| 879 |
+
"<extra_id_50>",
|
| 880 |
+
"<extra_id_51>",
|
| 881 |
+
"<extra_id_52>",
|
| 882 |
+
"<extra_id_53>",
|
| 883 |
+
"<extra_id_54>",
|
| 884 |
+
"<extra_id_55>",
|
| 885 |
+
"<extra_id_56>",
|
| 886 |
+
"<extra_id_57>",
|
| 887 |
+
"<extra_id_58>",
|
| 888 |
+
"<extra_id_59>",
|
| 889 |
+
"<extra_id_60>",
|
| 890 |
+
"<extra_id_61>",
|
| 891 |
+
"<extra_id_62>",
|
| 892 |
+
"<extra_id_63>",
|
| 893 |
+
"<extra_id_64>",
|
| 894 |
+
"<extra_id_65>",
|
| 895 |
+
"<extra_id_66>",
|
| 896 |
+
"<extra_id_67>",
|
| 897 |
+
"<extra_id_68>",
|
| 898 |
+
"<extra_id_69>",
|
| 899 |
+
"<extra_id_70>",
|
| 900 |
+
"<extra_id_71>",
|
| 901 |
+
"<extra_id_72>",
|
| 902 |
+
"<extra_id_73>",
|
| 903 |
+
"<extra_id_74>",
|
| 904 |
+
"<extra_id_75>",
|
| 905 |
+
"<extra_id_76>",
|
| 906 |
+
"<extra_id_77>",
|
| 907 |
+
"<extra_id_78>",
|
| 908 |
+
"<extra_id_79>",
|
| 909 |
+
"<extra_id_80>",
|
| 910 |
+
"<extra_id_81>",
|
| 911 |
+
"<extra_id_82>",
|
| 912 |
+
"<extra_id_83>",
|
| 913 |
+
"<extra_id_84>",
|
| 914 |
+
"<extra_id_85>",
|
| 915 |
+
"<extra_id_86>",
|
| 916 |
+
"<extra_id_87>",
|
| 917 |
+
"<extra_id_88>",
|
| 918 |
+
"<extra_id_89>",
|
| 919 |
+
"<extra_id_90>",
|
| 920 |
+
"<extra_id_91>",
|
| 921 |
+
"<extra_id_92>",
|
| 922 |
+
"<extra_id_93>",
|
| 923 |
+
"<extra_id_94>",
|
| 924 |
+
"<extra_id_95>",
|
| 925 |
+
"<extra_id_96>",
|
| 926 |
+
"<extra_id_97>",
|
| 927 |
+
"<extra_id_98>",
|
| 928 |
+
"<extra_id_99>"
|
| 929 |
+
],
|
| 930 |
+
"clean_up_tokenization_spaces": true,
|
| 931 |
+
"eos_token": "</s>",
|
| 932 |
+
"extra_ids": 100,
|
| 933 |
+
"legacy": false,
|
| 934 |
+
"model_max_length": 512,
|
| 935 |
+
"pad_token": "<pad>",
|
| 936 |
+
"sp_model_kwargs": {},
|
| 937 |
+
"tokenizer_class": "T5Tokenizer",
|
| 938 |
+
"unk_token": "<unk>"
|
| 939 |
+
}
|
comfy/text_encoders/umt5_config_base.json
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"d_ff": 2048,
|
| 3 |
+
"d_kv": 64,
|
| 4 |
+
"d_model": 768,
|
| 5 |
+
"decoder_start_token_id": 0,
|
| 6 |
+
"dropout_rate": 0.1,
|
| 7 |
+
"eos_token_id": 1,
|
| 8 |
+
"dense_act_fn": "gelu_pytorch_tanh",
|
| 9 |
+
"initializer_factor": 1.0,
|
| 10 |
+
"is_encoder_decoder": true,
|
| 11 |
+
"is_gated_act": true,
|
| 12 |
+
"layer_norm_epsilon": 1e-06,
|
| 13 |
+
"model_type": "umt5",
|
| 14 |
+
"num_decoder_layers": 12,
|
| 15 |
+
"num_heads": 12,
|
| 16 |
+
"num_layers": 12,
|
| 17 |
+
"output_past": true,
|
| 18 |
+
"pad_token_id": 0,
|
| 19 |
+
"relative_attention_num_buckets": 32,
|
| 20 |
+
"tie_word_embeddings": false,
|
| 21 |
+
"vocab_size": 256384
|
| 22 |
+
}
|
comfy/text_encoders/umt5_config_xxl.json
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"d_ff": 10240,
|
| 3 |
+
"d_kv": 64,
|
| 4 |
+
"d_model": 4096,
|
| 5 |
+
"decoder_start_token_id": 0,
|
| 6 |
+
"dropout_rate": 0.1,
|
| 7 |
+
"eos_token_id": 1,
|
| 8 |
+
"dense_act_fn": "gelu_pytorch_tanh",
|
| 9 |
+
"initializer_factor": 1.0,
|
| 10 |
+
"is_encoder_decoder": true,
|
| 11 |
+
"is_gated_act": true,
|
| 12 |
+
"layer_norm_epsilon": 1e-06,
|
| 13 |
+
"model_type": "umt5",
|
| 14 |
+
"num_decoder_layers": 24,
|
| 15 |
+
"num_heads": 64,
|
| 16 |
+
"num_layers": 24,
|
| 17 |
+
"output_past": true,
|
| 18 |
+
"pad_token_id": 0,
|
| 19 |
+
"relative_attention_num_buckets": 32,
|
| 20 |
+
"tie_word_embeddings": false,
|
| 21 |
+
"vocab_size": 256384
|
| 22 |
+
}
|
comfy/text_encoders/wan.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from comfy import sd1_clip
|
| 2 |
+
from .spiece_tokenizer import SPieceTokenizer
|
| 3 |
+
import comfy.text_encoders.t5
|
| 4 |
+
import os
|
| 5 |
+
|
| 6 |
+
class UMT5XXlModel(sd1_clip.SDClipModel):
|
| 7 |
+
def __init__(self, device="cpu", layer="last", layer_idx=None, dtype=None, model_options={}):
|
| 8 |
+
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "umt5_config_xxl.json")
|
| 9 |
+
super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, dtype=dtype, special_tokens={"end": 1, "pad": 0}, model_class=comfy.text_encoders.t5.T5, enable_attention_masks=True, zero_out_masked=True, model_options=model_options)
|
| 10 |
+
|
| 11 |
+
class UMT5XXlTokenizer(sd1_clip.SDTokenizer):
|
| 12 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 13 |
+
tokenizer = tokenizer_data.get("spiece_model", None)
|
| 14 |
+
super().__init__(tokenizer, pad_with_end=False, embedding_size=4096, embedding_key='umt5xxl', tokenizer_class=SPieceTokenizer, has_start_token=False, pad_to_max_length=False, max_length=99999999, min_length=512, pad_token=0, tokenizer_data=tokenizer_data)
|
| 15 |
+
|
| 16 |
+
def state_dict(self):
|
| 17 |
+
return {"spiece_model": self.tokenizer.serialize_model()}
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class WanT5Tokenizer(sd1_clip.SD1Tokenizer):
|
| 21 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 22 |
+
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, clip_name="umt5xxl", tokenizer=UMT5XXlTokenizer)
|
| 23 |
+
|
| 24 |
+
class WanT5Model(sd1_clip.SD1ClipModel):
|
| 25 |
+
def __init__(self, device="cpu", dtype=None, model_options={}, **kwargs):
|
| 26 |
+
super().__init__(device=device, dtype=dtype, model_options=model_options, name="umt5xxl", clip_model=UMT5XXlModel, **kwargs)
|
| 27 |
+
|
| 28 |
+
def te(dtype_t5=None, t5_quantization_metadata=None):
|
| 29 |
+
class WanTEModel(WanT5Model):
|
| 30 |
+
def __init__(self, device="cpu", dtype=None, model_options={}):
|
| 31 |
+
if t5_quantization_metadata is not None:
|
| 32 |
+
model_options = model_options.copy()
|
| 33 |
+
model_options["quantization_metadata"] = t5_quantization_metadata
|
| 34 |
+
if dtype_t5 is not None:
|
| 35 |
+
dtype = dtype_t5
|
| 36 |
+
super().__init__(device=device, dtype=dtype, model_options=model_options)
|
| 37 |
+
return WanTEModel
|
comfy/text_encoders/z_image.py
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from transformers import Qwen2Tokenizer
|
| 2 |
+
import comfy.text_encoders.llama
|
| 3 |
+
from comfy import sd1_clip
|
| 4 |
+
import os
|
| 5 |
+
|
| 6 |
+
class Qwen3Tokenizer(sd1_clip.SDTokenizer):
|
| 7 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 8 |
+
tokenizer_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "qwen25_tokenizer")
|
| 9 |
+
super().__init__(tokenizer_path, pad_with_end=False, embedding_directory=embedding_directory, embedding_size=2560, embedding_key='qwen3_4b', tokenizer_class=Qwen2Tokenizer, has_start_token=False, has_end_token=False, pad_to_max_length=False, max_length=99999999, min_length=1, pad_token=151643, tokenizer_data=tokenizer_data)
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class ZImageTokenizer(sd1_clip.SD1Tokenizer):
|
| 13 |
+
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
| 14 |
+
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="qwen3_4b", tokenizer=Qwen3Tokenizer)
|
| 15 |
+
self.llama_template = "<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
|
| 16 |
+
|
| 17 |
+
def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, **kwargs):
|
| 18 |
+
if llama_template is None:
|
| 19 |
+
llama_text = self.llama_template.format(text)
|
| 20 |
+
else:
|
| 21 |
+
llama_text = llama_template.format(text)
|
| 22 |
+
|
| 23 |
+
tokens = super().tokenize_with_weights(llama_text, return_word_ids=return_word_ids, disable_weights=True, **kwargs)
|
| 24 |
+
return tokens
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class Qwen3_4BModel(sd1_clip.SDClipModel):
|
| 28 |
+
def __init__(self, device="cpu", layer="hidden", layer_idx=-2, dtype=None, attention_mask=True, model_options={}):
|
| 29 |
+
super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, dtype=dtype, special_tokens={"pad": 151643}, layer_norm_hidden_state=False, model_class=comfy.text_encoders.llama.Qwen3_4B, enable_attention_masks=attention_mask, return_attention_masks=attention_mask, model_options=model_options)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
class ZImageTEModel(sd1_clip.SD1ClipModel):
|
| 33 |
+
def __init__(self, device="cpu", dtype=None, model_options={}):
|
| 34 |
+
super().__init__(device=device, dtype=dtype, name="qwen3_4b", clip_model=Qwen3_4BModel, model_options=model_options)
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def te(dtype_llama=None, llama_quantization_metadata=None):
|
| 38 |
+
class ZImageTEModel_(ZImageTEModel):
|
| 39 |
+
def __init__(self, device="cpu", dtype=None, model_options={}):
|
| 40 |
+
if dtype_llama is not None:
|
| 41 |
+
dtype = dtype_llama
|
| 42 |
+
if llama_quantization_metadata is not None:
|
| 43 |
+
model_options = model_options.copy()
|
| 44 |
+
model_options["quantization_metadata"] = llama_quantization_metadata
|
| 45 |
+
super().__init__(device=device, dtype=dtype, model_options=model_options)
|
| 46 |
+
return ZImageTEModel_
|
comfy/utils.py
ADDED
|
@@ -0,0 +1,1535 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
This file is part of ComfyUI.
|
| 3 |
+
Copyright (C) 2024 Comfy
|
| 4 |
+
|
| 5 |
+
This program is free software: you can redistribute it and/or modify
|
| 6 |
+
it under the terms of the GNU General Public License as published by
|
| 7 |
+
the Free Software Foundation, either version 3 of the License, or
|
| 8 |
+
(at your option) any later version.
|
| 9 |
+
|
| 10 |
+
This program is distributed in the hope that it will be useful,
|
| 11 |
+
but WITHOUT ANY WARRANTY; without even the implied warranty of
|
| 12 |
+
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
| 13 |
+
GNU General Public License for more details.
|
| 14 |
+
|
| 15 |
+
You should have received a copy of the GNU General Public License
|
| 16 |
+
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
| 17 |
+
"""
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
import torch
|
| 21 |
+
import math
|
| 22 |
+
import struct
|
| 23 |
+
import ctypes
|
| 24 |
+
import os
|
| 25 |
+
import comfy.memory_management
|
| 26 |
+
import safetensors.torch
|
| 27 |
+
import numpy as np
|
| 28 |
+
from PIL import Image
|
| 29 |
+
import logging
|
| 30 |
+
import itertools
|
| 31 |
+
from torch.nn.functional import interpolate
|
| 32 |
+
from tqdm.auto import trange
|
| 33 |
+
from einops import rearrange
|
| 34 |
+
from comfy.cli_args import args
|
| 35 |
+
import json
|
| 36 |
+
import time
|
| 37 |
+
import threading
|
| 38 |
+
import warnings
|
| 39 |
+
|
| 40 |
+
MMAP_TORCH_FILES = args.mmap_torch_files
|
| 41 |
+
DISABLE_MMAP = args.disable_mmap
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
if True: # ckpt/pt file whitelist for safe loading of old sd files
|
| 45 |
+
class ModelCheckpoint:
|
| 46 |
+
pass
|
| 47 |
+
ModelCheckpoint.__module__ = "pytorch_lightning.callbacks.model_checkpoint"
|
| 48 |
+
|
| 49 |
+
def scalar(*args, **kwargs):
|
| 50 |
+
return None
|
| 51 |
+
scalar.__module__ = "numpy.core.multiarray"
|
| 52 |
+
|
| 53 |
+
from numpy import dtype
|
| 54 |
+
from numpy.dtypes import Float64DType
|
| 55 |
+
|
| 56 |
+
def encode(*args, **kwargs): # no longer necessary on newer torch
|
| 57 |
+
return None
|
| 58 |
+
encode.__module__ = "_codecs"
|
| 59 |
+
|
| 60 |
+
torch.serialization.add_safe_globals([ModelCheckpoint, scalar, dtype, Float64DType, encode])
|
| 61 |
+
logging.info("Checkpoint files will always be loaded safely.")
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
# Current as of safetensors 0.7.0
|
| 65 |
+
_TYPES = {
|
| 66 |
+
"F64": torch.float64,
|
| 67 |
+
"F32": torch.float32,
|
| 68 |
+
"F16": torch.float16,
|
| 69 |
+
"BF16": torch.bfloat16,
|
| 70 |
+
"I64": torch.int64,
|
| 71 |
+
"I32": torch.int32,
|
| 72 |
+
"I16": torch.int16,
|
| 73 |
+
"I8": torch.int8,
|
| 74 |
+
"U8": torch.uint8,
|
| 75 |
+
"BOOL": torch.bool,
|
| 76 |
+
"F8_E4M3": torch.float8_e4m3fn,
|
| 77 |
+
"F8_E5M2": torch.float8_e5m2,
|
| 78 |
+
"C64": torch.complex64,
|
| 79 |
+
|
| 80 |
+
"U64": torch.uint64,
|
| 81 |
+
"U32": torch.uint32,
|
| 82 |
+
"U16": torch.uint16,
|
| 83 |
+
}
|
| 84 |
+
|
| 85 |
+
_SAFETENSORS_MAX_HEADER_SIZE = 100_000_000
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
def _invalid_safetensors_error(message, ckpt):
|
| 89 |
+
return ValueError("{}\n\nFile path: {}\n\nThe safetensors file is corrupt or invalid. Make sure this is actually a safetensors file and not a ckpt or pt or other filetype.".format(message, ckpt))
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
def _incomplete_safetensors_error(message, ckpt):
|
| 93 |
+
return ValueError("{}\n\nFile path: {}\n\nThe safetensors file is corrupt/incomplete. Check the file size and make sure you have copied/downloaded it correctly.".format(message, ckpt))
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
def load_safetensors(ckpt):
|
| 97 |
+
import comfy_aimdo.model_mmap
|
| 98 |
+
|
| 99 |
+
file_size = os.path.getsize(ckpt)
|
| 100 |
+
if file_size < 8:
|
| 101 |
+
raise _incomplete_safetensors_error("The safetensors header is incomplete.", ckpt)
|
| 102 |
+
|
| 103 |
+
file_lock = threading.Lock()
|
| 104 |
+
model_mmap = comfy_aimdo.model_mmap.ModelMMAP(ckpt)
|
| 105 |
+
f = model_mmap.get_file_handle()
|
| 106 |
+
mv = memoryview((ctypes.c_uint8 * file_size).from_address(model_mmap.get()))
|
| 107 |
+
|
| 108 |
+
header_size = struct.unpack("<Q", mv[:8])[0]
|
| 109 |
+
if header_size > _SAFETENSORS_MAX_HEADER_SIZE:
|
| 110 |
+
raise _invalid_safetensors_error("The safetensors header is too large.", ckpt)
|
| 111 |
+
|
| 112 |
+
data_base_offset = 8 + header_size
|
| 113 |
+
if data_base_offset > file_size:
|
| 114 |
+
raise _incomplete_safetensors_error("The safetensors header is incomplete.", ckpt)
|
| 115 |
+
|
| 116 |
+
try:
|
| 117 |
+
header = json.loads(mv[8:data_base_offset].tobytes().decode("utf-8"))
|
| 118 |
+
except (UnicodeDecodeError, json.JSONDecodeError) as e:
|
| 119 |
+
raise _invalid_safetensors_error(str(e), ckpt) from e
|
| 120 |
+
|
| 121 |
+
if not isinstance(header, dict):
|
| 122 |
+
raise _invalid_safetensors_error("The safetensors header is invalid.", ckpt)
|
| 123 |
+
|
| 124 |
+
mv = mv[data_base_offset:]
|
| 125 |
+
data_size = len(mv)
|
| 126 |
+
|
| 127 |
+
sd = {}
|
| 128 |
+
for name, info in header.items():
|
| 129 |
+
if name == "__metadata__":
|
| 130 |
+
continue
|
| 131 |
+
|
| 132 |
+
start, end = info["data_offsets"]
|
| 133 |
+
dtype = _TYPES[info["dtype"]]
|
| 134 |
+
if start < 0 or end < start:
|
| 135 |
+
raise _invalid_safetensors_error("Tensor '{}' has invalid data offsets.".format(name), ckpt)
|
| 136 |
+
if end > data_size:
|
| 137 |
+
raise _incomplete_safetensors_error("Tensor '{}' extends past the end of the file.".format(name), ckpt)
|
| 138 |
+
if math.prod(info["shape"]) * dtype.itemsize != end - start:
|
| 139 |
+
raise _invalid_safetensors_error("Tensor '{}' does not match its declared shape and dtype.".format(name), ckpt)
|
| 140 |
+
|
| 141 |
+
if start == end:
|
| 142 |
+
sd[name] = torch.empty(info["shape"], dtype=dtype)
|
| 143 |
+
else:
|
| 144 |
+
with warnings.catch_warnings():
|
| 145 |
+
#We are working with read-only RAM by design
|
| 146 |
+
warnings.filterwarnings("ignore", message="The given buffer is not writable")
|
| 147 |
+
tensor = torch.frombuffer(mv[start:end], dtype=dtype).view(info["shape"])
|
| 148 |
+
storage = tensor.untyped_storage()
|
| 149 |
+
setattr(storage,
|
| 150 |
+
"_comfy_tensor_file_slice",
|
| 151 |
+
comfy.memory_management.TensorFileSlice(f, file_lock, data_base_offset + start, end - start))
|
| 152 |
+
setattr(storage, "_comfy_tensor_mmap_refs", (model_mmap, mv))
|
| 153 |
+
sd[name] = tensor
|
| 154 |
+
|
| 155 |
+
return sd, header.get("__metadata__", {}),
|
| 156 |
+
|
| 157 |
+
|
| 158 |
+
def load_torch_file(ckpt, safe_load=False, device=None, return_metadata=False):
|
| 159 |
+
if device is None:
|
| 160 |
+
device = torch.device("cpu")
|
| 161 |
+
metadata = None
|
| 162 |
+
if ckpt.lower().endswith(".safetensors") or ckpt.lower().endswith(".sft"):
|
| 163 |
+
try:
|
| 164 |
+
if comfy.memory_management.aimdo_enabled:
|
| 165 |
+
sd, metadata = load_safetensors(ckpt)
|
| 166 |
+
if not return_metadata:
|
| 167 |
+
metadata = None
|
| 168 |
+
else:
|
| 169 |
+
with safetensors.safe_open(ckpt, framework="pt", device=device.type) as f:
|
| 170 |
+
sd = {}
|
| 171 |
+
for k in f.keys():
|
| 172 |
+
tensor = f.get_tensor(k)
|
| 173 |
+
if DISABLE_MMAP: # TODO: Not sure if this is the best way to bypass the mmap issues
|
| 174 |
+
tensor = tensor.to(device=device, copy=True)
|
| 175 |
+
sd[k] = tensor
|
| 176 |
+
if return_metadata:
|
| 177 |
+
metadata = f.metadata()
|
| 178 |
+
except Exception as e:
|
| 179 |
+
if len(e.args) > 0:
|
| 180 |
+
message = e.args[0]
|
| 181 |
+
if "HeaderTooLarge" in message:
|
| 182 |
+
raise _invalid_safetensors_error(message, ckpt)
|
| 183 |
+
if "MetadataIncompleteBuffer" in message:
|
| 184 |
+
raise _incomplete_safetensors_error(message, ckpt)
|
| 185 |
+
raise e
|
| 186 |
+
else:
|
| 187 |
+
torch_args = {}
|
| 188 |
+
if MMAP_TORCH_FILES:
|
| 189 |
+
torch_args["mmap"] = True
|
| 190 |
+
|
| 191 |
+
pl_sd = torch.load(ckpt, map_location=device, weights_only=True, **torch_args)
|
| 192 |
+
|
| 193 |
+
if "state_dict" in pl_sd:
|
| 194 |
+
sd = pl_sd["state_dict"]
|
| 195 |
+
else:
|
| 196 |
+
if len(pl_sd) == 1:
|
| 197 |
+
key = list(pl_sd.keys())[0]
|
| 198 |
+
sd = pl_sd[key]
|
| 199 |
+
if not isinstance(sd, dict):
|
| 200 |
+
sd = pl_sd
|
| 201 |
+
else:
|
| 202 |
+
sd = pl_sd
|
| 203 |
+
return (sd, metadata) if return_metadata else sd
|
| 204 |
+
|
| 205 |
+
def save_torch_file(sd, ckpt, metadata=None):
|
| 206 |
+
if metadata is not None:
|
| 207 |
+
safetensors.torch.save_file(sd, ckpt, metadata=metadata)
|
| 208 |
+
else:
|
| 209 |
+
safetensors.torch.save_file(sd, ckpt)
|
| 210 |
+
|
| 211 |
+
def calculate_parameters(sd, prefix=""):
|
| 212 |
+
params = 0
|
| 213 |
+
for k in sd.keys():
|
| 214 |
+
if k.startswith(prefix):
|
| 215 |
+
w = sd[k]
|
| 216 |
+
params += w.nelement()
|
| 217 |
+
return params
|
| 218 |
+
|
| 219 |
+
def weight_dtype(sd, prefix=""):
|
| 220 |
+
dtypes = {}
|
| 221 |
+
for k in sd.keys():
|
| 222 |
+
if k.startswith(prefix):
|
| 223 |
+
w = sd[k]
|
| 224 |
+
dtypes[w.dtype] = dtypes.get(w.dtype, 0) + w.numel()
|
| 225 |
+
|
| 226 |
+
if len(dtypes) == 0:
|
| 227 |
+
return None
|
| 228 |
+
|
| 229 |
+
return max(dtypes, key=dtypes.get)
|
| 230 |
+
|
| 231 |
+
def state_dict_key_replace(state_dict, keys_to_replace):
|
| 232 |
+
for x in keys_to_replace:
|
| 233 |
+
if x in state_dict:
|
| 234 |
+
state_dict[keys_to_replace[x]] = state_dict.pop(x)
|
| 235 |
+
return state_dict
|
| 236 |
+
|
| 237 |
+
def state_dict_prefix_replace(state_dict, replace_prefix, filter_keys=False):
|
| 238 |
+
if filter_keys:
|
| 239 |
+
out = {}
|
| 240 |
+
else:
|
| 241 |
+
out = state_dict
|
| 242 |
+
for rp in replace_prefix:
|
| 243 |
+
replace = list(map(lambda a: (a, "{}{}".format(replace_prefix[rp], a[len(rp):])), filter(lambda a: a.startswith(rp), state_dict.keys())))
|
| 244 |
+
for x in replace:
|
| 245 |
+
w = state_dict.pop(x[0])
|
| 246 |
+
out[x[1]] = w
|
| 247 |
+
return out
|
| 248 |
+
|
| 249 |
+
|
| 250 |
+
def transformers_convert(sd, prefix_from, prefix_to, number):
|
| 251 |
+
keys_to_replace = {
|
| 252 |
+
"{}positional_embedding": "{}embeddings.position_embedding.weight",
|
| 253 |
+
"{}token_embedding.weight": "{}embeddings.token_embedding.weight",
|
| 254 |
+
"{}ln_final.weight": "{}final_layer_norm.weight",
|
| 255 |
+
"{}ln_final.bias": "{}final_layer_norm.bias",
|
| 256 |
+
}
|
| 257 |
+
|
| 258 |
+
for k in keys_to_replace:
|
| 259 |
+
x = k.format(prefix_from)
|
| 260 |
+
if x in sd:
|
| 261 |
+
sd[keys_to_replace[k].format(prefix_to)] = sd.pop(x)
|
| 262 |
+
|
| 263 |
+
resblock_to_replace = {
|
| 264 |
+
"ln_1": "layer_norm1",
|
| 265 |
+
"ln_2": "layer_norm2",
|
| 266 |
+
"mlp.c_fc": "mlp.fc1",
|
| 267 |
+
"mlp.c_proj": "mlp.fc2",
|
| 268 |
+
"attn.out_proj": "self_attn.out_proj",
|
| 269 |
+
}
|
| 270 |
+
|
| 271 |
+
for resblock in range(number):
|
| 272 |
+
for x in resblock_to_replace:
|
| 273 |
+
for y in ["weight", "bias"]:
|
| 274 |
+
k = "{}transformer.resblocks.{}.{}.{}".format(prefix_from, resblock, x, y)
|
| 275 |
+
k_to = "{}encoder.layers.{}.{}.{}".format(prefix_to, resblock, resblock_to_replace[x], y)
|
| 276 |
+
if k in sd:
|
| 277 |
+
sd[k_to] = sd.pop(k)
|
| 278 |
+
|
| 279 |
+
for y in ["weight", "bias"]:
|
| 280 |
+
k_from = "{}transformer.resblocks.{}.attn.in_proj_{}".format(prefix_from, resblock, y)
|
| 281 |
+
if k_from in sd:
|
| 282 |
+
weights = sd.pop(k_from)
|
| 283 |
+
shape_from = weights.shape[0] // 3
|
| 284 |
+
for x in range(3):
|
| 285 |
+
p = ["self_attn.q_proj", "self_attn.k_proj", "self_attn.v_proj"]
|
| 286 |
+
k_to = "{}encoder.layers.{}.{}.{}".format(prefix_to, resblock, p[x], y)
|
| 287 |
+
sd[k_to] = weights[shape_from*x:shape_from*(x + 1)]
|
| 288 |
+
|
| 289 |
+
return sd
|
| 290 |
+
|
| 291 |
+
def clip_text_transformers_convert(sd, prefix_from, prefix_to):
|
| 292 |
+
sd = transformers_convert(sd, prefix_from, "{}text_model.".format(prefix_to), 32)
|
| 293 |
+
|
| 294 |
+
tp = "{}text_projection.weight".format(prefix_from)
|
| 295 |
+
if tp in sd:
|
| 296 |
+
sd["{}text_projection.weight".format(prefix_to)] = sd.pop(tp)
|
| 297 |
+
|
| 298 |
+
tp = "{}text_projection".format(prefix_from)
|
| 299 |
+
if tp in sd:
|
| 300 |
+
sd["{}text_projection.weight".format(prefix_to)] = sd.pop(tp).transpose(0, 1).contiguous()
|
| 301 |
+
return sd
|
| 302 |
+
|
| 303 |
+
|
| 304 |
+
UNET_MAP_ATTENTIONS = {
|
| 305 |
+
"proj_in.weight",
|
| 306 |
+
"proj_in.bias",
|
| 307 |
+
"proj_out.weight",
|
| 308 |
+
"proj_out.bias",
|
| 309 |
+
"norm.weight",
|
| 310 |
+
"norm.bias",
|
| 311 |
+
}
|
| 312 |
+
|
| 313 |
+
TRANSFORMER_BLOCKS = {
|
| 314 |
+
"norm1.weight",
|
| 315 |
+
"norm1.bias",
|
| 316 |
+
"norm2.weight",
|
| 317 |
+
"norm2.bias",
|
| 318 |
+
"norm3.weight",
|
| 319 |
+
"norm3.bias",
|
| 320 |
+
"attn1.to_q.weight",
|
| 321 |
+
"attn1.to_k.weight",
|
| 322 |
+
"attn1.to_v.weight",
|
| 323 |
+
"attn1.to_out.0.weight",
|
| 324 |
+
"attn1.to_out.0.bias",
|
| 325 |
+
"attn2.to_q.weight",
|
| 326 |
+
"attn2.to_k.weight",
|
| 327 |
+
"attn2.to_v.weight",
|
| 328 |
+
"attn2.to_out.0.weight",
|
| 329 |
+
"attn2.to_out.0.bias",
|
| 330 |
+
"ff.net.0.proj.weight",
|
| 331 |
+
"ff.net.0.proj.bias",
|
| 332 |
+
"ff.net.2.weight",
|
| 333 |
+
"ff.net.2.bias",
|
| 334 |
+
}
|
| 335 |
+
|
| 336 |
+
UNET_MAP_RESNET = {
|
| 337 |
+
"in_layers.2.weight": "conv1.weight",
|
| 338 |
+
"in_layers.2.bias": "conv1.bias",
|
| 339 |
+
"emb_layers.1.weight": "time_emb_proj.weight",
|
| 340 |
+
"emb_layers.1.bias": "time_emb_proj.bias",
|
| 341 |
+
"out_layers.3.weight": "conv2.weight",
|
| 342 |
+
"out_layers.3.bias": "conv2.bias",
|
| 343 |
+
"skip_connection.weight": "conv_shortcut.weight",
|
| 344 |
+
"skip_connection.bias": "conv_shortcut.bias",
|
| 345 |
+
"in_layers.0.weight": "norm1.weight",
|
| 346 |
+
"in_layers.0.bias": "norm1.bias",
|
| 347 |
+
"out_layers.0.weight": "norm2.weight",
|
| 348 |
+
"out_layers.0.bias": "norm2.bias",
|
| 349 |
+
}
|
| 350 |
+
|
| 351 |
+
UNET_MAP_BASIC = {
|
| 352 |
+
("label_emb.0.0.weight", "class_embedding.linear_1.weight"),
|
| 353 |
+
("label_emb.0.0.bias", "class_embedding.linear_1.bias"),
|
| 354 |
+
("label_emb.0.2.weight", "class_embedding.linear_2.weight"),
|
| 355 |
+
("label_emb.0.2.bias", "class_embedding.linear_2.bias"),
|
| 356 |
+
("label_emb.0.0.weight", "add_embedding.linear_1.weight"),
|
| 357 |
+
("label_emb.0.0.bias", "add_embedding.linear_1.bias"),
|
| 358 |
+
("label_emb.0.2.weight", "add_embedding.linear_2.weight"),
|
| 359 |
+
("label_emb.0.2.bias", "add_embedding.linear_2.bias"),
|
| 360 |
+
("input_blocks.0.0.weight", "conv_in.weight"),
|
| 361 |
+
("input_blocks.0.0.bias", "conv_in.bias"),
|
| 362 |
+
("out.0.weight", "conv_norm_out.weight"),
|
| 363 |
+
("out.0.bias", "conv_norm_out.bias"),
|
| 364 |
+
("out.2.weight", "conv_out.weight"),
|
| 365 |
+
("out.2.bias", "conv_out.bias"),
|
| 366 |
+
("time_embed.0.weight", "time_embedding.linear_1.weight"),
|
| 367 |
+
("time_embed.0.bias", "time_embedding.linear_1.bias"),
|
| 368 |
+
("time_embed.2.weight", "time_embedding.linear_2.weight"),
|
| 369 |
+
("time_embed.2.bias", "time_embedding.linear_2.bias")
|
| 370 |
+
}
|
| 371 |
+
|
| 372 |
+
def unet_to_diffusers(unet_config):
|
| 373 |
+
if "num_res_blocks" not in unet_config:
|
| 374 |
+
return {}
|
| 375 |
+
num_res_blocks = unet_config["num_res_blocks"]
|
| 376 |
+
channel_mult = unet_config["channel_mult"]
|
| 377 |
+
transformer_depth = unet_config["transformer_depth"][:]
|
| 378 |
+
transformer_depth_output = unet_config["transformer_depth_output"][:]
|
| 379 |
+
num_blocks = len(channel_mult)
|
| 380 |
+
|
| 381 |
+
transformers_mid = unet_config.get("transformer_depth_middle", None)
|
| 382 |
+
|
| 383 |
+
diffusers_unet_map = {}
|
| 384 |
+
for x in range(num_blocks):
|
| 385 |
+
n = 1 + (num_res_blocks[x] + 1) * x
|
| 386 |
+
for i in range(num_res_blocks[x]):
|
| 387 |
+
for b in UNET_MAP_RESNET:
|
| 388 |
+
diffusers_unet_map["down_blocks.{}.resnets.{}.{}".format(x, i, UNET_MAP_RESNET[b])] = "input_blocks.{}.0.{}".format(n, b)
|
| 389 |
+
num_transformers = transformer_depth.pop(0)
|
| 390 |
+
if num_transformers > 0:
|
| 391 |
+
for b in UNET_MAP_ATTENTIONS:
|
| 392 |
+
diffusers_unet_map["down_blocks.{}.attentions.{}.{}".format(x, i, b)] = "input_blocks.{}.1.{}".format(n, b)
|
| 393 |
+
for t in range(num_transformers):
|
| 394 |
+
for b in TRANSFORMER_BLOCKS:
|
| 395 |
+
diffusers_unet_map["down_blocks.{}.attentions.{}.transformer_blocks.{}.{}".format(x, i, t, b)] = "input_blocks.{}.1.transformer_blocks.{}.{}".format(n, t, b)
|
| 396 |
+
n += 1
|
| 397 |
+
for k in ["weight", "bias"]:
|
| 398 |
+
diffusers_unet_map["down_blocks.{}.downsamplers.0.conv.{}".format(x, k)] = "input_blocks.{}.0.op.{}".format(n, k)
|
| 399 |
+
|
| 400 |
+
i = 0
|
| 401 |
+
for b in UNET_MAP_ATTENTIONS:
|
| 402 |
+
diffusers_unet_map["mid_block.attentions.{}.{}".format(i, b)] = "middle_block.1.{}".format(b)
|
| 403 |
+
for t in range(transformers_mid):
|
| 404 |
+
for b in TRANSFORMER_BLOCKS:
|
| 405 |
+
diffusers_unet_map["mid_block.attentions.{}.transformer_blocks.{}.{}".format(i, t, b)] = "middle_block.1.transformer_blocks.{}.{}".format(t, b)
|
| 406 |
+
|
| 407 |
+
for i, n in enumerate([0, 2]):
|
| 408 |
+
for b in UNET_MAP_RESNET:
|
| 409 |
+
diffusers_unet_map["mid_block.resnets.{}.{}".format(i, UNET_MAP_RESNET[b])] = "middle_block.{}.{}".format(n, b)
|
| 410 |
+
|
| 411 |
+
num_res_blocks = list(reversed(num_res_blocks))
|
| 412 |
+
for x in range(num_blocks):
|
| 413 |
+
n = (num_res_blocks[x] + 1) * x
|
| 414 |
+
l = num_res_blocks[x] + 1
|
| 415 |
+
for i in range(l):
|
| 416 |
+
c = 0
|
| 417 |
+
for b in UNET_MAP_RESNET:
|
| 418 |
+
diffusers_unet_map["up_blocks.{}.resnets.{}.{}".format(x, i, UNET_MAP_RESNET[b])] = "output_blocks.{}.0.{}".format(n, b)
|
| 419 |
+
c += 1
|
| 420 |
+
num_transformers = transformer_depth_output.pop()
|
| 421 |
+
if num_transformers > 0:
|
| 422 |
+
c += 1
|
| 423 |
+
for b in UNET_MAP_ATTENTIONS:
|
| 424 |
+
diffusers_unet_map["up_blocks.{}.attentions.{}.{}".format(x, i, b)] = "output_blocks.{}.1.{}".format(n, b)
|
| 425 |
+
for t in range(num_transformers):
|
| 426 |
+
for b in TRANSFORMER_BLOCKS:
|
| 427 |
+
diffusers_unet_map["up_blocks.{}.attentions.{}.transformer_blocks.{}.{}".format(x, i, t, b)] = "output_blocks.{}.1.transformer_blocks.{}.{}".format(n, t, b)
|
| 428 |
+
if i == l - 1:
|
| 429 |
+
for k in ["weight", "bias"]:
|
| 430 |
+
diffusers_unet_map["up_blocks.{}.upsamplers.0.conv.{}".format(x, k)] = "output_blocks.{}.{}.conv.{}".format(n, c, k)
|
| 431 |
+
n += 1
|
| 432 |
+
|
| 433 |
+
for k in UNET_MAP_BASIC:
|
| 434 |
+
diffusers_unet_map[k[1]] = k[0]
|
| 435 |
+
|
| 436 |
+
return diffusers_unet_map
|
| 437 |
+
|
| 438 |
+
def swap_scale_shift(weight):
|
| 439 |
+
shift, scale = weight.chunk(2, dim=0)
|
| 440 |
+
new_weight = torch.cat([scale, shift], dim=0)
|
| 441 |
+
return new_weight
|
| 442 |
+
|
| 443 |
+
MMDIT_MAP_BASIC = {
|
| 444 |
+
("context_embedder.bias", "context_embedder.bias"),
|
| 445 |
+
("context_embedder.weight", "context_embedder.weight"),
|
| 446 |
+
("t_embedder.mlp.0.bias", "time_text_embed.timestep_embedder.linear_1.bias"),
|
| 447 |
+
("t_embedder.mlp.0.weight", "time_text_embed.timestep_embedder.linear_1.weight"),
|
| 448 |
+
("t_embedder.mlp.2.bias", "time_text_embed.timestep_embedder.linear_2.bias"),
|
| 449 |
+
("t_embedder.mlp.2.weight", "time_text_embed.timestep_embedder.linear_2.weight"),
|
| 450 |
+
("x_embedder.proj.bias", "pos_embed.proj.bias"),
|
| 451 |
+
("x_embedder.proj.weight", "pos_embed.proj.weight"),
|
| 452 |
+
("y_embedder.mlp.0.bias", "time_text_embed.text_embedder.linear_1.bias"),
|
| 453 |
+
("y_embedder.mlp.0.weight", "time_text_embed.text_embedder.linear_1.weight"),
|
| 454 |
+
("y_embedder.mlp.2.bias", "time_text_embed.text_embedder.linear_2.bias"),
|
| 455 |
+
("y_embedder.mlp.2.weight", "time_text_embed.text_embedder.linear_2.weight"),
|
| 456 |
+
("pos_embed", "pos_embed.pos_embed"),
|
| 457 |
+
("final_layer.adaLN_modulation.1.bias", "norm_out.linear.bias", swap_scale_shift),
|
| 458 |
+
("final_layer.adaLN_modulation.1.weight", "norm_out.linear.weight", swap_scale_shift),
|
| 459 |
+
("final_layer.linear.bias", "proj_out.bias"),
|
| 460 |
+
("final_layer.linear.weight", "proj_out.weight"),
|
| 461 |
+
}
|
| 462 |
+
|
| 463 |
+
MMDIT_MAP_BLOCK = {
|
| 464 |
+
("context_block.adaLN_modulation.1.bias", "norm1_context.linear.bias"),
|
| 465 |
+
("context_block.adaLN_modulation.1.weight", "norm1_context.linear.weight"),
|
| 466 |
+
("context_block.attn.proj.bias", "attn.to_add_out.bias"),
|
| 467 |
+
("context_block.attn.proj.weight", "attn.to_add_out.weight"),
|
| 468 |
+
("context_block.mlp.fc1.bias", "ff_context.net.0.proj.bias"),
|
| 469 |
+
("context_block.mlp.fc1.weight", "ff_context.net.0.proj.weight"),
|
| 470 |
+
("context_block.mlp.fc2.bias", "ff_context.net.2.bias"),
|
| 471 |
+
("context_block.mlp.fc2.weight", "ff_context.net.2.weight"),
|
| 472 |
+
("context_block.attn.ln_q.weight", "attn.norm_added_q.weight"),
|
| 473 |
+
("context_block.attn.ln_k.weight", "attn.norm_added_k.weight"),
|
| 474 |
+
("x_block.adaLN_modulation.1.bias", "norm1.linear.bias"),
|
| 475 |
+
("x_block.adaLN_modulation.1.weight", "norm1.linear.weight"),
|
| 476 |
+
("x_block.attn.proj.bias", "attn.to_out.0.bias"),
|
| 477 |
+
("x_block.attn.proj.weight", "attn.to_out.0.weight"),
|
| 478 |
+
("x_block.attn.ln_q.weight", "attn.norm_q.weight"),
|
| 479 |
+
("x_block.attn.ln_k.weight", "attn.norm_k.weight"),
|
| 480 |
+
("x_block.attn2.proj.bias", "attn2.to_out.0.bias"),
|
| 481 |
+
("x_block.attn2.proj.weight", "attn2.to_out.0.weight"),
|
| 482 |
+
("x_block.attn2.ln_q.weight", "attn2.norm_q.weight"),
|
| 483 |
+
("x_block.attn2.ln_k.weight", "attn2.norm_k.weight"),
|
| 484 |
+
("x_block.mlp.fc1.bias", "ff.net.0.proj.bias"),
|
| 485 |
+
("x_block.mlp.fc1.weight", "ff.net.0.proj.weight"),
|
| 486 |
+
("x_block.mlp.fc2.bias", "ff.net.2.bias"),
|
| 487 |
+
("x_block.mlp.fc2.weight", "ff.net.2.weight"),
|
| 488 |
+
}
|
| 489 |
+
|
| 490 |
+
def mmdit_to_diffusers(mmdit_config, output_prefix=""):
|
| 491 |
+
key_map = {}
|
| 492 |
+
|
| 493 |
+
depth = mmdit_config.get("depth", 0)
|
| 494 |
+
num_blocks = mmdit_config.get("num_blocks", depth)
|
| 495 |
+
for i in range(num_blocks):
|
| 496 |
+
block_from = "transformer_blocks.{}".format(i)
|
| 497 |
+
block_to = "{}joint_blocks.{}".format(output_prefix, i)
|
| 498 |
+
|
| 499 |
+
offset = depth * 64
|
| 500 |
+
|
| 501 |
+
for end in ("weight", "bias"):
|
| 502 |
+
k = "{}.attn.".format(block_from)
|
| 503 |
+
qkv = "{}.x_block.attn.qkv.{}".format(block_to, end)
|
| 504 |
+
key_map["{}to_q.{}".format(k, end)] = (qkv, (0, 0, offset))
|
| 505 |
+
key_map["{}to_k.{}".format(k, end)] = (qkv, (0, offset, offset))
|
| 506 |
+
key_map["{}to_v.{}".format(k, end)] = (qkv, (0, offset * 2, offset))
|
| 507 |
+
|
| 508 |
+
qkv = "{}.context_block.attn.qkv.{}".format(block_to, end)
|
| 509 |
+
key_map["{}add_q_proj.{}".format(k, end)] = (qkv, (0, 0, offset))
|
| 510 |
+
key_map["{}add_k_proj.{}".format(k, end)] = (qkv, (0, offset, offset))
|
| 511 |
+
key_map["{}add_v_proj.{}".format(k, end)] = (qkv, (0, offset * 2, offset))
|
| 512 |
+
|
| 513 |
+
k = "{}.attn2.".format(block_from)
|
| 514 |
+
qkv = "{}.x_block.attn2.qkv.{}".format(block_to, end)
|
| 515 |
+
key_map["{}to_q.{}".format(k, end)] = (qkv, (0, 0, offset))
|
| 516 |
+
key_map["{}to_k.{}".format(k, end)] = (qkv, (0, offset, offset))
|
| 517 |
+
key_map["{}to_v.{}".format(k, end)] = (qkv, (0, offset * 2, offset))
|
| 518 |
+
|
| 519 |
+
for k in MMDIT_MAP_BLOCK:
|
| 520 |
+
key_map["{}.{}".format(block_from, k[1])] = "{}.{}".format(block_to, k[0])
|
| 521 |
+
|
| 522 |
+
map_basic = MMDIT_MAP_BASIC.copy()
|
| 523 |
+
map_basic.add(("joint_blocks.{}.context_block.adaLN_modulation.1.bias".format(depth - 1), "transformer_blocks.{}.norm1_context.linear.bias".format(depth - 1), swap_scale_shift))
|
| 524 |
+
map_basic.add(("joint_blocks.{}.context_block.adaLN_modulation.1.weight".format(depth - 1), "transformer_blocks.{}.norm1_context.linear.weight".format(depth - 1), swap_scale_shift))
|
| 525 |
+
|
| 526 |
+
for k in map_basic:
|
| 527 |
+
if len(k) > 2:
|
| 528 |
+
key_map[k[1]] = ("{}{}".format(output_prefix, k[0]), None, k[2])
|
| 529 |
+
else:
|
| 530 |
+
key_map[k[1]] = "{}{}".format(output_prefix, k[0])
|
| 531 |
+
|
| 532 |
+
return key_map
|
| 533 |
+
|
| 534 |
+
PIXART_MAP_BASIC = {
|
| 535 |
+
("csize_embedder.mlp.0.weight", "adaln_single.emb.resolution_embedder.linear_1.weight"),
|
| 536 |
+
("csize_embedder.mlp.0.bias", "adaln_single.emb.resolution_embedder.linear_1.bias"),
|
| 537 |
+
("csize_embedder.mlp.2.weight", "adaln_single.emb.resolution_embedder.linear_2.weight"),
|
| 538 |
+
("csize_embedder.mlp.2.bias", "adaln_single.emb.resolution_embedder.linear_2.bias"),
|
| 539 |
+
("ar_embedder.mlp.0.weight", "adaln_single.emb.aspect_ratio_embedder.linear_1.weight"),
|
| 540 |
+
("ar_embedder.mlp.0.bias", "adaln_single.emb.aspect_ratio_embedder.linear_1.bias"),
|
| 541 |
+
("ar_embedder.mlp.2.weight", "adaln_single.emb.aspect_ratio_embedder.linear_2.weight"),
|
| 542 |
+
("ar_embedder.mlp.2.bias", "adaln_single.emb.aspect_ratio_embedder.linear_2.bias"),
|
| 543 |
+
("x_embedder.proj.weight", "pos_embed.proj.weight"),
|
| 544 |
+
("x_embedder.proj.bias", "pos_embed.proj.bias"),
|
| 545 |
+
("y_embedder.y_embedding", "caption_projection.y_embedding"),
|
| 546 |
+
("y_embedder.y_proj.fc1.weight", "caption_projection.linear_1.weight"),
|
| 547 |
+
("y_embedder.y_proj.fc1.bias", "caption_projection.linear_1.bias"),
|
| 548 |
+
("y_embedder.y_proj.fc2.weight", "caption_projection.linear_2.weight"),
|
| 549 |
+
("y_embedder.y_proj.fc2.bias", "caption_projection.linear_2.bias"),
|
| 550 |
+
("t_embedder.mlp.0.weight", "adaln_single.emb.timestep_embedder.linear_1.weight"),
|
| 551 |
+
("t_embedder.mlp.0.bias", "adaln_single.emb.timestep_embedder.linear_1.bias"),
|
| 552 |
+
("t_embedder.mlp.2.weight", "adaln_single.emb.timestep_embedder.linear_2.weight"),
|
| 553 |
+
("t_embedder.mlp.2.bias", "adaln_single.emb.timestep_embedder.linear_2.bias"),
|
| 554 |
+
("t_block.1.weight", "adaln_single.linear.weight"),
|
| 555 |
+
("t_block.1.bias", "adaln_single.linear.bias"),
|
| 556 |
+
("final_layer.linear.weight", "proj_out.weight"),
|
| 557 |
+
("final_layer.linear.bias", "proj_out.bias"),
|
| 558 |
+
("final_layer.scale_shift_table", "scale_shift_table"),
|
| 559 |
+
}
|
| 560 |
+
|
| 561 |
+
PIXART_MAP_BLOCK = {
|
| 562 |
+
("scale_shift_table", "scale_shift_table"),
|
| 563 |
+
("attn.proj.weight", "attn1.to_out.0.weight"),
|
| 564 |
+
("attn.proj.bias", "attn1.to_out.0.bias"),
|
| 565 |
+
("mlp.fc1.weight", "ff.net.0.proj.weight"),
|
| 566 |
+
("mlp.fc1.bias", "ff.net.0.proj.bias"),
|
| 567 |
+
("mlp.fc2.weight", "ff.net.2.weight"),
|
| 568 |
+
("mlp.fc2.bias", "ff.net.2.bias"),
|
| 569 |
+
("cross_attn.proj.weight" ,"attn2.to_out.0.weight"),
|
| 570 |
+
("cross_attn.proj.bias" ,"attn2.to_out.0.bias"),
|
| 571 |
+
}
|
| 572 |
+
|
| 573 |
+
def pixart_to_diffusers(mmdit_config, output_prefix=""):
|
| 574 |
+
key_map = {}
|
| 575 |
+
|
| 576 |
+
depth = mmdit_config.get("depth", 0)
|
| 577 |
+
offset = mmdit_config.get("hidden_size", 1152)
|
| 578 |
+
|
| 579 |
+
for i in range(depth):
|
| 580 |
+
block_from = "transformer_blocks.{}".format(i)
|
| 581 |
+
block_to = "{}blocks.{}".format(output_prefix, i)
|
| 582 |
+
|
| 583 |
+
for end in ("weight", "bias"):
|
| 584 |
+
s = "{}.attn1.".format(block_from)
|
| 585 |
+
qkv = "{}.attn.qkv.{}".format(block_to, end)
|
| 586 |
+
key_map["{}to_q.{}".format(s, end)] = (qkv, (0, 0, offset))
|
| 587 |
+
key_map["{}to_k.{}".format(s, end)] = (qkv, (0, offset, offset))
|
| 588 |
+
key_map["{}to_v.{}".format(s, end)] = (qkv, (0, offset * 2, offset))
|
| 589 |
+
|
| 590 |
+
s = "{}.attn2.".format(block_from)
|
| 591 |
+
q = "{}.cross_attn.q_linear.{}".format(block_to, end)
|
| 592 |
+
kv = "{}.cross_attn.kv_linear.{}".format(block_to, end)
|
| 593 |
+
|
| 594 |
+
key_map["{}to_q.{}".format(s, end)] = q
|
| 595 |
+
key_map["{}to_k.{}".format(s, end)] = (kv, (0, 0, offset))
|
| 596 |
+
key_map["{}to_v.{}".format(s, end)] = (kv, (0, offset, offset))
|
| 597 |
+
|
| 598 |
+
for k in PIXART_MAP_BLOCK:
|
| 599 |
+
key_map["{}.{}".format(block_from, k[1])] = "{}.{}".format(block_to, k[0])
|
| 600 |
+
|
| 601 |
+
for k in PIXART_MAP_BASIC:
|
| 602 |
+
key_map[k[1]] = "{}{}".format(output_prefix, k[0])
|
| 603 |
+
|
| 604 |
+
return key_map
|
| 605 |
+
|
| 606 |
+
def auraflow_to_diffusers(mmdit_config, output_prefix=""):
|
| 607 |
+
n_double_layers = mmdit_config.get("n_double_layers", 0)
|
| 608 |
+
n_layers = mmdit_config.get("n_layers", 0)
|
| 609 |
+
|
| 610 |
+
key_map = {}
|
| 611 |
+
for i in range(n_layers):
|
| 612 |
+
if i < n_double_layers:
|
| 613 |
+
index = i
|
| 614 |
+
prefix_from = "joint_transformer_blocks"
|
| 615 |
+
prefix_to = "{}double_layers".format(output_prefix)
|
| 616 |
+
block_map = {
|
| 617 |
+
"attn.to_q.weight": "attn.w2q.weight",
|
| 618 |
+
"attn.to_k.weight": "attn.w2k.weight",
|
| 619 |
+
"attn.to_v.weight": "attn.w2v.weight",
|
| 620 |
+
"attn.to_out.0.weight": "attn.w2o.weight",
|
| 621 |
+
"attn.add_q_proj.weight": "attn.w1q.weight",
|
| 622 |
+
"attn.add_k_proj.weight": "attn.w1k.weight",
|
| 623 |
+
"attn.add_v_proj.weight": "attn.w1v.weight",
|
| 624 |
+
"attn.to_add_out.weight": "attn.w1o.weight",
|
| 625 |
+
"ff.linear_1.weight": "mlpX.c_fc1.weight",
|
| 626 |
+
"ff.linear_2.weight": "mlpX.c_fc2.weight",
|
| 627 |
+
"ff.out_projection.weight": "mlpX.c_proj.weight",
|
| 628 |
+
"ff_context.linear_1.weight": "mlpC.c_fc1.weight",
|
| 629 |
+
"ff_context.linear_2.weight": "mlpC.c_fc2.weight",
|
| 630 |
+
"ff_context.out_projection.weight": "mlpC.c_proj.weight",
|
| 631 |
+
"norm1.linear.weight": "modX.1.weight",
|
| 632 |
+
"norm1_context.linear.weight": "modC.1.weight",
|
| 633 |
+
}
|
| 634 |
+
else:
|
| 635 |
+
index = i - n_double_layers
|
| 636 |
+
prefix_from = "single_transformer_blocks"
|
| 637 |
+
prefix_to = "{}single_layers".format(output_prefix)
|
| 638 |
+
|
| 639 |
+
block_map = {
|
| 640 |
+
"attn.to_q.weight": "attn.w1q.weight",
|
| 641 |
+
"attn.to_k.weight": "attn.w1k.weight",
|
| 642 |
+
"attn.to_v.weight": "attn.w1v.weight",
|
| 643 |
+
"attn.to_out.0.weight": "attn.w1o.weight",
|
| 644 |
+
"norm1.linear.weight": "modCX.1.weight",
|
| 645 |
+
"ff.linear_1.weight": "mlp.c_fc1.weight",
|
| 646 |
+
"ff.linear_2.weight": "mlp.c_fc2.weight",
|
| 647 |
+
"ff.out_projection.weight": "mlp.c_proj.weight"
|
| 648 |
+
}
|
| 649 |
+
|
| 650 |
+
for k in block_map:
|
| 651 |
+
key_map["{}.{}.{}".format(prefix_from, index, k)] = "{}.{}.{}".format(prefix_to, index, block_map[k])
|
| 652 |
+
|
| 653 |
+
MAP_BASIC = {
|
| 654 |
+
("positional_encoding", "pos_embed.pos_embed"),
|
| 655 |
+
("register_tokens", "register_tokens"),
|
| 656 |
+
("t_embedder.mlp.0.weight", "time_step_proj.linear_1.weight"),
|
| 657 |
+
("t_embedder.mlp.0.bias", "time_step_proj.linear_1.bias"),
|
| 658 |
+
("t_embedder.mlp.2.weight", "time_step_proj.linear_2.weight"),
|
| 659 |
+
("t_embedder.mlp.2.bias", "time_step_proj.linear_2.bias"),
|
| 660 |
+
("cond_seq_linear.weight", "context_embedder.weight"),
|
| 661 |
+
("init_x_linear.weight", "pos_embed.proj.weight"),
|
| 662 |
+
("init_x_linear.bias", "pos_embed.proj.bias"),
|
| 663 |
+
("final_linear.weight", "proj_out.weight"),
|
| 664 |
+
("modF.1.weight", "norm_out.linear.weight", swap_scale_shift),
|
| 665 |
+
}
|
| 666 |
+
|
| 667 |
+
for k in MAP_BASIC:
|
| 668 |
+
if len(k) > 2:
|
| 669 |
+
key_map[k[1]] = ("{}{}".format(output_prefix, k[0]), None, k[2])
|
| 670 |
+
else:
|
| 671 |
+
key_map[k[1]] = "{}{}".format(output_prefix, k[0])
|
| 672 |
+
|
| 673 |
+
return key_map
|
| 674 |
+
|
| 675 |
+
def flux_to_diffusers(mmdit_config, output_prefix=""):
|
| 676 |
+
n_double_layers = mmdit_config.get("depth", 0)
|
| 677 |
+
n_single_layers = mmdit_config.get("depth_single_blocks", 0)
|
| 678 |
+
hidden_size = mmdit_config.get("hidden_size", 0)
|
| 679 |
+
|
| 680 |
+
key_map = {}
|
| 681 |
+
for index in range(n_double_layers):
|
| 682 |
+
prefix_from = "transformer_blocks.{}".format(index)
|
| 683 |
+
prefix_to = "{}double_blocks.{}".format(output_prefix, index)
|
| 684 |
+
|
| 685 |
+
for end in ("weight", "bias"):
|
| 686 |
+
k = "{}.attn.".format(prefix_from)
|
| 687 |
+
qkv = "{}.img_attn.qkv.{}".format(prefix_to, end)
|
| 688 |
+
key_map["{}to_q.{}".format(k, end)] = (qkv, (0, 0, hidden_size))
|
| 689 |
+
key_map["{}to_k.{}".format(k, end)] = (qkv, (0, hidden_size, hidden_size))
|
| 690 |
+
key_map["{}to_v.{}".format(k, end)] = (qkv, (0, hidden_size * 2, hidden_size))
|
| 691 |
+
|
| 692 |
+
k = "{}.attn.".format(prefix_from)
|
| 693 |
+
qkv = "{}.txt_attn.qkv.{}".format(prefix_to, end)
|
| 694 |
+
key_map["{}add_q_proj.{}".format(k, end)] = (qkv, (0, 0, hidden_size))
|
| 695 |
+
key_map["{}add_k_proj.{}".format(k, end)] = (qkv, (0, hidden_size, hidden_size))
|
| 696 |
+
key_map["{}add_v_proj.{}".format(k, end)] = (qkv, (0, hidden_size * 2, hidden_size))
|
| 697 |
+
|
| 698 |
+
block_map = {
|
| 699 |
+
"attn.to_out.0.weight": "img_attn.proj.weight",
|
| 700 |
+
"attn.to_out.0.bias": "img_attn.proj.bias",
|
| 701 |
+
"norm1.linear.weight": "img_mod.lin.weight",
|
| 702 |
+
"norm1.linear.bias": "img_mod.lin.bias",
|
| 703 |
+
"norm1_context.linear.weight": "txt_mod.lin.weight",
|
| 704 |
+
"norm1_context.linear.bias": "txt_mod.lin.bias",
|
| 705 |
+
"attn.to_add_out.weight": "txt_attn.proj.weight",
|
| 706 |
+
"attn.to_add_out.bias": "txt_attn.proj.bias",
|
| 707 |
+
"ff.net.0.proj.weight": "img_mlp.0.weight",
|
| 708 |
+
"ff.net.0.proj.bias": "img_mlp.0.bias",
|
| 709 |
+
"ff.net.2.weight": "img_mlp.2.weight",
|
| 710 |
+
"ff.net.2.bias": "img_mlp.2.bias",
|
| 711 |
+
"ff_context.net.0.proj.weight": "txt_mlp.0.weight",
|
| 712 |
+
"ff_context.net.0.proj.bias": "txt_mlp.0.bias",
|
| 713 |
+
"ff_context.net.2.weight": "txt_mlp.2.weight",
|
| 714 |
+
"ff_context.net.2.bias": "txt_mlp.2.bias",
|
| 715 |
+
"ff.linear_in.weight": "img_mlp.0.weight", # LyCoris LoKr
|
| 716 |
+
"ff.linear_in.bias": "img_mlp.0.bias",
|
| 717 |
+
"ff.linear_out.weight": "img_mlp.2.weight",
|
| 718 |
+
"ff.linear_out.bias": "img_mlp.2.bias",
|
| 719 |
+
"ff_context.linear_in.weight": "txt_mlp.0.weight",
|
| 720 |
+
"ff_context.linear_in.bias": "txt_mlp.0.bias",
|
| 721 |
+
"ff_context.linear_out.weight": "txt_mlp.2.weight",
|
| 722 |
+
"ff_context.linear_out.bias": "txt_mlp.2.bias",
|
| 723 |
+
"attn.norm_q.weight": "img_attn.norm.query_norm.weight",
|
| 724 |
+
"attn.norm_k.weight": "img_attn.norm.key_norm.weight",
|
| 725 |
+
"attn.norm_added_q.weight": "txt_attn.norm.query_norm.weight",
|
| 726 |
+
"attn.norm_added_k.weight": "txt_attn.norm.key_norm.weight",
|
| 727 |
+
}
|
| 728 |
+
|
| 729 |
+
for k in block_map:
|
| 730 |
+
key_map["{}.{}".format(prefix_from, k)] = "{}.{}".format(prefix_to, block_map[k])
|
| 731 |
+
|
| 732 |
+
for index in range(n_single_layers):
|
| 733 |
+
prefix_from = "single_transformer_blocks.{}".format(index)
|
| 734 |
+
prefix_to = "{}single_blocks.{}".format(output_prefix, index)
|
| 735 |
+
|
| 736 |
+
for end in ("weight", "bias"):
|
| 737 |
+
k = "{}.attn.".format(prefix_from)
|
| 738 |
+
qkv = "{}.linear1.{}".format(prefix_to, end)
|
| 739 |
+
key_map["{}to_q.{}".format(k, end)] = (qkv, (0, 0, hidden_size))
|
| 740 |
+
key_map["{}to_k.{}".format(k, end)] = (qkv, (0, hidden_size, hidden_size))
|
| 741 |
+
key_map["{}to_v.{}".format(k, end)] = (qkv, (0, hidden_size * 2, hidden_size))
|
| 742 |
+
key_map["{}.proj_mlp.{}".format(prefix_from, end)] = (qkv, (0, hidden_size * 3, hidden_size * 4))
|
| 743 |
+
|
| 744 |
+
block_map = {
|
| 745 |
+
"norm.linear.weight": "modulation.lin.weight",
|
| 746 |
+
"norm.linear.bias": "modulation.lin.bias",
|
| 747 |
+
"proj_out.weight": "linear2.weight",
|
| 748 |
+
"proj_out.bias": "linear2.bias",
|
| 749 |
+
"attn.norm_q.weight": "norm.query_norm.weight",
|
| 750 |
+
"attn.norm_k.weight": "norm.key_norm.weight",
|
| 751 |
+
"attn.to_qkv_mlp_proj.weight": "linear1.weight", # Flux 2
|
| 752 |
+
"attn.to_out.weight": "linear2.weight", # Flux 2
|
| 753 |
+
}
|
| 754 |
+
|
| 755 |
+
for k in block_map:
|
| 756 |
+
key_map["{}.{}".format(prefix_from, k)] = "{}.{}".format(prefix_to, block_map[k])
|
| 757 |
+
|
| 758 |
+
MAP_BASIC = {
|
| 759 |
+
("final_layer.linear.bias", "proj_out.bias"),
|
| 760 |
+
("final_layer.linear.weight", "proj_out.weight"),
|
| 761 |
+
("img_in.bias", "x_embedder.bias"),
|
| 762 |
+
("img_in.weight", "x_embedder.weight"),
|
| 763 |
+
("time_in.in_layer.bias", "time_text_embed.timestep_embedder.linear_1.bias"),
|
| 764 |
+
("time_in.in_layer.weight", "time_text_embed.timestep_embedder.linear_1.weight"),
|
| 765 |
+
("time_in.out_layer.bias", "time_text_embed.timestep_embedder.linear_2.bias"),
|
| 766 |
+
("time_in.out_layer.weight", "time_text_embed.timestep_embedder.linear_2.weight"),
|
| 767 |
+
("txt_in.bias", "context_embedder.bias"),
|
| 768 |
+
("txt_in.weight", "context_embedder.weight"),
|
| 769 |
+
("vector_in.in_layer.bias", "time_text_embed.text_embedder.linear_1.bias"),
|
| 770 |
+
("vector_in.in_layer.weight", "time_text_embed.text_embedder.linear_1.weight"),
|
| 771 |
+
("vector_in.out_layer.bias", "time_text_embed.text_embedder.linear_2.bias"),
|
| 772 |
+
("vector_in.out_layer.weight", "time_text_embed.text_embedder.linear_2.weight"),
|
| 773 |
+
("guidance_in.in_layer.bias", "time_text_embed.guidance_embedder.linear_1.bias"),
|
| 774 |
+
("guidance_in.in_layer.weight", "time_text_embed.guidance_embedder.linear_1.weight"),
|
| 775 |
+
("guidance_in.out_layer.bias", "time_text_embed.guidance_embedder.linear_2.bias"),
|
| 776 |
+
("guidance_in.out_layer.weight", "time_text_embed.guidance_embedder.linear_2.weight"),
|
| 777 |
+
("final_layer.adaLN_modulation.1.bias", "norm_out.linear.bias", swap_scale_shift),
|
| 778 |
+
("final_layer.adaLN_modulation.1.weight", "norm_out.linear.weight", swap_scale_shift),
|
| 779 |
+
("pos_embed_input.bias", "controlnet_x_embedder.bias"),
|
| 780 |
+
("pos_embed_input.weight", "controlnet_x_embedder.weight"),
|
| 781 |
+
}
|
| 782 |
+
|
| 783 |
+
for k in MAP_BASIC:
|
| 784 |
+
if len(k) > 2:
|
| 785 |
+
key_map[k[1]] = ("{}{}".format(output_prefix, k[0]), None, k[2])
|
| 786 |
+
else:
|
| 787 |
+
key_map[k[1]] = "{}{}".format(output_prefix, k[0])
|
| 788 |
+
|
| 789 |
+
return key_map
|
| 790 |
+
|
| 791 |
+
def z_image_to_diffusers(mmdit_config, output_prefix=""):
|
| 792 |
+
n_layers = mmdit_config.get("n_layers", 0)
|
| 793 |
+
hidden_size = mmdit_config.get("dim", 0)
|
| 794 |
+
n_context_refiner = mmdit_config.get("n_refiner_layers", 2)
|
| 795 |
+
n_noise_refiner = mmdit_config.get("n_refiner_layers", 2)
|
| 796 |
+
key_map = {}
|
| 797 |
+
|
| 798 |
+
def add_block_keys(prefix_from, prefix_to, has_adaln=True):
|
| 799 |
+
for end in ("weight", "bias"):
|
| 800 |
+
k = "{}.attention.".format(prefix_from)
|
| 801 |
+
qkv = "{}.attention.qkv.{}".format(prefix_to, end)
|
| 802 |
+
key_map["{}to_q.{}".format(k, end)] = (qkv, (0, 0, hidden_size))
|
| 803 |
+
key_map["{}to_k.{}".format(k, end)] = (qkv, (0, hidden_size, hidden_size))
|
| 804 |
+
key_map["{}to_v.{}".format(k, end)] = (qkv, (0, hidden_size * 2, hidden_size))
|
| 805 |
+
|
| 806 |
+
block_map = {
|
| 807 |
+
"attention.norm_q.weight": "attention.q_norm.weight",
|
| 808 |
+
"attention.norm_k.weight": "attention.k_norm.weight",
|
| 809 |
+
"attention.to_out.0.weight": "attention.out.weight",
|
| 810 |
+
"attention.to_out.0.bias": "attention.out.bias",
|
| 811 |
+
"attention_norm1.weight": "attention_norm1.weight",
|
| 812 |
+
"attention_norm2.weight": "attention_norm2.weight",
|
| 813 |
+
"feed_forward.w1.weight": "feed_forward.w1.weight",
|
| 814 |
+
"feed_forward.w2.weight": "feed_forward.w2.weight",
|
| 815 |
+
"feed_forward.w3.weight": "feed_forward.w3.weight",
|
| 816 |
+
"ffn_norm1.weight": "ffn_norm1.weight",
|
| 817 |
+
"ffn_norm2.weight": "ffn_norm2.weight",
|
| 818 |
+
}
|
| 819 |
+
if has_adaln:
|
| 820 |
+
block_map["adaLN_modulation.0.weight"] = "adaLN_modulation.0.weight"
|
| 821 |
+
block_map["adaLN_modulation.0.bias"] = "adaLN_modulation.0.bias"
|
| 822 |
+
for k, v in block_map.items():
|
| 823 |
+
key_map["{}.{}".format(prefix_from, k)] = "{}.{}".format(prefix_to, v)
|
| 824 |
+
|
| 825 |
+
for i in range(n_layers):
|
| 826 |
+
add_block_keys("layers.{}".format(i), "{}layers.{}".format(output_prefix, i))
|
| 827 |
+
|
| 828 |
+
for i in range(n_context_refiner):
|
| 829 |
+
add_block_keys("context_refiner.{}".format(i), "{}context_refiner.{}".format(output_prefix, i))
|
| 830 |
+
|
| 831 |
+
for i in range(n_noise_refiner):
|
| 832 |
+
add_block_keys("noise_refiner.{}".format(i), "{}noise_refiner.{}".format(output_prefix, i))
|
| 833 |
+
|
| 834 |
+
MAP_BASIC = [
|
| 835 |
+
("final_layer.linear.weight", "all_final_layer.2-1.linear.weight"),
|
| 836 |
+
("final_layer.linear.bias", "all_final_layer.2-1.linear.bias"),
|
| 837 |
+
("final_layer.adaLN_modulation.1.weight", "all_final_layer.2-1.adaLN_modulation.1.weight"),
|
| 838 |
+
("final_layer.adaLN_modulation.1.bias", "all_final_layer.2-1.adaLN_modulation.1.bias"),
|
| 839 |
+
("x_embedder.weight", "all_x_embedder.2-1.weight"),
|
| 840 |
+
("x_embedder.bias", "all_x_embedder.2-1.bias"),
|
| 841 |
+
("x_pad_token", "x_pad_token"),
|
| 842 |
+
("cap_embedder.0.weight", "cap_embedder.0.weight"),
|
| 843 |
+
("cap_embedder.1.weight", "cap_embedder.1.weight"),
|
| 844 |
+
("cap_embedder.1.bias", "cap_embedder.1.bias"),
|
| 845 |
+
("cap_pad_token", "cap_pad_token"),
|
| 846 |
+
("t_embedder.mlp.0.weight", "t_embedder.mlp.0.weight"),
|
| 847 |
+
("t_embedder.mlp.0.bias", "t_embedder.mlp.0.bias"),
|
| 848 |
+
("t_embedder.mlp.2.weight", "t_embedder.mlp.2.weight"),
|
| 849 |
+
("t_embedder.mlp.2.bias", "t_embedder.mlp.2.bias"),
|
| 850 |
+
]
|
| 851 |
+
|
| 852 |
+
for c, diffusers in MAP_BASIC:
|
| 853 |
+
key_map[diffusers] = "{}{}".format(output_prefix, c)
|
| 854 |
+
|
| 855 |
+
return key_map
|
| 856 |
+
|
| 857 |
+
def krea2_to_diffusers(mmdit_config, output_prefix=""):
|
| 858 |
+
n_layers = mmdit_config.get("layers", 0)
|
| 859 |
+
n_txt_layerwise = 2 # TextFusionTransformer hardcodes 2 layerwise + 2 refiner blocks
|
| 860 |
+
n_txt_refiner = 2
|
| 861 |
+
key_map = {}
|
| 862 |
+
|
| 863 |
+
def add_block(prefix_to, prefix_from):
|
| 864 |
+
block_map = {
|
| 865 |
+
"attn.to_q": "attn.wq", "attn.to_k": "attn.wk", "attn.to_v": "attn.wv",
|
| 866 |
+
"attn.to_gate": "attn.gate", "attn.to_out.0": "attn.wo",
|
| 867 |
+
"attn.to_out": "attn.wo", # some tools drop the ".0" on to_out
|
| 868 |
+
"ff.gate": "mlp.gate", "ff.up": "mlp.up", "ff.down": "mlp.down",
|
| 869 |
+
}
|
| 870 |
+
for d, c in block_map.items():
|
| 871 |
+
key_map["{}.{}.weight".format(prefix_to, d)] = "{}{}.{}.weight".format(output_prefix, prefix_from, c)
|
| 872 |
+
|
| 873 |
+
for i in range(n_layers):
|
| 874 |
+
add_block("transformer_blocks.{}".format(i), "blocks.{}".format(i))
|
| 875 |
+
for i in range(n_txt_layerwise):
|
| 876 |
+
add_block("text_fusion.layerwise_blocks.{}".format(i), "txtfusion.layerwise_blocks.{}".format(i))
|
| 877 |
+
for i in range(n_txt_refiner):
|
| 878 |
+
add_block("text_fusion.refiner_blocks.{}".format(i), "txtfusion.refiner_blocks.{}".format(i))
|
| 879 |
+
|
| 880 |
+
MAP_BASIC = [
|
| 881 |
+
("img_in", "first"),
|
| 882 |
+
("time_embed.linear_1", "tmlp.0"),
|
| 883 |
+
("time_embed.linear_2", "tmlp.2"),
|
| 884 |
+
("time_mod_proj", "tproj.1"),
|
| 885 |
+
("txt_in.linear_1", "txtmlp.1"),
|
| 886 |
+
("txt_in.linear_2", "txtmlp.3"),
|
| 887 |
+
("text_fusion.projector", "txtfusion.projector"),
|
| 888 |
+
("final_layer.linear", "last.linear"),
|
| 889 |
+
]
|
| 890 |
+
for d, c in MAP_BASIC:
|
| 891 |
+
key_map["{}.weight".format(d)] = "{}{}.weight".format(output_prefix, c)
|
| 892 |
+
|
| 893 |
+
return key_map
|
| 894 |
+
|
| 895 |
+
def repeat_to_batch_size(tensor, batch_size, dim=0):
|
| 896 |
+
if tensor.shape[dim] > batch_size:
|
| 897 |
+
return tensor.narrow(dim, 0, batch_size)
|
| 898 |
+
elif tensor.shape[dim] < batch_size:
|
| 899 |
+
return tensor.repeat(dim * [1] + [math.ceil(batch_size / tensor.shape[dim])] + [1] * (len(tensor.shape) - 1 - dim)).narrow(dim, 0, batch_size)
|
| 900 |
+
return tensor
|
| 901 |
+
|
| 902 |
+
def resize_to_batch_size(tensor, batch_size):
|
| 903 |
+
in_batch_size = tensor.shape[0]
|
| 904 |
+
if in_batch_size == batch_size:
|
| 905 |
+
return tensor
|
| 906 |
+
|
| 907 |
+
if batch_size <= 1:
|
| 908 |
+
return tensor[:batch_size]
|
| 909 |
+
|
| 910 |
+
output = torch.empty([batch_size] + list(tensor.shape)[1:], dtype=tensor.dtype, device=tensor.device)
|
| 911 |
+
if batch_size < in_batch_size:
|
| 912 |
+
scale = (in_batch_size - 1) / (batch_size - 1)
|
| 913 |
+
for i in range(batch_size):
|
| 914 |
+
output[i] = tensor[min(round(i * scale), in_batch_size - 1)]
|
| 915 |
+
else:
|
| 916 |
+
scale = in_batch_size / batch_size
|
| 917 |
+
for i in range(batch_size):
|
| 918 |
+
output[i] = tensor[min(math.floor((i + 0.5) * scale), in_batch_size - 1)]
|
| 919 |
+
|
| 920 |
+
return output
|
| 921 |
+
|
| 922 |
+
def resize_list_to_batch_size(l, batch_size):
|
| 923 |
+
in_batch_size = len(l)
|
| 924 |
+
if in_batch_size == batch_size or in_batch_size == 0:
|
| 925 |
+
return l
|
| 926 |
+
|
| 927 |
+
if batch_size <= 1:
|
| 928 |
+
return l[:batch_size]
|
| 929 |
+
|
| 930 |
+
output = []
|
| 931 |
+
if batch_size < in_batch_size:
|
| 932 |
+
scale = (in_batch_size - 1) / (batch_size - 1)
|
| 933 |
+
for i in range(batch_size):
|
| 934 |
+
output.append(l[min(round(i * scale), in_batch_size - 1)])
|
| 935 |
+
else:
|
| 936 |
+
scale = in_batch_size / batch_size
|
| 937 |
+
for i in range(batch_size):
|
| 938 |
+
output.append(l[min(math.floor((i + 0.5) * scale), in_batch_size - 1)])
|
| 939 |
+
|
| 940 |
+
return output
|
| 941 |
+
|
| 942 |
+
def convert_sd_to(state_dict, dtype):
|
| 943 |
+
keys = list(state_dict.keys())
|
| 944 |
+
for k in keys:
|
| 945 |
+
state_dict[k] = state_dict[k].to(dtype)
|
| 946 |
+
return state_dict
|
| 947 |
+
|
| 948 |
+
def safetensors_header(safetensors_path, max_size=100*1024*1024):
|
| 949 |
+
with open(safetensors_path, "rb") as f:
|
| 950 |
+
header = f.read(8)
|
| 951 |
+
length_of_header = struct.unpack('<Q', header)[0]
|
| 952 |
+
if length_of_header > max_size:
|
| 953 |
+
return None
|
| 954 |
+
return f.read(length_of_header)
|
| 955 |
+
|
| 956 |
+
ATTR_UNSET={}
|
| 957 |
+
|
| 958 |
+
def resolve_attr(obj, attr):
|
| 959 |
+
attrs = attr.split(".")
|
| 960 |
+
for name in attrs[:-1]:
|
| 961 |
+
obj = getattr(obj, name)
|
| 962 |
+
return obj, attrs[-1]
|
| 963 |
+
|
| 964 |
+
def set_attr(obj, attr, value):
|
| 965 |
+
obj, name = resolve_attr(obj, attr)
|
| 966 |
+
prev = getattr(obj, name, ATTR_UNSET)
|
| 967 |
+
if value is ATTR_UNSET:
|
| 968 |
+
delattr(obj, name)
|
| 969 |
+
else:
|
| 970 |
+
setattr(obj, name, value)
|
| 971 |
+
return prev
|
| 972 |
+
|
| 973 |
+
def set_attr_param(obj, attr, value):
|
| 974 |
+
# Clone inference tensors (created under torch.inference_mode) since
|
| 975 |
+
# their version counter is frozen and nn.Parameter() cannot wrap them.
|
| 976 |
+
if (not torch.is_inference_mode_enabled()) and value.is_inference():
|
| 977 |
+
value = value.clone()
|
| 978 |
+
return set_attr(obj, attr, torch.nn.Parameter(value, requires_grad=False))
|
| 979 |
+
|
| 980 |
+
def set_attr_buffer(obj, attr, value):
|
| 981 |
+
obj, name = resolve_attr(obj, attr)
|
| 982 |
+
prev = getattr(obj, name, ATTR_UNSET)
|
| 983 |
+
persistent = name not in getattr(obj, "_non_persistent_buffers_set", set())
|
| 984 |
+
obj.register_buffer(name, value, persistent=persistent)
|
| 985 |
+
return prev
|
| 986 |
+
|
| 987 |
+
def copy_to_param(obj, attr, value):
|
| 988 |
+
# inplace update tensor instead of replacing it
|
| 989 |
+
attrs = attr.split(".")
|
| 990 |
+
for name in attrs[:-1]:
|
| 991 |
+
obj = getattr(obj, name)
|
| 992 |
+
prev = getattr(obj, attrs[-1])
|
| 993 |
+
prev.data.copy_(value)
|
| 994 |
+
|
| 995 |
+
def get_attr(obj, attr: str):
|
| 996 |
+
"""Retrieves a nested attribute from an object using dot notation.
|
| 997 |
+
|
| 998 |
+
Args:
|
| 999 |
+
obj: The object to get the attribute from
|
| 1000 |
+
attr (str): The attribute path using dot notation (e.g. "model.layer.weight")
|
| 1001 |
+
|
| 1002 |
+
Returns:
|
| 1003 |
+
The value of the requested attribute
|
| 1004 |
+
|
| 1005 |
+
Example:
|
| 1006 |
+
model = MyModel()
|
| 1007 |
+
weight = get_attr(model, "layer1.conv.weight")
|
| 1008 |
+
# Equivalent to: model.layer1.conv.weight
|
| 1009 |
+
|
| 1010 |
+
Important:
|
| 1011 |
+
Always prefer `comfy.model_patcher.ModelPatcher.get_model_object` when
|
| 1012 |
+
accessing nested model objects under `ModelPatcher.model`.
|
| 1013 |
+
"""
|
| 1014 |
+
attrs = attr.split(".")
|
| 1015 |
+
for name in attrs:
|
| 1016 |
+
obj = getattr(obj, name)
|
| 1017 |
+
return obj
|
| 1018 |
+
|
| 1019 |
+
def bislerp(samples, width, height):
|
| 1020 |
+
def slerp(b1, b2, r):
|
| 1021 |
+
'''slerps batches b1, b2 according to ratio r, batches should be flat e.g. NxC'''
|
| 1022 |
+
|
| 1023 |
+
c = b1.shape[-1]
|
| 1024 |
+
|
| 1025 |
+
#norms
|
| 1026 |
+
b1_norms = torch.norm(b1, dim=-1, keepdim=True)
|
| 1027 |
+
b2_norms = torch.norm(b2, dim=-1, keepdim=True)
|
| 1028 |
+
|
| 1029 |
+
#normalize
|
| 1030 |
+
b1_normalized = b1 / b1_norms
|
| 1031 |
+
b2_normalized = b2 / b2_norms
|
| 1032 |
+
|
| 1033 |
+
#zero when norms are zero
|
| 1034 |
+
b1_normalized[b1_norms.expand(-1,c) == 0.0] = 0.0
|
| 1035 |
+
b2_normalized[b2_norms.expand(-1,c) == 0.0] = 0.0
|
| 1036 |
+
|
| 1037 |
+
#slerp
|
| 1038 |
+
dot = (b1_normalized*b2_normalized).sum(1)
|
| 1039 |
+
omega = torch.acos(dot)
|
| 1040 |
+
so = torch.sin(omega)
|
| 1041 |
+
|
| 1042 |
+
#technically not mathematically correct, but more pleasing?
|
| 1043 |
+
res = (torch.sin((1.0-r.squeeze(1))*omega)/so).unsqueeze(1)*b1_normalized + (torch.sin(r.squeeze(1)*omega)/so).unsqueeze(1) * b2_normalized
|
| 1044 |
+
res *= (b1_norms * (1.0-r) + b2_norms * r).expand(-1,c)
|
| 1045 |
+
|
| 1046 |
+
#edge cases for same or polar opposites
|
| 1047 |
+
res[dot > 1 - 1e-5] = b1[dot > 1 - 1e-5]
|
| 1048 |
+
res[dot < 1e-5 - 1] = (b1 * (1.0-r) + b2 * r)[dot < 1e-5 - 1]
|
| 1049 |
+
return res
|
| 1050 |
+
|
| 1051 |
+
def generate_bilinear_data(length_old, length_new, device):
|
| 1052 |
+
coords_1 = torch.arange(length_old, dtype=torch.float32, device=device).reshape((1,1,1,-1))
|
| 1053 |
+
coords_1 = torch.nn.functional.interpolate(coords_1, size=(1, length_new), mode="bilinear")
|
| 1054 |
+
ratios = coords_1 - coords_1.floor()
|
| 1055 |
+
coords_1 = coords_1.to(torch.int64)
|
| 1056 |
+
|
| 1057 |
+
coords_2 = torch.arange(length_old, dtype=torch.float32, device=device).reshape((1,1,1,-1)) + 1
|
| 1058 |
+
coords_2[:,:,:,-1] -= 1
|
| 1059 |
+
coords_2 = torch.nn.functional.interpolate(coords_2, size=(1, length_new), mode="bilinear")
|
| 1060 |
+
coords_2 = coords_2.to(torch.int64)
|
| 1061 |
+
return ratios, coords_1, coords_2
|
| 1062 |
+
|
| 1063 |
+
orig_dtype = samples.dtype
|
| 1064 |
+
samples = samples.float()
|
| 1065 |
+
n,c,h,w = samples.shape
|
| 1066 |
+
h_new, w_new = (height, width)
|
| 1067 |
+
|
| 1068 |
+
#linear w
|
| 1069 |
+
ratios, coords_1, coords_2 = generate_bilinear_data(w, w_new, samples.device)
|
| 1070 |
+
coords_1 = coords_1.expand((n, c, h, -1))
|
| 1071 |
+
coords_2 = coords_2.expand((n, c, h, -1))
|
| 1072 |
+
ratios = ratios.expand((n, 1, h, -1))
|
| 1073 |
+
|
| 1074 |
+
pass_1 = samples.gather(-1,coords_1).movedim(1, -1).reshape((-1,c))
|
| 1075 |
+
pass_2 = samples.gather(-1,coords_2).movedim(1, -1).reshape((-1,c))
|
| 1076 |
+
ratios = ratios.movedim(1, -1).reshape((-1,1))
|
| 1077 |
+
|
| 1078 |
+
result = slerp(pass_1, pass_2, ratios)
|
| 1079 |
+
result = result.reshape(n, h, w_new, c).movedim(-1, 1)
|
| 1080 |
+
|
| 1081 |
+
#linear h
|
| 1082 |
+
ratios, coords_1, coords_2 = generate_bilinear_data(h, h_new, samples.device)
|
| 1083 |
+
coords_1 = coords_1.reshape((1,1,-1,1)).expand((n, c, -1, w_new))
|
| 1084 |
+
coords_2 = coords_2.reshape((1,1,-1,1)).expand((n, c, -1, w_new))
|
| 1085 |
+
ratios = ratios.reshape((1,1,-1,1)).expand((n, 1, -1, w_new))
|
| 1086 |
+
|
| 1087 |
+
pass_1 = result.gather(-2,coords_1).movedim(1, -1).reshape((-1,c))
|
| 1088 |
+
pass_2 = result.gather(-2,coords_2).movedim(1, -1).reshape((-1,c))
|
| 1089 |
+
ratios = ratios.movedim(1, -1).reshape((-1,1))
|
| 1090 |
+
|
| 1091 |
+
result = slerp(pass_1, pass_2, ratios)
|
| 1092 |
+
result = result.reshape(n, h_new, w_new, c).movedim(-1, 1)
|
| 1093 |
+
return result.to(orig_dtype)
|
| 1094 |
+
|
| 1095 |
+
def lanczos(samples, width, height):
|
| 1096 |
+
#the below API is strict and expects grayscale to be squeezed
|
| 1097 |
+
if samples.ndim == 4:
|
| 1098 |
+
samples = samples.squeeze(1) if samples.shape[1] == 1 else samples.movedim(1, -1)
|
| 1099 |
+
images = [Image.fromarray(np.clip(255. * image.cpu().numpy(), 0, 255).astype(np.uint8)) for image in samples]
|
| 1100 |
+
images = [image.resize((width, height), resample=Image.Resampling.LANCZOS) for image in images]
|
| 1101 |
+
images = [torch.from_numpy(t).movedim(-1, 0) if (t := np.array(image).astype(np.float32) / 255.0).ndim == 3 else torch.from_numpy(t) for image in images]
|
| 1102 |
+
result = torch.stack(images)
|
| 1103 |
+
return result.to(samples.device, samples.dtype)
|
| 1104 |
+
|
| 1105 |
+
def common_upscale(samples, width, height, upscale_method, crop):
|
| 1106 |
+
orig_shape = tuple(samples.shape)
|
| 1107 |
+
if len(orig_shape) > 4:
|
| 1108 |
+
samples = samples.reshape(samples.shape[0], samples.shape[1], -1, samples.shape[-2], samples.shape[-1])
|
| 1109 |
+
samples = samples.movedim(2, 1)
|
| 1110 |
+
samples = samples.reshape(-1, orig_shape[1], orig_shape[-2], orig_shape[-1])
|
| 1111 |
+
if crop == "center":
|
| 1112 |
+
old_width = samples.shape[-1]
|
| 1113 |
+
old_height = samples.shape[-2]
|
| 1114 |
+
old_aspect = old_width / old_height
|
| 1115 |
+
new_aspect = width / height
|
| 1116 |
+
x = 0
|
| 1117 |
+
y = 0
|
| 1118 |
+
if old_aspect > new_aspect:
|
| 1119 |
+
x = round((old_width - old_width * (new_aspect / old_aspect)) / 2)
|
| 1120 |
+
elif old_aspect < new_aspect:
|
| 1121 |
+
y = round((old_height - old_height * (old_aspect / new_aspect)) / 2)
|
| 1122 |
+
s = samples.narrow(-2, y, old_height - y * 2).narrow(-1, x, old_width - x * 2)
|
| 1123 |
+
else:
|
| 1124 |
+
s = samples
|
| 1125 |
+
|
| 1126 |
+
if upscale_method == "bislerp":
|
| 1127 |
+
out = bislerp(s, width, height)
|
| 1128 |
+
elif upscale_method == "lanczos":
|
| 1129 |
+
out = lanczos(s, width, height)
|
| 1130 |
+
else:
|
| 1131 |
+
out = torch.nn.functional.interpolate(s, size=(height, width), mode=upscale_method)
|
| 1132 |
+
|
| 1133 |
+
if len(orig_shape) == 4:
|
| 1134 |
+
return out
|
| 1135 |
+
|
| 1136 |
+
out = out.reshape((orig_shape[0], -1, orig_shape[1]) + (height, width))
|
| 1137 |
+
return out.movedim(2, 1).reshape(orig_shape[:-2] + (height, width))
|
| 1138 |
+
|
| 1139 |
+
def get_tiled_scale_steps(width, height, tile_x, tile_y, overlap):
|
| 1140 |
+
rows = 1 if height <= tile_y else math.ceil((height - overlap) / (tile_y - overlap))
|
| 1141 |
+
cols = 1 if width <= tile_x else math.ceil((width - overlap) / (tile_x - overlap))
|
| 1142 |
+
return rows * cols
|
| 1143 |
+
|
| 1144 |
+
@torch.inference_mode()
|
| 1145 |
+
def tiled_scale_multidim(samples, function, tile=(64, 64), overlap=8, upscale_amount=4, out_channels=3, output_device="cpu", downscale=False, index_formulas=None, pbar=None):
|
| 1146 |
+
dims = len(tile)
|
| 1147 |
+
|
| 1148 |
+
if not (isinstance(upscale_amount, (tuple, list))):
|
| 1149 |
+
upscale_amount = [upscale_amount] * dims
|
| 1150 |
+
|
| 1151 |
+
if not (isinstance(overlap, (tuple, list))):
|
| 1152 |
+
overlap = [overlap] * dims
|
| 1153 |
+
|
| 1154 |
+
if index_formulas is None:
|
| 1155 |
+
index_formulas = upscale_amount
|
| 1156 |
+
|
| 1157 |
+
if not (isinstance(index_formulas, (tuple, list))):
|
| 1158 |
+
index_formulas = [index_formulas] * dims
|
| 1159 |
+
|
| 1160 |
+
def get_upscale(dim, val):
|
| 1161 |
+
up = upscale_amount[dim]
|
| 1162 |
+
if callable(up):
|
| 1163 |
+
return up(val)
|
| 1164 |
+
else:
|
| 1165 |
+
return up * val
|
| 1166 |
+
|
| 1167 |
+
def get_downscale(dim, val):
|
| 1168 |
+
up = upscale_amount[dim]
|
| 1169 |
+
if callable(up):
|
| 1170 |
+
return up(val)
|
| 1171 |
+
else:
|
| 1172 |
+
return val / up
|
| 1173 |
+
|
| 1174 |
+
def get_upscale_pos(dim, val):
|
| 1175 |
+
up = index_formulas[dim]
|
| 1176 |
+
if callable(up):
|
| 1177 |
+
return up(val)
|
| 1178 |
+
else:
|
| 1179 |
+
return up * val
|
| 1180 |
+
|
| 1181 |
+
def get_downscale_pos(dim, val):
|
| 1182 |
+
up = index_formulas[dim]
|
| 1183 |
+
if callable(up):
|
| 1184 |
+
return up(val)
|
| 1185 |
+
else:
|
| 1186 |
+
return val / up
|
| 1187 |
+
|
| 1188 |
+
if downscale:
|
| 1189 |
+
get_scale = get_downscale
|
| 1190 |
+
get_pos = get_downscale_pos
|
| 1191 |
+
else:
|
| 1192 |
+
get_scale = get_upscale
|
| 1193 |
+
get_pos = get_upscale_pos
|
| 1194 |
+
|
| 1195 |
+
def mult_list_upscale(a):
|
| 1196 |
+
out = []
|
| 1197 |
+
for i in range(len(a)):
|
| 1198 |
+
out.append(round(get_scale(i, a[i])))
|
| 1199 |
+
return out
|
| 1200 |
+
|
| 1201 |
+
output = torch.empty([samples.shape[0], out_channels] + mult_list_upscale(samples.shape[2:]), device=output_device)
|
| 1202 |
+
|
| 1203 |
+
for b in range(samples.shape[0]):
|
| 1204 |
+
s = samples[b:b+1]
|
| 1205 |
+
|
| 1206 |
+
# handle entire input fitting in a single tile
|
| 1207 |
+
if all(s.shape[d+2] <= tile[d] for d in range(dims)):
|
| 1208 |
+
output[b:b+1] = function(s).to(output_device)
|
| 1209 |
+
if pbar is not None:
|
| 1210 |
+
pbar.update(1)
|
| 1211 |
+
continue
|
| 1212 |
+
|
| 1213 |
+
out = output[b:b+1].zero_()
|
| 1214 |
+
out_div = torch.zeros([s.shape[0], 1] + mult_list_upscale(s.shape[2:]), device=output_device)
|
| 1215 |
+
|
| 1216 |
+
positions = [range(0, s.shape[d+2] - overlap[d], tile[d] - overlap[d]) if s.shape[d+2] > tile[d] else [0] for d in range(dims)]
|
| 1217 |
+
|
| 1218 |
+
for it in itertools.product(*positions):
|
| 1219 |
+
s_in = s
|
| 1220 |
+
upscaled = []
|
| 1221 |
+
|
| 1222 |
+
for d in range(dims):
|
| 1223 |
+
pos = max(0, min(s.shape[d + 2] - overlap[d], it[d]))
|
| 1224 |
+
l = min(tile[d], s.shape[d + 2] - pos)
|
| 1225 |
+
s_in = s_in.narrow(d + 2, pos, l)
|
| 1226 |
+
upscaled.append(round(get_pos(d, pos)))
|
| 1227 |
+
|
| 1228 |
+
ps = function(s_in).to(output_device)
|
| 1229 |
+
mask = torch.ones([1, 1] + list(ps.shape[2:]), device=output_device)
|
| 1230 |
+
|
| 1231 |
+
for d in range(2, dims + 2):
|
| 1232 |
+
feather = round(get_scale(d - 2, overlap[d - 2]))
|
| 1233 |
+
if feather >= mask.shape[d]:
|
| 1234 |
+
continue
|
| 1235 |
+
for t in range(feather):
|
| 1236 |
+
a = (t + 1) / feather
|
| 1237 |
+
mask.narrow(d, t, 1).mul_(a)
|
| 1238 |
+
mask.narrow(d, mask.shape[d] - 1 - t, 1).mul_(a)
|
| 1239 |
+
|
| 1240 |
+
o = out
|
| 1241 |
+
o_d = out_div
|
| 1242 |
+
ps_view = ps
|
| 1243 |
+
mask_view = mask
|
| 1244 |
+
for d in range(dims):
|
| 1245 |
+
l = min(ps_view.shape[d + 2], o.shape[d + 2] - upscaled[d])
|
| 1246 |
+
o = o.narrow(d + 2, upscaled[d], l)
|
| 1247 |
+
o_d = o_d.narrow(d + 2, upscaled[d], l)
|
| 1248 |
+
if l < ps_view.shape[d + 2]:
|
| 1249 |
+
ps_view = ps_view.narrow(d + 2, 0, l)
|
| 1250 |
+
mask_view = mask_view.narrow(d + 2, 0, l)
|
| 1251 |
+
|
| 1252 |
+
o.add_(ps_view * mask_view)
|
| 1253 |
+
o_d.add_(mask_view)
|
| 1254 |
+
|
| 1255 |
+
if pbar is not None:
|
| 1256 |
+
pbar.update(1)
|
| 1257 |
+
|
| 1258 |
+
out.div_(out_div)
|
| 1259 |
+
return output
|
| 1260 |
+
|
| 1261 |
+
def tiled_scale(samples, function, tile_x=64, tile_y=64, overlap = 8, upscale_amount = 4, out_channels = 3, output_device="cpu", pbar = None):
|
| 1262 |
+
return tiled_scale_multidim(samples, function, (tile_y, tile_x), overlap=overlap, upscale_amount=upscale_amount, out_channels=out_channels, output_device=output_device, pbar=pbar)
|
| 1263 |
+
|
| 1264 |
+
def model_trange(*args, **kwargs):
|
| 1265 |
+
if not comfy.memory_management.aimdo_enabled:
|
| 1266 |
+
return trange(*args, **kwargs)
|
| 1267 |
+
|
| 1268 |
+
pbar = trange(*args, **kwargs, smoothing=1.0)
|
| 1269 |
+
pbar._i = 0
|
| 1270 |
+
pbar.set_postfix_str(" Model Initializing ... ")
|
| 1271 |
+
|
| 1272 |
+
_update = pbar.update
|
| 1273 |
+
|
| 1274 |
+
def warmup_update(n=1):
|
| 1275 |
+
pbar._i += 1
|
| 1276 |
+
if pbar._i == 1:
|
| 1277 |
+
pbar.i1_time = time.time()
|
| 1278 |
+
pbar.set_postfix_str(" Model Initialization complete! ")
|
| 1279 |
+
elif pbar._i == 2:
|
| 1280 |
+
#bring forward the effective start time based the diff between first and second iteration
|
| 1281 |
+
#to attempt to remove load overhead from the final step rate estimate.
|
| 1282 |
+
pbar.start_t = pbar.i1_time - (time.time() - pbar.i1_time)
|
| 1283 |
+
pbar.set_postfix_str("")
|
| 1284 |
+
|
| 1285 |
+
_update(n)
|
| 1286 |
+
|
| 1287 |
+
pbar.update = warmup_update
|
| 1288 |
+
return pbar
|
| 1289 |
+
|
| 1290 |
+
PROGRESS_BAR_ENABLED = True
|
| 1291 |
+
def set_progress_bar_enabled(enabled):
|
| 1292 |
+
global PROGRESS_BAR_ENABLED
|
| 1293 |
+
PROGRESS_BAR_ENABLED = enabled
|
| 1294 |
+
|
| 1295 |
+
PROGRESS_BAR_HOOK = None
|
| 1296 |
+
def set_progress_bar_global_hook(function):
|
| 1297 |
+
global PROGRESS_BAR_HOOK
|
| 1298 |
+
PROGRESS_BAR_HOOK = function
|
| 1299 |
+
|
| 1300 |
+
# Throttle settings for progress bar updates to reduce WebSocket flooding
|
| 1301 |
+
PROGRESS_THROTTLE_MIN_INTERVAL = 0.1 # 100ms minimum between updates
|
| 1302 |
+
PROGRESS_THROTTLE_MIN_PERCENT = 0.5 # 0.5% minimum progress change
|
| 1303 |
+
|
| 1304 |
+
class ProgressBar:
|
| 1305 |
+
def __init__(self, total, node_id=None):
|
| 1306 |
+
global PROGRESS_BAR_HOOK
|
| 1307 |
+
self.total = total
|
| 1308 |
+
self.current = 0
|
| 1309 |
+
self.hook = PROGRESS_BAR_HOOK
|
| 1310 |
+
self.node_id = node_id
|
| 1311 |
+
self._last_update_time = 0.0
|
| 1312 |
+
self._last_sent_value = -1
|
| 1313 |
+
|
| 1314 |
+
def update_absolute(self, value, total=None, preview=None):
|
| 1315 |
+
if total is not None:
|
| 1316 |
+
self.total = total
|
| 1317 |
+
if value > self.total:
|
| 1318 |
+
value = self.total
|
| 1319 |
+
self.current = value
|
| 1320 |
+
if self.hook is not None:
|
| 1321 |
+
current_time = time.perf_counter()
|
| 1322 |
+
is_first = (self._last_sent_value < 0)
|
| 1323 |
+
is_final = (value >= self.total)
|
| 1324 |
+
has_preview = (preview is not None)
|
| 1325 |
+
|
| 1326 |
+
# Always send immediately for previews, first update, or final update
|
| 1327 |
+
if has_preview or is_first or is_final:
|
| 1328 |
+
self.hook(self.current, self.total, preview, node_id=self.node_id)
|
| 1329 |
+
self._last_update_time = current_time
|
| 1330 |
+
self._last_sent_value = value
|
| 1331 |
+
return
|
| 1332 |
+
|
| 1333 |
+
# Apply throttling for regular progress updates
|
| 1334 |
+
if self.total > 0:
|
| 1335 |
+
percent_changed = ((value - max(0, self._last_sent_value)) / self.total) * 100
|
| 1336 |
+
else:
|
| 1337 |
+
percent_changed = 100
|
| 1338 |
+
time_elapsed = current_time - self._last_update_time
|
| 1339 |
+
|
| 1340 |
+
if time_elapsed >= PROGRESS_THROTTLE_MIN_INTERVAL and percent_changed >= PROGRESS_THROTTLE_MIN_PERCENT:
|
| 1341 |
+
self.hook(self.current, self.total, preview, node_id=self.node_id)
|
| 1342 |
+
self._last_update_time = current_time
|
| 1343 |
+
self._last_sent_value = value
|
| 1344 |
+
|
| 1345 |
+
def update(self, value):
|
| 1346 |
+
self.update_absolute(self.current + value)
|
| 1347 |
+
|
| 1348 |
+
def reshape_mask(input_mask, output_shape):
|
| 1349 |
+
dims = len(output_shape) - 2
|
| 1350 |
+
|
| 1351 |
+
if dims == 1:
|
| 1352 |
+
scale_mode = "linear"
|
| 1353 |
+
|
| 1354 |
+
if dims == 2:
|
| 1355 |
+
input_mask = input_mask.reshape((-1, 1, input_mask.shape[-2], input_mask.shape[-1]))
|
| 1356 |
+
scale_mode = "bilinear"
|
| 1357 |
+
|
| 1358 |
+
if dims == 3:
|
| 1359 |
+
if len(input_mask.shape) < 5:
|
| 1360 |
+
input_mask = input_mask.reshape((1, 1, -1, input_mask.shape[-2], input_mask.shape[-1]))
|
| 1361 |
+
scale_mode = "trilinear"
|
| 1362 |
+
|
| 1363 |
+
mask = torch.nn.functional.interpolate(input_mask, size=output_shape[2:], mode=scale_mode)
|
| 1364 |
+
if mask.shape[1] < output_shape[1]:
|
| 1365 |
+
mask = mask.repeat((1, output_shape[1]) + (1,) * dims)[:,:output_shape[1]]
|
| 1366 |
+
mask = repeat_to_batch_size(mask, output_shape[0])
|
| 1367 |
+
return mask
|
| 1368 |
+
|
| 1369 |
+
def upscale_dit_mask(mask: torch.Tensor, img_size_in, img_size_out):
|
| 1370 |
+
hi, wi = img_size_in
|
| 1371 |
+
ho, wo = img_size_out
|
| 1372 |
+
# if it's already the correct size, no need to do anything
|
| 1373 |
+
if (hi, wi) == (ho, wo):
|
| 1374 |
+
return mask
|
| 1375 |
+
if mask.ndim == 2:
|
| 1376 |
+
mask = mask.unsqueeze(0)
|
| 1377 |
+
if mask.ndim != 3:
|
| 1378 |
+
raise ValueError(f"Got a mask of shape {list(mask.shape)}, expected [b, q, k] or [q, k]")
|
| 1379 |
+
txt_tokens = mask.shape[1] - (hi * wi)
|
| 1380 |
+
# quadrants of the mask
|
| 1381 |
+
txt_to_txt = mask[:, :txt_tokens, :txt_tokens]
|
| 1382 |
+
txt_to_img = mask[:, :txt_tokens, txt_tokens:]
|
| 1383 |
+
img_to_img = mask[:, txt_tokens:, txt_tokens:]
|
| 1384 |
+
img_to_txt = mask[:, txt_tokens:, :txt_tokens]
|
| 1385 |
+
|
| 1386 |
+
# convert to 1d x 2d, interpolate, then back to 1d x 1d
|
| 1387 |
+
txt_to_img = rearrange (txt_to_img, "b t (h w) -> b t h w", h=hi, w=wi)
|
| 1388 |
+
txt_to_img = interpolate(txt_to_img, size=img_size_out, mode="bilinear")
|
| 1389 |
+
txt_to_img = rearrange (txt_to_img, "b t h w -> b t (h w)")
|
| 1390 |
+
# this one is hard because we have to do it twice
|
| 1391 |
+
# convert to 1d x 2d, interpolate, then to 2d x 1d, interpolate, then 1d x 1d
|
| 1392 |
+
img_to_img = rearrange (img_to_img, "b hw (h w) -> b hw h w", h=hi, w=wi)
|
| 1393 |
+
img_to_img = interpolate(img_to_img, size=img_size_out, mode="bilinear")
|
| 1394 |
+
img_to_img = rearrange (img_to_img, "b (hk wk) hq wq -> b (hq wq) hk wk", hk=hi, wk=wi)
|
| 1395 |
+
img_to_img = interpolate(img_to_img, size=img_size_out, mode="bilinear")
|
| 1396 |
+
img_to_img = rearrange (img_to_img, "b (hq wq) hk wk -> b (hk wk) (hq wq)", hq=ho, wq=wo)
|
| 1397 |
+
# convert to 2d x 1d, interpolate, then back to 1d x 1d
|
| 1398 |
+
img_to_txt = rearrange (img_to_txt, "b (h w) t -> b t h w", h=hi, w=wi)
|
| 1399 |
+
img_to_txt = interpolate(img_to_txt, size=img_size_out, mode="bilinear")
|
| 1400 |
+
img_to_txt = rearrange (img_to_txt, "b t h w -> b (h w) t")
|
| 1401 |
+
|
| 1402 |
+
# reassemble the mask from blocks
|
| 1403 |
+
out = torch.cat([
|
| 1404 |
+
torch.cat([txt_to_txt, txt_to_img], dim=2),
|
| 1405 |
+
torch.cat([img_to_txt, img_to_img], dim=2)],
|
| 1406 |
+
dim=1
|
| 1407 |
+
)
|
| 1408 |
+
return out
|
| 1409 |
+
|
| 1410 |
+
def pack_latents(latents):
|
| 1411 |
+
latent_shapes = []
|
| 1412 |
+
tensors = []
|
| 1413 |
+
for tensor in latents:
|
| 1414 |
+
latent_shapes.append(tensor.shape)
|
| 1415 |
+
tensors.append(tensor.reshape(tensor.shape[0], 1, -1))
|
| 1416 |
+
|
| 1417 |
+
latent = torch.cat(tensors, dim=-1)
|
| 1418 |
+
return latent, latent_shapes
|
| 1419 |
+
|
| 1420 |
+
def unpack_latents(combined_latent, latent_shapes):
|
| 1421 |
+
if len(latent_shapes) > 1:
|
| 1422 |
+
output_tensors = []
|
| 1423 |
+
for shape in latent_shapes:
|
| 1424 |
+
cut = math.prod(shape[1:])
|
| 1425 |
+
tens = combined_latent[:, :, :cut]
|
| 1426 |
+
combined_latent = combined_latent[:, :, cut:]
|
| 1427 |
+
output_tensors.append(tens.reshape([tens.shape[0]] + list(shape)[1:]))
|
| 1428 |
+
else:
|
| 1429 |
+
output_tensors = [combined_latent]
|
| 1430 |
+
return output_tensors
|
| 1431 |
+
|
| 1432 |
+
def detect_layer_quantization(state_dict, prefix):
|
| 1433 |
+
for k in state_dict:
|
| 1434 |
+
if k.startswith(prefix) and k.endswith(".comfy_quant"):
|
| 1435 |
+
logging.info("Found quantization metadata version 1")
|
| 1436 |
+
return {"mixed_ops": True}
|
| 1437 |
+
return None
|
| 1438 |
+
|
| 1439 |
+
def convert_old_quants(state_dict, model_prefix="", metadata={}):
|
| 1440 |
+
if metadata is None:
|
| 1441 |
+
metadata = {}
|
| 1442 |
+
|
| 1443 |
+
quant_metadata = None
|
| 1444 |
+
if "_quantization_metadata" not in metadata:
|
| 1445 |
+
scaled_fp8_key = "{}scaled_fp8".format(model_prefix)
|
| 1446 |
+
|
| 1447 |
+
if scaled_fp8_key in state_dict:
|
| 1448 |
+
scaled_fp8_weight = state_dict[scaled_fp8_key]
|
| 1449 |
+
scaled_fp8_dtype = scaled_fp8_weight.dtype
|
| 1450 |
+
if scaled_fp8_dtype == torch.float32:
|
| 1451 |
+
scaled_fp8_dtype = torch.float8_e4m3fn
|
| 1452 |
+
|
| 1453 |
+
if scaled_fp8_weight.nelement() == 2:
|
| 1454 |
+
full_precision_matrix_mult = True
|
| 1455 |
+
else:
|
| 1456 |
+
full_precision_matrix_mult = False
|
| 1457 |
+
|
| 1458 |
+
out_sd = {}
|
| 1459 |
+
layers = {}
|
| 1460 |
+
for k in list(state_dict.keys()):
|
| 1461 |
+
if k == scaled_fp8_key:
|
| 1462 |
+
continue
|
| 1463 |
+
if not k.startswith(model_prefix):
|
| 1464 |
+
out_sd[k] = state_dict[k]
|
| 1465 |
+
continue
|
| 1466 |
+
k_out = k
|
| 1467 |
+
w = state_dict.pop(k)
|
| 1468 |
+
layer = None
|
| 1469 |
+
if k_out.endswith(".scale_weight"):
|
| 1470 |
+
layer = k_out[:-len(".scale_weight")]
|
| 1471 |
+
k_out = "{}.weight_scale".format(layer)
|
| 1472 |
+
|
| 1473 |
+
if layer is not None:
|
| 1474 |
+
layer_conf = {"format": "float8_e4m3fn"}
|
| 1475 |
+
if full_precision_matrix_mult:
|
| 1476 |
+
layer_conf["full_precision_matrix_mult"] = full_precision_matrix_mult
|
| 1477 |
+
layers[layer] = layer_conf
|
| 1478 |
+
|
| 1479 |
+
if k_out.endswith(".scale_input"):
|
| 1480 |
+
layer = k_out[:-len(".scale_input")]
|
| 1481 |
+
k_out = "{}.input_scale".format(layer)
|
| 1482 |
+
if w.item() == 1.0:
|
| 1483 |
+
continue
|
| 1484 |
+
|
| 1485 |
+
out_sd[k_out] = w
|
| 1486 |
+
|
| 1487 |
+
state_dict = out_sd
|
| 1488 |
+
quant_metadata = {"layers": layers}
|
| 1489 |
+
else:
|
| 1490 |
+
quant_metadata = json.loads(metadata["_quantization_metadata"])
|
| 1491 |
+
|
| 1492 |
+
if quant_metadata is not None:
|
| 1493 |
+
layers = quant_metadata["layers"]
|
| 1494 |
+
for k, v in layers.items():
|
| 1495 |
+
state_dict["{}.comfy_quant".format(k)] = torch.tensor(list(json.dumps(v).encode('utf-8')), dtype=torch.uint8)
|
| 1496 |
+
|
| 1497 |
+
return state_dict, metadata
|
| 1498 |
+
|
| 1499 |
+
def string_to_seed(data):
|
| 1500 |
+
crc = 0xFFFFFFFF
|
| 1501 |
+
for byte in data:
|
| 1502 |
+
if isinstance(byte, str):
|
| 1503 |
+
byte = ord(byte)
|
| 1504 |
+
crc ^= byte
|
| 1505 |
+
for _ in range(8):
|
| 1506 |
+
if crc & 1:
|
| 1507 |
+
crc = (crc >> 1) ^ 0xEDB88320
|
| 1508 |
+
else:
|
| 1509 |
+
crc >>= 1
|
| 1510 |
+
return crc ^ 0xFFFFFFFF
|
| 1511 |
+
|
| 1512 |
+
def deepcopy_list_dict(obj, memo=None):
|
| 1513 |
+
if memo is None:
|
| 1514 |
+
memo = {}
|
| 1515 |
+
|
| 1516 |
+
obj_id = id(obj)
|
| 1517 |
+
if obj_id in memo:
|
| 1518 |
+
return memo[obj_id]
|
| 1519 |
+
|
| 1520 |
+
if isinstance(obj, dict):
|
| 1521 |
+
res = {deepcopy_list_dict(k, memo): deepcopy_list_dict(v, memo) for k, v in obj.items()}
|
| 1522 |
+
elif isinstance(obj, list):
|
| 1523 |
+
res = [deepcopy_list_dict(i, memo) for i in obj]
|
| 1524 |
+
else:
|
| 1525 |
+
res = obj
|
| 1526 |
+
|
| 1527 |
+
memo[obj_id] = res
|
| 1528 |
+
return res
|
| 1529 |
+
|
| 1530 |
+
def bit_reverse_range(index, bits):
|
| 1531 |
+
result = 0
|
| 1532 |
+
for _ in range(bits):
|
| 1533 |
+
result = (result << 1) | (index & 1)
|
| 1534 |
+
index >>= 1
|
| 1535 |
+
return result
|
comfy/weight_adapter/__init__.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .base import WeightAdapterBase, WeightAdapterTrainBase
|
| 2 |
+
from .lora import LoRAAdapter
|
| 3 |
+
from .loha import LoHaAdapter
|
| 4 |
+
from .lokr import LoKrAdapter
|
| 5 |
+
from .glora import GLoRAAdapter
|
| 6 |
+
from .oft import OFTAdapter
|
| 7 |
+
from .boft import BOFTAdapter
|
| 8 |
+
from .bypass import (
|
| 9 |
+
BypassInjectionManager,
|
| 10 |
+
BypassForwardHook,
|
| 11 |
+
create_bypass_injections_from_patches,
|
| 12 |
+
)
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
adapters: list[type[WeightAdapterBase]] = [
|
| 16 |
+
LoRAAdapter,
|
| 17 |
+
LoHaAdapter,
|
| 18 |
+
LoKrAdapter,
|
| 19 |
+
GLoRAAdapter,
|
| 20 |
+
OFTAdapter,
|
| 21 |
+
BOFTAdapter,
|
| 22 |
+
]
|
| 23 |
+
adapter_maps: dict[str, type[WeightAdapterBase]] = {
|
| 24 |
+
"LoRA": LoRAAdapter,
|
| 25 |
+
"LoHa": LoHaAdapter,
|
| 26 |
+
"LoKr": LoKrAdapter,
|
| 27 |
+
"OFT": OFTAdapter,
|
| 28 |
+
## We disable not implemented algo for now
|
| 29 |
+
# "GLoRA": GLoRAAdapter,
|
| 30 |
+
# "BOFT": BOFTAdapter,
|
| 31 |
+
}
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
__all__ = [
|
| 35 |
+
"WeightAdapterBase",
|
| 36 |
+
"WeightAdapterTrainBase",
|
| 37 |
+
"adapters",
|
| 38 |
+
"adapter_maps",
|
| 39 |
+
"BypassInjectionManager",
|
| 40 |
+
"BypassForwardHook",
|
| 41 |
+
"create_bypass_injections_from_patches",
|
| 42 |
+
] + [a.__name__ for a in adapters]
|
comfy/weight_adapter/base.py
ADDED
|
@@ -0,0 +1,396 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Callable, Optional
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn as nn
|
| 5 |
+
|
| 6 |
+
import comfy.model_management
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
class WeightAdapterBase:
|
| 10 |
+
"""
|
| 11 |
+
Base class for weight adapters (LoRA, LoHa, LoKr, OFT, etc.)
|
| 12 |
+
|
| 13 |
+
Bypass Mode:
|
| 14 |
+
All adapters follow the pattern: bypass(f)(x) = g(f(x) + h(x))
|
| 15 |
+
|
| 16 |
+
- h(x): Additive component (LoRA path). Returns delta to add to base output.
|
| 17 |
+
- g(y): Output transformation. Applied after base + h(x).
|
| 18 |
+
|
| 19 |
+
For LoRA/LoHa/LoKr: g = identity, h = adapter(x)
|
| 20 |
+
For OFT/BOFT: g = transform, h = 0
|
| 21 |
+
"""
|
| 22 |
+
|
| 23 |
+
name: str
|
| 24 |
+
loaded_keys: set[str]
|
| 25 |
+
weights: list[torch.Tensor]
|
| 26 |
+
|
| 27 |
+
# Attributes set by bypass system
|
| 28 |
+
multiplier: float = 1.0
|
| 29 |
+
shape: tuple = None # (out_features, in_features) or (out_ch, in_ch, *kernel)
|
| 30 |
+
|
| 31 |
+
@classmethod
|
| 32 |
+
def load(
|
| 33 |
+
cls,
|
| 34 |
+
x: str,
|
| 35 |
+
lora: dict[str, torch.Tensor],
|
| 36 |
+
alpha: float,
|
| 37 |
+
dora_scale: torch.Tensor,
|
| 38 |
+
) -> Optional["WeightAdapterBase"]:
|
| 39 |
+
raise NotImplementedError
|
| 40 |
+
|
| 41 |
+
def to_train(self) -> "WeightAdapterTrainBase":
|
| 42 |
+
raise NotImplementedError
|
| 43 |
+
|
| 44 |
+
@classmethod
|
| 45 |
+
def create_train(cls, weight, *args) -> "WeightAdapterTrainBase":
|
| 46 |
+
"""
|
| 47 |
+
weight: The original weight tensor to be modified.
|
| 48 |
+
*args: Additional arguments for configuration, such as rank, alpha etc.
|
| 49 |
+
"""
|
| 50 |
+
raise NotImplementedError
|
| 51 |
+
|
| 52 |
+
def calculate_shape(
|
| 53 |
+
self,
|
| 54 |
+
key
|
| 55 |
+
):
|
| 56 |
+
return None
|
| 57 |
+
|
| 58 |
+
def calculate_weight(
|
| 59 |
+
self,
|
| 60 |
+
weight,
|
| 61 |
+
key,
|
| 62 |
+
strength,
|
| 63 |
+
strength_model,
|
| 64 |
+
offset,
|
| 65 |
+
function,
|
| 66 |
+
intermediate_dtype=torch.float32,
|
| 67 |
+
original_weight=None,
|
| 68 |
+
):
|
| 69 |
+
raise NotImplementedError
|
| 70 |
+
|
| 71 |
+
# ===== Bypass Mode Methods =====
|
| 72 |
+
#
|
| 73 |
+
# IMPORTANT: Bypass mode is designed for quantized models where original weights
|
| 74 |
+
# may not be accessible in a usable format. Therefore, h() and bypass_forward()
|
| 75 |
+
# do NOT take org_weight as a parameter. All necessary information (out_channels,
|
| 76 |
+
# in_channels, conv params, etc.) is provided via attributes set by BypassForwardHook.
|
| 77 |
+
|
| 78 |
+
def h(self, x: torch.Tensor, base_out: torch.Tensor) -> torch.Tensor:
|
| 79 |
+
"""
|
| 80 |
+
Additive bypass component: h(x, base_out)
|
| 81 |
+
|
| 82 |
+
Computes the adapter's contribution to be added to base forward output.
|
| 83 |
+
For adapters that only transform output (OFT/BOFT), returns zeros.
|
| 84 |
+
|
| 85 |
+
Note:
|
| 86 |
+
This method does NOT access original model weights. Bypass mode is
|
| 87 |
+
designed for quantized models where weights may not be in a usable format.
|
| 88 |
+
All shape info comes from module attributes set by BypassForwardHook.
|
| 89 |
+
|
| 90 |
+
Args:
|
| 91 |
+
x: Input tensor
|
| 92 |
+
base_out: Output from base forward f(x), can be used for shape reference
|
| 93 |
+
|
| 94 |
+
Returns:
|
| 95 |
+
Delta tensor to add to base output. Shape matches base output.
|
| 96 |
+
|
| 97 |
+
Reference: LyCORIS LoConModule.bypass_forward_diff
|
| 98 |
+
"""
|
| 99 |
+
# Default: no additive component (for OFT/BOFT)
|
| 100 |
+
# Simply return zeros matching base_out shape
|
| 101 |
+
return torch.zeros_like(base_out)
|
| 102 |
+
|
| 103 |
+
def g(self, y: torch.Tensor) -> torch.Tensor:
|
| 104 |
+
"""
|
| 105 |
+
Output transformation: g(y)
|
| 106 |
+
|
| 107 |
+
Applied after base forward + h(x). For most adapters this is identity.
|
| 108 |
+
OFT/BOFT override this to apply orthogonal transformation.
|
| 109 |
+
|
| 110 |
+
Args:
|
| 111 |
+
y: Combined output (base + h(x))
|
| 112 |
+
|
| 113 |
+
Returns:
|
| 114 |
+
Transformed output
|
| 115 |
+
|
| 116 |
+
Reference: LyCORIS OFTModule applies orthogonal transform here
|
| 117 |
+
"""
|
| 118 |
+
# Default: identity (for LoRA/LoHa/LoKr)
|
| 119 |
+
return y
|
| 120 |
+
|
| 121 |
+
def bypass_forward(
|
| 122 |
+
self,
|
| 123 |
+
org_forward: Callable,
|
| 124 |
+
x: torch.Tensor,
|
| 125 |
+
*args,
|
| 126 |
+
**kwargs,
|
| 127 |
+
) -> torch.Tensor:
|
| 128 |
+
"""
|
| 129 |
+
Full bypass forward: g(f(x) + h(x, f(x)))
|
| 130 |
+
|
| 131 |
+
Note:
|
| 132 |
+
This method does NOT take org_weight/org_bias parameters. Bypass mode
|
| 133 |
+
is designed for quantized models where weights may not be accessible.
|
| 134 |
+
The original forward function handles weight access internally.
|
| 135 |
+
|
| 136 |
+
Args:
|
| 137 |
+
org_forward: Original module forward function
|
| 138 |
+
x: Input tensor
|
| 139 |
+
*args, **kwargs: Additional arguments for org_forward
|
| 140 |
+
|
| 141 |
+
Returns:
|
| 142 |
+
Output with adapter applied in bypass mode
|
| 143 |
+
|
| 144 |
+
Reference: LyCORIS LoConModule.bypass_forward
|
| 145 |
+
"""
|
| 146 |
+
# Base forward: f(x)
|
| 147 |
+
base_out = org_forward(x, *args, **kwargs)
|
| 148 |
+
|
| 149 |
+
# Additive component: h(x, base_out) - base_out provided for shape reference
|
| 150 |
+
h_out = self.h(x, base_out)
|
| 151 |
+
|
| 152 |
+
# Output transformation: g(base + h)
|
| 153 |
+
return self.g(base_out + h_out)
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
class WeightAdapterTrainBase(nn.Module):
|
| 157 |
+
"""
|
| 158 |
+
Base class for trainable weight adapters (LoRA, LoHa, LoKr, OFT, etc.)
|
| 159 |
+
|
| 160 |
+
Bypass Mode:
|
| 161 |
+
All adapters follow the pattern: bypass(f)(x) = g(f(x) + h(x))
|
| 162 |
+
|
| 163 |
+
- h(x): Additive component (LoRA path). Returns delta to add to base output.
|
| 164 |
+
- g(y): Output transformation. Applied after base + h(x).
|
| 165 |
+
|
| 166 |
+
For LoRA/LoHa/LoKr: g = identity, h = adapter(x)
|
| 167 |
+
For OFT: g = transform, h = 0
|
| 168 |
+
|
| 169 |
+
Note:
|
| 170 |
+
Unlike WeightAdapterBase, TrainBase classes have simplified weight formats
|
| 171 |
+
with fewer branches (e.g., LoKr only has w1/w2, not w1_a/w1_b decomposition).
|
| 172 |
+
|
| 173 |
+
We follow the scheme of PR #7032
|
| 174 |
+
"""
|
| 175 |
+
|
| 176 |
+
# Attributes set by bypass system (BypassForwardHook)
|
| 177 |
+
# These are set before h()/g()/bypass_forward() are called
|
| 178 |
+
multiplier: float = 1.0
|
| 179 |
+
is_conv: bool = False
|
| 180 |
+
conv_dim: int = 0 # 0=linear, 1=conv1d, 2=conv2d, 3=conv3d
|
| 181 |
+
kw_dict: dict = {} # Conv kwargs: stride, padding, dilation, groups
|
| 182 |
+
kernel_size: tuple = ()
|
| 183 |
+
in_channels: int = None
|
| 184 |
+
out_channels: int = None
|
| 185 |
+
|
| 186 |
+
def __init__(self):
|
| 187 |
+
super().__init__()
|
| 188 |
+
|
| 189 |
+
def __call__(self, w):
|
| 190 |
+
"""
|
| 191 |
+
Weight modification mode: returns modified weight.
|
| 192 |
+
|
| 193 |
+
Args:
|
| 194 |
+
w: The original weight tensor to be modified.
|
| 195 |
+
|
| 196 |
+
Returns:
|
| 197 |
+
Modified weight tensor.
|
| 198 |
+
"""
|
| 199 |
+
raise NotImplementedError
|
| 200 |
+
|
| 201 |
+
# ===== Bypass Mode Methods =====
|
| 202 |
+
|
| 203 |
+
def h(self, x: torch.Tensor, base_out: torch.Tensor) -> torch.Tensor:
|
| 204 |
+
"""
|
| 205 |
+
Additive bypass component: h(x, base_out)
|
| 206 |
+
|
| 207 |
+
Computes the adapter's contribution to be added to base forward output.
|
| 208 |
+
For adapters that only transform output (OFT), returns zeros.
|
| 209 |
+
|
| 210 |
+
Args:
|
| 211 |
+
x: Input tensor
|
| 212 |
+
base_out: Output from base forward f(x), can be used for shape reference
|
| 213 |
+
|
| 214 |
+
Returns:
|
| 215 |
+
Delta tensor to add to base output. Shape matches base output.
|
| 216 |
+
|
| 217 |
+
Subclasses should override this method.
|
| 218 |
+
"""
|
| 219 |
+
raise NotImplementedError(
|
| 220 |
+
f"{self.__class__.__name__}.h() not implemented. "
|
| 221 |
+
"Subclasses must implement h() for bypass mode."
|
| 222 |
+
)
|
| 223 |
+
|
| 224 |
+
def g(self, y: torch.Tensor) -> torch.Tensor:
|
| 225 |
+
"""
|
| 226 |
+
Output transformation: g(y)
|
| 227 |
+
|
| 228 |
+
Applied after base forward + h(x). For most adapters this is identity.
|
| 229 |
+
OFT overrides this to apply orthogonal transformation.
|
| 230 |
+
|
| 231 |
+
Args:
|
| 232 |
+
y: Combined output (base + h(x))
|
| 233 |
+
|
| 234 |
+
Returns:
|
| 235 |
+
Transformed output
|
| 236 |
+
"""
|
| 237 |
+
# Default: identity (for LoRA/LoHa/LoKr)
|
| 238 |
+
return y
|
| 239 |
+
|
| 240 |
+
def bypass_forward(
|
| 241 |
+
self,
|
| 242 |
+
org_forward: Callable,
|
| 243 |
+
x: torch.Tensor,
|
| 244 |
+
*args,
|
| 245 |
+
**kwargs,
|
| 246 |
+
) -> torch.Tensor:
|
| 247 |
+
"""
|
| 248 |
+
Full bypass forward: g(f(x) + h(x, f(x)))
|
| 249 |
+
|
| 250 |
+
Args:
|
| 251 |
+
org_forward: Original module forward function
|
| 252 |
+
x: Input tensor
|
| 253 |
+
*args, **kwargs: Additional arguments for org_forward
|
| 254 |
+
|
| 255 |
+
Returns:
|
| 256 |
+
Output with adapter applied in bypass mode
|
| 257 |
+
"""
|
| 258 |
+
# Base forward: f(x)
|
| 259 |
+
base_out = org_forward(x, *args, **kwargs)
|
| 260 |
+
|
| 261 |
+
# Additive component: h(x, base_out) - base_out provided for shape reference
|
| 262 |
+
h_out = self.h(x, base_out)
|
| 263 |
+
|
| 264 |
+
# Output transformation: g(base + h)
|
| 265 |
+
return self.g(base_out + h_out)
|
| 266 |
+
|
| 267 |
+
def passive_memory_usage(self):
|
| 268 |
+
raise NotImplementedError("passive_memory_usage is not implemented")
|
| 269 |
+
|
| 270 |
+
def move_to(self, device):
|
| 271 |
+
self.to(device)
|
| 272 |
+
return self.passive_memory_usage()
|
| 273 |
+
|
| 274 |
+
|
| 275 |
+
def weight_decompose(
|
| 276 |
+
dora_scale, weight, lora_diff, alpha, strength, intermediate_dtype, function
|
| 277 |
+
):
|
| 278 |
+
dora_scale = comfy.model_management.cast_to_device(
|
| 279 |
+
dora_scale, weight.device, intermediate_dtype
|
| 280 |
+
)
|
| 281 |
+
lora_diff *= alpha
|
| 282 |
+
weight_calc = weight + function(lora_diff).type(weight.dtype)
|
| 283 |
+
|
| 284 |
+
wd_on_output_axis = dora_scale.shape[0] == weight_calc.shape[0]
|
| 285 |
+
if wd_on_output_axis:
|
| 286 |
+
weight_norm = (
|
| 287 |
+
weight.reshape(weight.shape[0], -1)
|
| 288 |
+
.norm(dim=1, keepdim=True)
|
| 289 |
+
.reshape(weight.shape[0], *[1] * (weight.dim() - 1))
|
| 290 |
+
)
|
| 291 |
+
else:
|
| 292 |
+
weight_norm = (
|
| 293 |
+
weight_calc.transpose(0, 1)
|
| 294 |
+
.reshape(weight_calc.shape[1], -1)
|
| 295 |
+
.norm(dim=1, keepdim=True)
|
| 296 |
+
.reshape(weight_calc.shape[1], *[1] * (weight_calc.dim() - 1))
|
| 297 |
+
.transpose(0, 1)
|
| 298 |
+
)
|
| 299 |
+
weight_norm = weight_norm + torch.finfo(weight.dtype).eps
|
| 300 |
+
|
| 301 |
+
weight_calc *= (dora_scale / weight_norm).type(weight.dtype)
|
| 302 |
+
if strength != 1.0:
|
| 303 |
+
weight_calc -= weight
|
| 304 |
+
weight += strength * (weight_calc)
|
| 305 |
+
else:
|
| 306 |
+
weight[:] = weight_calc
|
| 307 |
+
return weight
|
| 308 |
+
|
| 309 |
+
|
| 310 |
+
def pad_tensor_to_shape(tensor: torch.Tensor, new_shape: list[int]) -> torch.Tensor:
|
| 311 |
+
"""
|
| 312 |
+
Pad a tensor to a new shape with zeros.
|
| 313 |
+
|
| 314 |
+
Args:
|
| 315 |
+
tensor (torch.Tensor): The original tensor to be padded.
|
| 316 |
+
new_shape (List[int]): The desired shape of the padded tensor.
|
| 317 |
+
|
| 318 |
+
Returns:
|
| 319 |
+
torch.Tensor: A new tensor padded with zeros to the specified shape.
|
| 320 |
+
|
| 321 |
+
Note:
|
| 322 |
+
If the new shape is smaller than the original tensor in any dimension,
|
| 323 |
+
the original tensor will be truncated in that dimension.
|
| 324 |
+
"""
|
| 325 |
+
if any([new_shape[i] < tensor.shape[i] for i in range(len(new_shape))]):
|
| 326 |
+
raise ValueError(
|
| 327 |
+
"The new shape must be larger than the original tensor in all dimensions"
|
| 328 |
+
)
|
| 329 |
+
|
| 330 |
+
if len(new_shape) != len(tensor.shape):
|
| 331 |
+
raise ValueError(
|
| 332 |
+
"The new shape must have the same number of dimensions as the original tensor"
|
| 333 |
+
)
|
| 334 |
+
|
| 335 |
+
# Create a new tensor filled with zeros
|
| 336 |
+
padded_tensor = torch.zeros(new_shape, dtype=tensor.dtype, device=tensor.device)
|
| 337 |
+
|
| 338 |
+
# Create slicing tuples for both tensors
|
| 339 |
+
orig_slices = tuple(slice(0, dim) for dim in tensor.shape)
|
| 340 |
+
new_slices = tuple(slice(0, dim) for dim in tensor.shape)
|
| 341 |
+
|
| 342 |
+
# Copy the original tensor into the new tensor
|
| 343 |
+
padded_tensor[new_slices] = tensor[orig_slices]
|
| 344 |
+
|
| 345 |
+
return padded_tensor
|
| 346 |
+
|
| 347 |
+
|
| 348 |
+
def tucker_weight_from_conv(up, down, mid):
|
| 349 |
+
up = up.reshape(up.size(0), up.size(1))
|
| 350 |
+
down = down.reshape(down.size(0), down.size(1))
|
| 351 |
+
return torch.einsum("m n ..., i m, n j -> i j ...", mid, up, down)
|
| 352 |
+
|
| 353 |
+
|
| 354 |
+
def tucker_weight(wa, wb, t):
|
| 355 |
+
temp = torch.einsum("i j ..., j r -> i r ...", t, wb)
|
| 356 |
+
return torch.einsum("i j ..., i r -> r j ...", temp, wa)
|
| 357 |
+
|
| 358 |
+
|
| 359 |
+
def factorization(dimension: int, factor: int = -1) -> tuple[int, int]:
|
| 360 |
+
"""
|
| 361 |
+
return a tuple of two value of input dimension decomposed by the number closest to factor
|
| 362 |
+
second value is higher or equal than first value.
|
| 363 |
+
|
| 364 |
+
examples)
|
| 365 |
+
factor
|
| 366 |
+
-1 2 4 8 16 ...
|
| 367 |
+
127 -> 1, 127 127 -> 1, 127 127 -> 1, 127 127 -> 1, 127 127 -> 1, 127
|
| 368 |
+
128 -> 8, 16 128 -> 2, 64 128 -> 4, 32 128 -> 8, 16 128 -> 8, 16
|
| 369 |
+
250 -> 10, 25 250 -> 2, 125 250 -> 2, 125 250 -> 5, 50 250 -> 10, 25
|
| 370 |
+
360 -> 8, 45 360 -> 2, 180 360 -> 4, 90 360 -> 8, 45 360 -> 12, 30
|
| 371 |
+
512 -> 16, 32 512 -> 2, 256 512 -> 4, 128 512 -> 8, 64 512 -> 16, 32
|
| 372 |
+
1024 -> 32, 32 1024 -> 2, 512 1024 -> 4, 256 1024 -> 8, 128 1024 -> 16, 64
|
| 373 |
+
"""
|
| 374 |
+
|
| 375 |
+
if factor > 0 and (dimension % factor) == 0 and dimension >= factor**2:
|
| 376 |
+
m = factor
|
| 377 |
+
n = dimension // factor
|
| 378 |
+
if m > n:
|
| 379 |
+
n, m = m, n
|
| 380 |
+
return m, n
|
| 381 |
+
if factor < 0:
|
| 382 |
+
factor = dimension
|
| 383 |
+
m, n = 1, dimension
|
| 384 |
+
length = m + n
|
| 385 |
+
while m < n:
|
| 386 |
+
new_m = m + 1
|
| 387 |
+
while dimension % new_m != 0:
|
| 388 |
+
new_m += 1
|
| 389 |
+
new_n = dimension // new_m
|
| 390 |
+
if new_m + new_n > length or new_m > factor:
|
| 391 |
+
break
|
| 392 |
+
else:
|
| 393 |
+
m, n = new_m, new_n
|
| 394 |
+
if m > n:
|
| 395 |
+
n, m = m, n
|
| 396 |
+
return m, n
|
comfy/weight_adapter/boft.py
ADDED
|
@@ -0,0 +1,218 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from typing import Optional
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import comfy.model_management
|
| 6 |
+
from .base import WeightAdapterBase, weight_decompose
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
class BOFTAdapter(WeightAdapterBase):
|
| 10 |
+
name = "boft"
|
| 11 |
+
|
| 12 |
+
def __init__(self, loaded_keys, weights):
|
| 13 |
+
self.loaded_keys = loaded_keys
|
| 14 |
+
self.weights = weights
|
| 15 |
+
|
| 16 |
+
@classmethod
|
| 17 |
+
def load(
|
| 18 |
+
cls,
|
| 19 |
+
x: str,
|
| 20 |
+
lora: dict[str, torch.Tensor],
|
| 21 |
+
alpha: float,
|
| 22 |
+
dora_scale: torch.Tensor,
|
| 23 |
+
loaded_keys: set[str] = None,
|
| 24 |
+
) -> Optional["BOFTAdapter"]:
|
| 25 |
+
if loaded_keys is None:
|
| 26 |
+
loaded_keys = set()
|
| 27 |
+
blocks_name = "{}.oft_blocks".format(x)
|
| 28 |
+
rescale_name = "{}.rescale".format(x)
|
| 29 |
+
|
| 30 |
+
blocks = None
|
| 31 |
+
if blocks_name in lora.keys():
|
| 32 |
+
blocks = lora[blocks_name]
|
| 33 |
+
if blocks.ndim == 4:
|
| 34 |
+
loaded_keys.add(blocks_name)
|
| 35 |
+
else:
|
| 36 |
+
blocks = None
|
| 37 |
+
if blocks is None:
|
| 38 |
+
return None
|
| 39 |
+
|
| 40 |
+
rescale = None
|
| 41 |
+
if rescale_name in lora.keys():
|
| 42 |
+
rescale = lora[rescale_name]
|
| 43 |
+
loaded_keys.add(rescale_name)
|
| 44 |
+
|
| 45 |
+
weights = (blocks, rescale, alpha, dora_scale)
|
| 46 |
+
return cls(loaded_keys, weights)
|
| 47 |
+
|
| 48 |
+
def calculate_weight(
|
| 49 |
+
self,
|
| 50 |
+
weight,
|
| 51 |
+
key,
|
| 52 |
+
strength,
|
| 53 |
+
strength_model,
|
| 54 |
+
offset,
|
| 55 |
+
function,
|
| 56 |
+
intermediate_dtype=torch.float32,
|
| 57 |
+
original_weight=None,
|
| 58 |
+
):
|
| 59 |
+
v = self.weights
|
| 60 |
+
blocks = v[0]
|
| 61 |
+
rescale = v[1]
|
| 62 |
+
alpha = v[2]
|
| 63 |
+
dora_scale = v[3]
|
| 64 |
+
|
| 65 |
+
blocks = comfy.model_management.cast_to_device(
|
| 66 |
+
blocks, weight.device, intermediate_dtype
|
| 67 |
+
)
|
| 68 |
+
if rescale is not None:
|
| 69 |
+
rescale = comfy.model_management.cast_to_device(
|
| 70 |
+
rescale, weight.device, intermediate_dtype
|
| 71 |
+
)
|
| 72 |
+
|
| 73 |
+
boft_m, block_num, boft_b, *_ = blocks.shape
|
| 74 |
+
|
| 75 |
+
try:
|
| 76 |
+
# Get r
|
| 77 |
+
I = torch.eye(boft_b, device=blocks.device, dtype=blocks.dtype)
|
| 78 |
+
# for Q = -Q^T
|
| 79 |
+
q = blocks - blocks.transpose(-1, -2)
|
| 80 |
+
normed_q = q
|
| 81 |
+
if alpha > 0: # alpha in boft/bboft is for constraint
|
| 82 |
+
q_norm = torch.norm(q) + 1e-8
|
| 83 |
+
if q_norm > alpha:
|
| 84 |
+
normed_q = q * alpha / q_norm
|
| 85 |
+
# use float() to prevent unsupported type in .inverse()
|
| 86 |
+
r = (I + normed_q) @ (I - normed_q).float().inverse()
|
| 87 |
+
r = r.to(weight)
|
| 88 |
+
inp = org = weight
|
| 89 |
+
|
| 90 |
+
r_b = boft_b // 2
|
| 91 |
+
for i in range(boft_m):
|
| 92 |
+
bi = r[i]
|
| 93 |
+
g = 2
|
| 94 |
+
k = 2**i * r_b
|
| 95 |
+
if strength != 1:
|
| 96 |
+
bi = bi * strength + (1 - strength) * I
|
| 97 |
+
inp = (
|
| 98 |
+
inp.unflatten(0, (-1, g, k))
|
| 99 |
+
.transpose(1, 2)
|
| 100 |
+
.flatten(0, 2)
|
| 101 |
+
.unflatten(0, (-1, boft_b))
|
| 102 |
+
)
|
| 103 |
+
inp = torch.einsum("b i j, b j ...-> b i ...", bi, inp)
|
| 104 |
+
inp = (
|
| 105 |
+
inp.flatten(0, 1)
|
| 106 |
+
.unflatten(0, (-1, k, g))
|
| 107 |
+
.transpose(1, 2)
|
| 108 |
+
.flatten(0, 2)
|
| 109 |
+
)
|
| 110 |
+
|
| 111 |
+
if rescale is not None:
|
| 112 |
+
inp = inp * rescale
|
| 113 |
+
|
| 114 |
+
lora_diff = inp - org
|
| 115 |
+
lora_diff = comfy.model_management.cast_to_device(
|
| 116 |
+
lora_diff, weight.device, intermediate_dtype
|
| 117 |
+
)
|
| 118 |
+
if dora_scale is not None:
|
| 119 |
+
weight = weight_decompose(
|
| 120 |
+
dora_scale,
|
| 121 |
+
weight,
|
| 122 |
+
lora_diff,
|
| 123 |
+
alpha,
|
| 124 |
+
strength,
|
| 125 |
+
intermediate_dtype,
|
| 126 |
+
function,
|
| 127 |
+
)
|
| 128 |
+
else:
|
| 129 |
+
weight += function((strength * lora_diff).type(weight.dtype))
|
| 130 |
+
except Exception as e:
|
| 131 |
+
logging.error("ERROR {} {} {}".format(self.name, key, e))
|
| 132 |
+
return weight
|
| 133 |
+
|
| 134 |
+
def _get_orthogonal_matrices(self, device, dtype):
|
| 135 |
+
"""Compute the orthogonal rotation matrices R from BOFT blocks."""
|
| 136 |
+
v = self.weights
|
| 137 |
+
blocks = v[0].to(device=device, dtype=dtype)
|
| 138 |
+
alpha = v[2]
|
| 139 |
+
if alpha is None:
|
| 140 |
+
alpha = 0
|
| 141 |
+
|
| 142 |
+
boft_m, block_num, boft_b, _ = blocks.shape
|
| 143 |
+
I = torch.eye(boft_b, device=device, dtype=dtype)
|
| 144 |
+
|
| 145 |
+
# Q = blocks - blocks^T (skew-symmetric)
|
| 146 |
+
q = blocks - blocks.transpose(-1, -2)
|
| 147 |
+
normed_q = q
|
| 148 |
+
|
| 149 |
+
# Apply constraint if alpha > 0
|
| 150 |
+
if alpha > 0:
|
| 151 |
+
q_norm = torch.norm(q) + 1e-8
|
| 152 |
+
if q_norm > alpha:
|
| 153 |
+
normed_q = q * alpha / q_norm
|
| 154 |
+
|
| 155 |
+
# Cayley transform: R = (I + Q)(I - Q)^-1
|
| 156 |
+
r = (I + normed_q) @ (I - normed_q).float().inverse()
|
| 157 |
+
return r, boft_m, boft_b
|
| 158 |
+
|
| 159 |
+
def g(self, y: torch.Tensor) -> torch.Tensor:
|
| 160 |
+
"""
|
| 161 |
+
Output transformation for BOFT: applies butterfly orthogonal transform.
|
| 162 |
+
|
| 163 |
+
BOFT uses multiple stages of butterfly-structured orthogonal transforms.
|
| 164 |
+
|
| 165 |
+
Reference: LyCORIS ButterflyOFTModule._bypass_forward
|
| 166 |
+
"""
|
| 167 |
+
v = self.weights
|
| 168 |
+
rescale = v[1]
|
| 169 |
+
|
| 170 |
+
r, boft_m, boft_b = self._get_orthogonal_matrices(y.device, y.dtype)
|
| 171 |
+
r_b = boft_b // 2
|
| 172 |
+
|
| 173 |
+
# Apply multiplier
|
| 174 |
+
multiplier = getattr(self, "multiplier", 1.0)
|
| 175 |
+
I = torch.eye(boft_b, device=y.device, dtype=y.dtype)
|
| 176 |
+
|
| 177 |
+
# Use module info from bypass injection to determine conv vs linear
|
| 178 |
+
is_conv = getattr(self, "is_conv", y.dim() > 2)
|
| 179 |
+
|
| 180 |
+
if is_conv:
|
| 181 |
+
# Conv output: (N, C, H, W, ...) -> transpose to (N, H, W, ..., C)
|
| 182 |
+
y = y.transpose(1, -1)
|
| 183 |
+
|
| 184 |
+
# Apply butterfly transform stages
|
| 185 |
+
inp = y
|
| 186 |
+
for i in range(boft_m):
|
| 187 |
+
bi = r[i] # (block_num, boft_b, boft_b)
|
| 188 |
+
g = 2
|
| 189 |
+
k = 2**i * r_b
|
| 190 |
+
|
| 191 |
+
# Interpolate with identity based on multiplier
|
| 192 |
+
if multiplier != 1:
|
| 193 |
+
bi = bi * multiplier + (1 - multiplier) * I
|
| 194 |
+
|
| 195 |
+
# Reshape for butterfly: unflatten last dim, transpose, flatten, unflatten
|
| 196 |
+
inp = (
|
| 197 |
+
inp.unflatten(-1, (-1, g, k))
|
| 198 |
+
.transpose(-2, -1)
|
| 199 |
+
.flatten(-3)
|
| 200 |
+
.unflatten(-1, (-1, boft_b))
|
| 201 |
+
)
|
| 202 |
+
# Apply block-diagonal orthogonal transform
|
| 203 |
+
inp = torch.einsum("b i j, ... b j -> ... b i", bi, inp)
|
| 204 |
+
# Reshape back
|
| 205 |
+
inp = (
|
| 206 |
+
inp.flatten(-2).unflatten(-1, (-1, k, g)).transpose(-2, -1).flatten(-3)
|
| 207 |
+
)
|
| 208 |
+
|
| 209 |
+
# Apply rescale if present
|
| 210 |
+
if rescale is not None:
|
| 211 |
+
rescale = rescale.to(device=y.device, dtype=y.dtype)
|
| 212 |
+
inp = inp * rescale.transpose(0, -1)
|
| 213 |
+
|
| 214 |
+
if is_conv:
|
| 215 |
+
# Transpose back: (N, H, W, ..., C) -> (N, C, H, W, ...)
|
| 216 |
+
inp = inp.transpose(1, -1)
|
| 217 |
+
|
| 218 |
+
return inp
|
comfy/weight_adapter/bypass.py
ADDED
|
@@ -0,0 +1,441 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Bypass mode implementation for weight adapters (LoRA, LoKr, LoHa, etc.)
|
| 3 |
+
|
| 4 |
+
Bypass mode applies adapters during forward pass without modifying base weights:
|
| 5 |
+
bypass(f)(x) = g(f(x) + h(x))
|
| 6 |
+
|
| 7 |
+
Where:
|
| 8 |
+
- f(x): Original layer forward
|
| 9 |
+
- h(x): Additive component from adapter (LoRA path)
|
| 10 |
+
- g(y): Output transformation (identity for most adapters)
|
| 11 |
+
|
| 12 |
+
This is useful for:
|
| 13 |
+
- Training with gradient checkpointing
|
| 14 |
+
- Avoiding weight modifications when weights are offloaded
|
| 15 |
+
- Supporting multiple adapters with different strengths dynamically
|
| 16 |
+
"""
|
| 17 |
+
|
| 18 |
+
import logging
|
| 19 |
+
from typing import Optional, Union
|
| 20 |
+
|
| 21 |
+
import torch
|
| 22 |
+
import torch.nn as nn
|
| 23 |
+
|
| 24 |
+
import comfy.model_management
|
| 25 |
+
from .base import WeightAdapterBase, WeightAdapterTrainBase
|
| 26 |
+
from comfy.patcher_extension import PatcherInjection
|
| 27 |
+
|
| 28 |
+
# Type alias for adapters that support bypass mode
|
| 29 |
+
BypassAdapter = Union[WeightAdapterBase, WeightAdapterTrainBase]
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def get_module_type_info(module: nn.Module) -> dict:
|
| 33 |
+
"""
|
| 34 |
+
Determine module type and extract conv parameters from module class.
|
| 35 |
+
|
| 36 |
+
This is more reliable than checking weight.ndim, especially for quantized layers
|
| 37 |
+
where weight shape might be different.
|
| 38 |
+
|
| 39 |
+
Returns:
|
| 40 |
+
dict with keys: is_conv, conv_dim, stride, padding, dilation, groups
|
| 41 |
+
"""
|
| 42 |
+
info = {
|
| 43 |
+
"is_conv": False,
|
| 44 |
+
"conv_dim": 0,
|
| 45 |
+
"stride": (1,),
|
| 46 |
+
"padding": (0,),
|
| 47 |
+
"dilation": (1,),
|
| 48 |
+
"groups": 1,
|
| 49 |
+
"kernel_size": (1,),
|
| 50 |
+
"in_channels": None,
|
| 51 |
+
"out_channels": None,
|
| 52 |
+
}
|
| 53 |
+
|
| 54 |
+
# Determine conv type
|
| 55 |
+
if isinstance(module, nn.Conv1d):
|
| 56 |
+
info["is_conv"] = True
|
| 57 |
+
info["conv_dim"] = 1
|
| 58 |
+
elif isinstance(module, nn.Conv2d):
|
| 59 |
+
info["is_conv"] = True
|
| 60 |
+
info["conv_dim"] = 2
|
| 61 |
+
elif isinstance(module, nn.Conv3d):
|
| 62 |
+
info["is_conv"] = True
|
| 63 |
+
info["conv_dim"] = 3
|
| 64 |
+
elif isinstance(module, nn.Linear):
|
| 65 |
+
info["is_conv"] = False
|
| 66 |
+
info["conv_dim"] = 0
|
| 67 |
+
else:
|
| 68 |
+
# Try to infer from class name for custom/quantized layers
|
| 69 |
+
class_name = type(module).__name__.lower()
|
| 70 |
+
if "conv3d" in class_name:
|
| 71 |
+
info["is_conv"] = True
|
| 72 |
+
info["conv_dim"] = 3
|
| 73 |
+
elif "conv2d" in class_name:
|
| 74 |
+
info["is_conv"] = True
|
| 75 |
+
info["conv_dim"] = 2
|
| 76 |
+
elif "conv1d" in class_name:
|
| 77 |
+
info["is_conv"] = True
|
| 78 |
+
info["conv_dim"] = 1
|
| 79 |
+
elif "conv" in class_name:
|
| 80 |
+
info["is_conv"] = True
|
| 81 |
+
info["conv_dim"] = 2
|
| 82 |
+
|
| 83 |
+
# Extract conv parameters if it's a conv layer
|
| 84 |
+
if info["is_conv"]:
|
| 85 |
+
# Try to get stride, padding, dilation, groups, kernel_size from module
|
| 86 |
+
info["stride"] = getattr(module, "stride", (1,) * info["conv_dim"])
|
| 87 |
+
info["padding"] = getattr(module, "padding", (0,) * info["conv_dim"])
|
| 88 |
+
info["dilation"] = getattr(module, "dilation", (1,) * info["conv_dim"])
|
| 89 |
+
info["groups"] = getattr(module, "groups", 1)
|
| 90 |
+
info["kernel_size"] = getattr(module, "kernel_size", (1,) * info["conv_dim"])
|
| 91 |
+
info["in_channels"] = getattr(module, "in_channels", None)
|
| 92 |
+
info["out_channels"] = getattr(module, "out_channels", None)
|
| 93 |
+
|
| 94 |
+
# Ensure they're tuples
|
| 95 |
+
if isinstance(info["stride"], int):
|
| 96 |
+
info["stride"] = (info["stride"],) * info["conv_dim"]
|
| 97 |
+
if isinstance(info["padding"], int):
|
| 98 |
+
info["padding"] = (info["padding"],) * info["conv_dim"]
|
| 99 |
+
if isinstance(info["dilation"], int):
|
| 100 |
+
info["dilation"] = (info["dilation"],) * info["conv_dim"]
|
| 101 |
+
if isinstance(info["kernel_size"], int):
|
| 102 |
+
info["kernel_size"] = (info["kernel_size"],) * info["conv_dim"]
|
| 103 |
+
|
| 104 |
+
return info
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
class BypassForwardHook:
|
| 108 |
+
"""
|
| 109 |
+
Hook that wraps a layer's forward to apply adapter in bypass mode.
|
| 110 |
+
|
| 111 |
+
Stores the original forward and replaces it with bypass version.
|
| 112 |
+
|
| 113 |
+
Supports both:
|
| 114 |
+
- WeightAdapterBase: Inference adapters (uses self.weights tuple)
|
| 115 |
+
- WeightAdapterTrainBase: Training adapters (nn.Module with parameters)
|
| 116 |
+
"""
|
| 117 |
+
|
| 118 |
+
def __init__(
|
| 119 |
+
self,
|
| 120 |
+
module: nn.Module,
|
| 121 |
+
adapter: BypassAdapter,
|
| 122 |
+
multiplier: float = 1.0,
|
| 123 |
+
):
|
| 124 |
+
self.module = module
|
| 125 |
+
self.adapter = adapter
|
| 126 |
+
self.multiplier = multiplier
|
| 127 |
+
self.original_forward = None
|
| 128 |
+
|
| 129 |
+
# Determine layer type and conv params from module class (works for quantized layers)
|
| 130 |
+
module_info = get_module_type_info(module)
|
| 131 |
+
|
| 132 |
+
# Set multiplier and layer type info on adapter for use in h()
|
| 133 |
+
adapter.multiplier = multiplier
|
| 134 |
+
adapter.is_conv = module_info["is_conv"]
|
| 135 |
+
adapter.conv_dim = module_info["conv_dim"]
|
| 136 |
+
adapter.kernel_size = module_info["kernel_size"]
|
| 137 |
+
adapter.in_channels = module_info["in_channels"]
|
| 138 |
+
adapter.out_channels = module_info["out_channels"]
|
| 139 |
+
# Store kw_dict for conv operations (like LyCORIS extra_args)
|
| 140 |
+
if module_info["is_conv"]:
|
| 141 |
+
adapter.kw_dict = {
|
| 142 |
+
"stride": module_info["stride"],
|
| 143 |
+
"padding": module_info["padding"],
|
| 144 |
+
"dilation": module_info["dilation"],
|
| 145 |
+
"groups": module_info["groups"],
|
| 146 |
+
}
|
| 147 |
+
else:
|
| 148 |
+
adapter.kw_dict = {}
|
| 149 |
+
|
| 150 |
+
def _bypass_forward(self, x: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
| 151 |
+
"""Bypass forward: uses adapter's bypass_forward or default g(f(x) + h(x))
|
| 152 |
+
|
| 153 |
+
Note:
|
| 154 |
+
Bypass mode does NOT access original model weights (org_weight).
|
| 155 |
+
This is intentional - bypass mode is designed for quantized models
|
| 156 |
+
where weights may not be in a usable format. All necessary shape
|
| 157 |
+
information is provided via adapter attributes set during inject().
|
| 158 |
+
"""
|
| 159 |
+
# Check if adapter has custom bypass_forward (e.g., GLoRA)
|
| 160 |
+
adapter_bypass = getattr(self.adapter, "bypass_forward", None)
|
| 161 |
+
if adapter_bypass is not None:
|
| 162 |
+
# Check if it's overridden (not the base class default)
|
| 163 |
+
# Need to check both base classes since adapter could be either type
|
| 164 |
+
adapter_type = type(self.adapter)
|
| 165 |
+
is_default_bypass = (
|
| 166 |
+
adapter_type.bypass_forward is WeightAdapterBase.bypass_forward
|
| 167 |
+
or adapter_type.bypass_forward is WeightAdapterTrainBase.bypass_forward
|
| 168 |
+
)
|
| 169 |
+
if not is_default_bypass:
|
| 170 |
+
return adapter_bypass(self.original_forward, x, *args, **kwargs)
|
| 171 |
+
|
| 172 |
+
# Default bypass: g(f(x) + h(x, f(x)))
|
| 173 |
+
base_out = self.original_forward(x, *args, **kwargs)
|
| 174 |
+
h_out = self.adapter.h(x, base_out)
|
| 175 |
+
return self.adapter.g(base_out + h_out)
|
| 176 |
+
|
| 177 |
+
def inject(self):
|
| 178 |
+
"""Replace module forward with bypass version."""
|
| 179 |
+
if self.original_forward is not None:
|
| 180 |
+
logging.debug(
|
| 181 |
+
f"[BypassHook] Already injected for {type(self.module).__name__}"
|
| 182 |
+
)
|
| 183 |
+
return # Already injected
|
| 184 |
+
|
| 185 |
+
# Move adapter weights to compute device (GPU)
|
| 186 |
+
# Use get_torch_device() instead of module.weight.device because
|
| 187 |
+
# with offloading, module weights may be on CPU while compute happens on GPU
|
| 188 |
+
device = comfy.model_management.get_torch_device()
|
| 189 |
+
|
| 190 |
+
# Get dtype from module weight if available
|
| 191 |
+
dtype = None
|
| 192 |
+
if hasattr(self.module, "weight") and self.module.weight is not None:
|
| 193 |
+
dtype = self.module.weight.dtype
|
| 194 |
+
|
| 195 |
+
# Only use dtype if it's a standard float type, not quantized
|
| 196 |
+
if dtype is not None and dtype not in (torch.float32, torch.float16, torch.bfloat16):
|
| 197 |
+
dtype = None
|
| 198 |
+
|
| 199 |
+
self._move_adapter_weights_to_device(device, dtype)
|
| 200 |
+
|
| 201 |
+
self.original_forward = self.module.forward
|
| 202 |
+
self.module.forward = self._bypass_forward
|
| 203 |
+
logging.debug(
|
| 204 |
+
f"[BypassHook] Injected bypass forward for {type(self.module).__name__} (adapter={type(self.adapter).__name__})"
|
| 205 |
+
)
|
| 206 |
+
|
| 207 |
+
def _move_adapter_weights_to_device(self, device, dtype=None):
|
| 208 |
+
"""Move adapter weights to specified device to avoid per-forward transfers.
|
| 209 |
+
|
| 210 |
+
Handles both:
|
| 211 |
+
- WeightAdapterBase: has self.weights tuple of tensors
|
| 212 |
+
- WeightAdapterTrainBase: nn.Module with parameters, uses .to() method
|
| 213 |
+
"""
|
| 214 |
+
adapter = self.adapter
|
| 215 |
+
|
| 216 |
+
# Check if adapter is an nn.Module (WeightAdapterTrainBase)
|
| 217 |
+
if isinstance(adapter, nn.Module):
|
| 218 |
+
# In training mode we don't touch dtype as trainer will handle it
|
| 219 |
+
adapter.to(device=device)
|
| 220 |
+
logging.debug(
|
| 221 |
+
f"[BypassHook] Moved training adapter (nn.Module) to {device}"
|
| 222 |
+
)
|
| 223 |
+
return
|
| 224 |
+
|
| 225 |
+
# WeightAdapterBase: handle self.weights tuple
|
| 226 |
+
if not hasattr(adapter, "weights") or adapter.weights is None:
|
| 227 |
+
return
|
| 228 |
+
|
| 229 |
+
weights = adapter.weights
|
| 230 |
+
if isinstance(weights, (list, tuple)):
|
| 231 |
+
new_weights = []
|
| 232 |
+
for w in weights:
|
| 233 |
+
if isinstance(w, torch.Tensor):
|
| 234 |
+
if dtype is not None:
|
| 235 |
+
new_weights.append(w.to(device=device, dtype=dtype))
|
| 236 |
+
else:
|
| 237 |
+
new_weights.append(w.to(device=device))
|
| 238 |
+
else:
|
| 239 |
+
new_weights.append(w)
|
| 240 |
+
adapter.weights = (
|
| 241 |
+
tuple(new_weights) if isinstance(weights, tuple) else new_weights
|
| 242 |
+
)
|
| 243 |
+
elif isinstance(weights, torch.Tensor):
|
| 244 |
+
if dtype is not None:
|
| 245 |
+
adapter.weights = weights.to(device=device, dtype=dtype)
|
| 246 |
+
else:
|
| 247 |
+
adapter.weights = weights.to(device=device)
|
| 248 |
+
|
| 249 |
+
logging.debug(f"[BypassHook] Moved adapter weights to {device}")
|
| 250 |
+
|
| 251 |
+
def eject(self):
|
| 252 |
+
"""Restore original module forward."""
|
| 253 |
+
if self.original_forward is None:
|
| 254 |
+
logging.debug(f"[BypassHook] Not injected for {type(self.module).__name__}")
|
| 255 |
+
return # Not injected
|
| 256 |
+
|
| 257 |
+
self.module.forward = self.original_forward
|
| 258 |
+
self.original_forward = None
|
| 259 |
+
logging.debug(
|
| 260 |
+
f"[BypassHook] Ejected bypass forward for {type(self.module).__name__}"
|
| 261 |
+
)
|
| 262 |
+
|
| 263 |
+
|
| 264 |
+
class BypassInjectionManager:
|
| 265 |
+
"""
|
| 266 |
+
Manages bypass mode injection for a collection of adapters.
|
| 267 |
+
|
| 268 |
+
Creates PatcherInjection objects that can be used with ModelPatcher.
|
| 269 |
+
|
| 270 |
+
Supports both inference adapters (WeightAdapterBase) and training adapters
|
| 271 |
+
(WeightAdapterTrainBase).
|
| 272 |
+
|
| 273 |
+
Usage:
|
| 274 |
+
manager = BypassInjectionManager()
|
| 275 |
+
manager.add_adapter("model.layers.0.self_attn.q_proj", lora_adapter, strength=0.8)
|
| 276 |
+
manager.add_adapter("model.layers.0.self_attn.k_proj", lora_adapter, strength=0.8)
|
| 277 |
+
|
| 278 |
+
injections = manager.create_injections(model)
|
| 279 |
+
model_patcher.set_injections("bypass_lora", injections)
|
| 280 |
+
"""
|
| 281 |
+
|
| 282 |
+
def __init__(self):
|
| 283 |
+
self.adapters: dict[str, tuple[BypassAdapter, float]] = {}
|
| 284 |
+
self.hooks: list[BypassForwardHook] = []
|
| 285 |
+
|
| 286 |
+
def add_adapter(
|
| 287 |
+
self,
|
| 288 |
+
key: str,
|
| 289 |
+
adapter: BypassAdapter,
|
| 290 |
+
strength: float = 1.0,
|
| 291 |
+
):
|
| 292 |
+
"""
|
| 293 |
+
Add an adapter for a specific weight key.
|
| 294 |
+
|
| 295 |
+
Args:
|
| 296 |
+
key: Weight key (e.g., "model.layers.0.self_attn.q_proj.weight")
|
| 297 |
+
adapter: The weight adapter (LoRAAdapter, LoKrAdapter, etc.)
|
| 298 |
+
strength: Multiplier for adapter effect
|
| 299 |
+
"""
|
| 300 |
+
# Remove .weight suffix if present for module lookup
|
| 301 |
+
module_key = key
|
| 302 |
+
if module_key.endswith(".weight"):
|
| 303 |
+
module_key = module_key[:-7]
|
| 304 |
+
logging.debug(
|
| 305 |
+
f"[BypassManager] Stripped .weight suffix: {key} -> {module_key}"
|
| 306 |
+
)
|
| 307 |
+
|
| 308 |
+
self.adapters[module_key] = (adapter, strength)
|
| 309 |
+
logging.debug(
|
| 310 |
+
f"[BypassManager] Added adapter: {module_key} (type={type(adapter).__name__}, strength={strength})"
|
| 311 |
+
)
|
| 312 |
+
|
| 313 |
+
def clear_adapters(self):
|
| 314 |
+
"""Remove all adapters."""
|
| 315 |
+
self.adapters.clear()
|
| 316 |
+
|
| 317 |
+
def _get_module_by_key(self, model: nn.Module, key: str) -> Optional[nn.Module]:
|
| 318 |
+
"""Get a submodule by dot-separated key."""
|
| 319 |
+
parts = key.split(".")
|
| 320 |
+
module = model
|
| 321 |
+
try:
|
| 322 |
+
for i, part in enumerate(parts):
|
| 323 |
+
if part.isdigit():
|
| 324 |
+
module = module[int(part)]
|
| 325 |
+
else:
|
| 326 |
+
module = getattr(module, part)
|
| 327 |
+
logging.debug(
|
| 328 |
+
f"[BypassManager] Found module for key {key}: {type(module).__name__}"
|
| 329 |
+
)
|
| 330 |
+
return module
|
| 331 |
+
except (AttributeError, IndexError, KeyError) as e:
|
| 332 |
+
logging.error(f"[BypassManager] Failed to find module for key {key}: {e}")
|
| 333 |
+
logging.error(
|
| 334 |
+
f"[BypassManager] Failed at part index {i}, part={part}, current module type={type(module).__name__}"
|
| 335 |
+
)
|
| 336 |
+
return None
|
| 337 |
+
|
| 338 |
+
def create_injections(self, model: nn.Module) -> list[PatcherInjection]:
|
| 339 |
+
"""
|
| 340 |
+
Create PatcherInjection objects for all registered adapters.
|
| 341 |
+
|
| 342 |
+
Args:
|
| 343 |
+
model: The model to inject into (e.g., model_patcher.model)
|
| 344 |
+
|
| 345 |
+
Returns:
|
| 346 |
+
List of PatcherInjection objects to use with model_patcher.set_injections()
|
| 347 |
+
"""
|
| 348 |
+
self.hooks.clear()
|
| 349 |
+
|
| 350 |
+
logging.debug(
|
| 351 |
+
f"[BypassManager] create_injections called with {len(self.adapters)} adapters"
|
| 352 |
+
)
|
| 353 |
+
logging.debug(f"[BypassManager] Model type: {type(model).__name__}")
|
| 354 |
+
|
| 355 |
+
for key, (adapter, strength) in self.adapters.items():
|
| 356 |
+
logging.debug(f"[BypassManager] Looking for module: {key}")
|
| 357 |
+
module = self._get_module_by_key(model, key)
|
| 358 |
+
|
| 359 |
+
if module is None:
|
| 360 |
+
logging.warning(f"[BypassManager] Module not found for key {key}")
|
| 361 |
+
continue
|
| 362 |
+
|
| 363 |
+
if not hasattr(module, "weight"):
|
| 364 |
+
logging.warning(
|
| 365 |
+
f"[BypassManager] Module {key} has no weight attribute (type={type(module).__name__})"
|
| 366 |
+
)
|
| 367 |
+
continue
|
| 368 |
+
|
| 369 |
+
logging.debug(
|
| 370 |
+
f"[BypassManager] Creating hook for {key} (module type={type(module).__name__}, weight shape={module.weight.shape})"
|
| 371 |
+
)
|
| 372 |
+
hook = BypassForwardHook(module, adapter, multiplier=strength)
|
| 373 |
+
self.hooks.append(hook)
|
| 374 |
+
|
| 375 |
+
logging.debug(f"[BypassManager] Created {len(self.hooks)} hooks")
|
| 376 |
+
|
| 377 |
+
# Create single injection that manages all hooks
|
| 378 |
+
def inject_all(model_patcher):
|
| 379 |
+
logging.debug(
|
| 380 |
+
f"[BypassManager] inject_all called, injecting {len(self.hooks)} hooks"
|
| 381 |
+
)
|
| 382 |
+
for hook in self.hooks:
|
| 383 |
+
hook.inject()
|
| 384 |
+
logging.debug(
|
| 385 |
+
f"[BypassManager] Injected hook for {type(hook.module).__name__}"
|
| 386 |
+
)
|
| 387 |
+
|
| 388 |
+
def eject_all(model_patcher):
|
| 389 |
+
logging.debug(
|
| 390 |
+
f"[BypassManager] eject_all called, ejecting {len(self.hooks)} hooks"
|
| 391 |
+
)
|
| 392 |
+
for hook in self.hooks:
|
| 393 |
+
hook.eject()
|
| 394 |
+
|
| 395 |
+
return [PatcherInjection(inject=inject_all, eject=eject_all)]
|
| 396 |
+
|
| 397 |
+
def get_hook_count(self) -> int:
|
| 398 |
+
"""Return number of hooks that will be/are injected."""
|
| 399 |
+
return len(self.hooks)
|
| 400 |
+
|
| 401 |
+
|
| 402 |
+
def create_bypass_injections_from_patches(
|
| 403 |
+
model: nn.Module,
|
| 404 |
+
patches: dict,
|
| 405 |
+
strength: float = 1.0,
|
| 406 |
+
) -> list[PatcherInjection]:
|
| 407 |
+
"""
|
| 408 |
+
Convenience function to create bypass injections from a patches dict.
|
| 409 |
+
|
| 410 |
+
This is useful when you have patches in the format used by model_patcher.add_patches()
|
| 411 |
+
and want to apply them in bypass mode instead.
|
| 412 |
+
|
| 413 |
+
Args:
|
| 414 |
+
model: The model to inject into
|
| 415 |
+
patches: Dict mapping weight keys to adapter data
|
| 416 |
+
strength: Global strength multiplier
|
| 417 |
+
|
| 418 |
+
Returns:
|
| 419 |
+
List of PatcherInjection objects
|
| 420 |
+
"""
|
| 421 |
+
manager = BypassInjectionManager()
|
| 422 |
+
|
| 423 |
+
for key, patch_list in patches.items():
|
| 424 |
+
if not patch_list:
|
| 425 |
+
continue
|
| 426 |
+
|
| 427 |
+
# patches format: list of (strength_patch, patch_data, strength_model, offset, function)
|
| 428 |
+
for patch in patch_list:
|
| 429 |
+
patch_strength, patch_data, strength_model, offset, function = patch
|
| 430 |
+
|
| 431 |
+
# patch_data should be a WeightAdapterBase/WeightAdapterTrainBase or tuple
|
| 432 |
+
if isinstance(patch_data, (WeightAdapterBase, WeightAdapterTrainBase)):
|
| 433 |
+
adapter = patch_data
|
| 434 |
+
else:
|
| 435 |
+
# Skip non-adapter patches
|
| 436 |
+
continue
|
| 437 |
+
|
| 438 |
+
combined_strength = strength * patch_strength
|
| 439 |
+
manager.add_adapter(key, adapter, strength=combined_strength)
|
| 440 |
+
|
| 441 |
+
return manager.create_injections(model)
|
comfy/weight_adapter/glora.py
ADDED
|
@@ -0,0 +1,290 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from typing import Callable, Optional
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import torch.nn.functional as F
|
| 6 |
+
import comfy.model_management
|
| 7 |
+
from .base import WeightAdapterBase, weight_decompose
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class GLoRAAdapter(WeightAdapterBase):
|
| 11 |
+
name = "glora"
|
| 12 |
+
|
| 13 |
+
def __init__(self, loaded_keys, weights):
|
| 14 |
+
self.loaded_keys = loaded_keys
|
| 15 |
+
self.weights = weights
|
| 16 |
+
|
| 17 |
+
@classmethod
|
| 18 |
+
def load(
|
| 19 |
+
cls,
|
| 20 |
+
x: str,
|
| 21 |
+
lora: dict[str, torch.Tensor],
|
| 22 |
+
alpha: float,
|
| 23 |
+
dora_scale: torch.Tensor,
|
| 24 |
+
loaded_keys: set[str] = None,
|
| 25 |
+
) -> Optional["GLoRAAdapter"]:
|
| 26 |
+
if loaded_keys is None:
|
| 27 |
+
loaded_keys = set()
|
| 28 |
+
a1_name = "{}.a1.weight".format(x)
|
| 29 |
+
a2_name = "{}.a2.weight".format(x)
|
| 30 |
+
b1_name = "{}.b1.weight".format(x)
|
| 31 |
+
b2_name = "{}.b2.weight".format(x)
|
| 32 |
+
if a1_name in lora:
|
| 33 |
+
weights = (
|
| 34 |
+
lora[a1_name],
|
| 35 |
+
lora[a2_name],
|
| 36 |
+
lora[b1_name],
|
| 37 |
+
lora[b2_name],
|
| 38 |
+
alpha,
|
| 39 |
+
dora_scale,
|
| 40 |
+
)
|
| 41 |
+
loaded_keys.add(a1_name)
|
| 42 |
+
loaded_keys.add(a2_name)
|
| 43 |
+
loaded_keys.add(b1_name)
|
| 44 |
+
loaded_keys.add(b2_name)
|
| 45 |
+
return cls(loaded_keys, weights)
|
| 46 |
+
else:
|
| 47 |
+
return None
|
| 48 |
+
|
| 49 |
+
def calculate_weight(
|
| 50 |
+
self,
|
| 51 |
+
weight,
|
| 52 |
+
key,
|
| 53 |
+
strength,
|
| 54 |
+
strength_model,
|
| 55 |
+
offset,
|
| 56 |
+
function,
|
| 57 |
+
intermediate_dtype=torch.float32,
|
| 58 |
+
original_weight=None,
|
| 59 |
+
):
|
| 60 |
+
v = self.weights
|
| 61 |
+
dora_scale = v[5]
|
| 62 |
+
|
| 63 |
+
old_glora = False
|
| 64 |
+
if v[3].shape[1] == v[2].shape[0] == v[0].shape[0] == v[1].shape[1]:
|
| 65 |
+
rank = v[0].shape[0]
|
| 66 |
+
old_glora = True
|
| 67 |
+
|
| 68 |
+
if v[3].shape[0] == v[2].shape[1] == v[0].shape[1] == v[1].shape[0]:
|
| 69 |
+
if (
|
| 70 |
+
old_glora
|
| 71 |
+
and v[1].shape[0] == weight.shape[0]
|
| 72 |
+
and weight.shape[0] == weight.shape[1]
|
| 73 |
+
):
|
| 74 |
+
pass
|
| 75 |
+
else:
|
| 76 |
+
old_glora = False
|
| 77 |
+
rank = v[1].shape[0]
|
| 78 |
+
|
| 79 |
+
a1 = comfy.model_management.cast_to_device(
|
| 80 |
+
v[0].flatten(start_dim=1), weight.device, intermediate_dtype
|
| 81 |
+
)
|
| 82 |
+
a2 = comfy.model_management.cast_to_device(
|
| 83 |
+
v[1].flatten(start_dim=1), weight.device, intermediate_dtype
|
| 84 |
+
)
|
| 85 |
+
b1 = comfy.model_management.cast_to_device(
|
| 86 |
+
v[2].flatten(start_dim=1), weight.device, intermediate_dtype
|
| 87 |
+
)
|
| 88 |
+
b2 = comfy.model_management.cast_to_device(
|
| 89 |
+
v[3].flatten(start_dim=1), weight.device, intermediate_dtype
|
| 90 |
+
)
|
| 91 |
+
|
| 92 |
+
if v[4] is not None:
|
| 93 |
+
alpha = v[4] / rank
|
| 94 |
+
else:
|
| 95 |
+
alpha = 1.0
|
| 96 |
+
|
| 97 |
+
try:
|
| 98 |
+
if old_glora:
|
| 99 |
+
lora_diff = (
|
| 100 |
+
torch.mm(b2, b1)
|
| 101 |
+
+ torch.mm(
|
| 102 |
+
torch.mm(
|
| 103 |
+
weight.flatten(start_dim=1).to(dtype=intermediate_dtype), a2
|
| 104 |
+
),
|
| 105 |
+
a1,
|
| 106 |
+
)
|
| 107 |
+
).reshape(
|
| 108 |
+
weight.shape
|
| 109 |
+
) # old lycoris glora
|
| 110 |
+
else:
|
| 111 |
+
if weight.dim() > 2:
|
| 112 |
+
lora_diff = torch.einsum(
|
| 113 |
+
"o i ..., i j -> o j ...",
|
| 114 |
+
torch.einsum(
|
| 115 |
+
"o i ..., i j -> o j ...",
|
| 116 |
+
weight.to(dtype=intermediate_dtype),
|
| 117 |
+
a1,
|
| 118 |
+
),
|
| 119 |
+
a2,
|
| 120 |
+
).reshape(weight.shape)
|
| 121 |
+
else:
|
| 122 |
+
lora_diff = torch.mm(
|
| 123 |
+
torch.mm(weight.to(dtype=intermediate_dtype), a1), a2
|
| 124 |
+
).reshape(weight.shape)
|
| 125 |
+
lora_diff += torch.mm(b1, b2).reshape(weight.shape)
|
| 126 |
+
|
| 127 |
+
if dora_scale is not None:
|
| 128 |
+
weight = weight_decompose(
|
| 129 |
+
dora_scale,
|
| 130 |
+
weight,
|
| 131 |
+
lora_diff,
|
| 132 |
+
alpha,
|
| 133 |
+
strength,
|
| 134 |
+
intermediate_dtype,
|
| 135 |
+
function,
|
| 136 |
+
)
|
| 137 |
+
else:
|
| 138 |
+
weight += function(((strength * alpha) * lora_diff).type(weight.dtype))
|
| 139 |
+
except Exception as e:
|
| 140 |
+
logging.error("ERROR {} {} {}".format(self.name, key, e))
|
| 141 |
+
return weight
|
| 142 |
+
|
| 143 |
+
def _compute_paths(self, x: torch.Tensor):
|
| 144 |
+
"""
|
| 145 |
+
Compute A path and B path outputs for GLoRA bypass.
|
| 146 |
+
|
| 147 |
+
GLoRA: f(x) = Wx + WAx + Bx
|
| 148 |
+
- A path: a1(a2(x)) - modifies input to base forward
|
| 149 |
+
- B path: b1(b2(x)) - additive component
|
| 150 |
+
|
| 151 |
+
Note:
|
| 152 |
+
Does not access original model weights - bypass mode is designed
|
| 153 |
+
for quantized models where weights may not be accessible.
|
| 154 |
+
|
| 155 |
+
Returns: (a_out, b_out)
|
| 156 |
+
"""
|
| 157 |
+
v = self.weights
|
| 158 |
+
# v = (a1, a2, b1, b2, alpha, dora_scale)
|
| 159 |
+
a1 = v[0]
|
| 160 |
+
a2 = v[1]
|
| 161 |
+
b1 = v[2]
|
| 162 |
+
b2 = v[3]
|
| 163 |
+
alpha = v[4]
|
| 164 |
+
|
| 165 |
+
dtype = x.dtype
|
| 166 |
+
|
| 167 |
+
# Cast dtype (weights should already be on correct device from inject())
|
| 168 |
+
a1 = a1.to(dtype=dtype)
|
| 169 |
+
a2 = a2.to(dtype=dtype)
|
| 170 |
+
b1 = b1.to(dtype=dtype)
|
| 171 |
+
b2 = b2.to(dtype=dtype)
|
| 172 |
+
|
| 173 |
+
# Determine rank and scale
|
| 174 |
+
# Check for old vs new glora format
|
| 175 |
+
old_glora = False
|
| 176 |
+
if b2.shape[1] == b1.shape[0] == a1.shape[0] == a2.shape[1]:
|
| 177 |
+
rank = a1.shape[0]
|
| 178 |
+
old_glora = True
|
| 179 |
+
|
| 180 |
+
if b2.shape[0] == b1.shape[1] == a1.shape[1] == a2.shape[0]:
|
| 181 |
+
if old_glora and a2.shape[0] == x.shape[-1] and x.shape[-1] == x.shape[-1]:
|
| 182 |
+
pass
|
| 183 |
+
else:
|
| 184 |
+
old_glora = False
|
| 185 |
+
rank = a2.shape[0]
|
| 186 |
+
|
| 187 |
+
if alpha is not None:
|
| 188 |
+
scale = alpha / rank
|
| 189 |
+
else:
|
| 190 |
+
scale = 1.0
|
| 191 |
+
|
| 192 |
+
# Apply multiplier
|
| 193 |
+
multiplier = getattr(self, "multiplier", 1.0)
|
| 194 |
+
scale = scale * multiplier
|
| 195 |
+
|
| 196 |
+
# Use module info from bypass injection, not input tensor shape
|
| 197 |
+
is_conv = getattr(self, "is_conv", False)
|
| 198 |
+
conv_dim = getattr(self, "conv_dim", 0)
|
| 199 |
+
kw_dict = getattr(self, "kw_dict", {})
|
| 200 |
+
|
| 201 |
+
if is_conv:
|
| 202 |
+
# Conv case - conv_dim is 1/2/3 for conv1d/2d/3d
|
| 203 |
+
conv_fn = (F.conv1d, F.conv2d, F.conv3d)[conv_dim - 1]
|
| 204 |
+
|
| 205 |
+
# Get module's stride/padding for spatial dimension handling
|
| 206 |
+
module_stride = kw_dict.get("stride", (1,) * conv_dim)
|
| 207 |
+
module_padding = kw_dict.get("padding", (0,) * conv_dim)
|
| 208 |
+
kernel_size = getattr(self, "kernel_size", (1,) * conv_dim)
|
| 209 |
+
in_channels = getattr(self, "in_channels", None)
|
| 210 |
+
|
| 211 |
+
# Ensure weights are in conv shape
|
| 212 |
+
# a1, a2, b1 are always 1x1 kernels
|
| 213 |
+
if a1.ndim == 2:
|
| 214 |
+
a1 = a1.view(*a1.shape, *([1] * conv_dim))
|
| 215 |
+
if a2.ndim == 2:
|
| 216 |
+
a2 = a2.view(*a2.shape, *([1] * conv_dim))
|
| 217 |
+
if b1.ndim == 2:
|
| 218 |
+
b1 = b1.view(*b1.shape, *([1] * conv_dim))
|
| 219 |
+
# b2 has actual kernel_size (like LoRA down)
|
| 220 |
+
if b2.ndim == 2:
|
| 221 |
+
if in_channels is not None:
|
| 222 |
+
b2 = b2.view(b2.shape[0], in_channels, *kernel_size)
|
| 223 |
+
else:
|
| 224 |
+
b2 = b2.view(*b2.shape, *([1] * conv_dim))
|
| 225 |
+
|
| 226 |
+
# A path: a2(x) -> a1(...) - 1x1 convs, no stride/padding needed, a_out is added to x
|
| 227 |
+
a2_out = conv_fn(x, a2)
|
| 228 |
+
a_out = conv_fn(a2_out, a1) * scale
|
| 229 |
+
|
| 230 |
+
# B path: b2(x) with kernel/stride/padding -> b1(...) 1x1
|
| 231 |
+
b2_out = conv_fn(x, b2, stride=module_stride, padding=module_padding)
|
| 232 |
+
b_out = conv_fn(b2_out, b1) * scale
|
| 233 |
+
else:
|
| 234 |
+
# Linear case
|
| 235 |
+
if old_glora:
|
| 236 |
+
# Old format: a1 @ a2 @ x, b2 @ b1
|
| 237 |
+
a_out = F.linear(F.linear(x, a2), a1) * scale
|
| 238 |
+
b_out = F.linear(F.linear(x, b1), b2) * scale
|
| 239 |
+
else:
|
| 240 |
+
# New format: x @ a1 @ a2, b1 @ b2
|
| 241 |
+
a_out = F.linear(F.linear(x, a1), a2) * scale
|
| 242 |
+
b_out = F.linear(F.linear(x, b2), b1) * scale
|
| 243 |
+
|
| 244 |
+
return a_out, b_out
|
| 245 |
+
|
| 246 |
+
def bypass_forward(
|
| 247 |
+
self,
|
| 248 |
+
org_forward: Callable,
|
| 249 |
+
x: torch.Tensor,
|
| 250 |
+
*args,
|
| 251 |
+
**kwargs,
|
| 252 |
+
) -> torch.Tensor:
|
| 253 |
+
"""
|
| 254 |
+
GLoRA bypass forward: f(x + a(x)) + b(x)
|
| 255 |
+
|
| 256 |
+
Unlike standard adapters, GLoRA modifies the input to the base forward
|
| 257 |
+
AND adds the B path output.
|
| 258 |
+
|
| 259 |
+
Note:
|
| 260 |
+
Does not access original model weights - bypass mode is designed
|
| 261 |
+
for quantized models where weights may not be accessible.
|
| 262 |
+
|
| 263 |
+
Reference: LyCORIS GLoRAModule._bypass_forward
|
| 264 |
+
"""
|
| 265 |
+
a_out, b_out = self._compute_paths(x)
|
| 266 |
+
|
| 267 |
+
# Call base forward with modified input
|
| 268 |
+
base_out = org_forward(x + a_out, *args, **kwargs)
|
| 269 |
+
|
| 270 |
+
# Add B path
|
| 271 |
+
return base_out + b_out
|
| 272 |
+
|
| 273 |
+
def h(self, x: torch.Tensor, base_out: torch.Tensor) -> torch.Tensor:
|
| 274 |
+
"""
|
| 275 |
+
For GLoRA, h() returns the B path output.
|
| 276 |
+
|
| 277 |
+
Note:
|
| 278 |
+
GLoRA's full bypass requires overriding bypass_forward() since
|
| 279 |
+
it also modifies the input to org_forward. This h() is provided for
|
| 280 |
+
compatibility but bypass_forward() should be used for correct behavior.
|
| 281 |
+
|
| 282 |
+
Does not access original model weights - bypass mode is designed
|
| 283 |
+
for quantized models where weights may not be accessible.
|
| 284 |
+
|
| 285 |
+
Args:
|
| 286 |
+
x: Input tensor
|
| 287 |
+
base_out: Output from base forward (unused, for API consistency)
|
| 288 |
+
"""
|
| 289 |
+
_, b_out = self._compute_paths(x)
|
| 290 |
+
return b_out
|
comfy/weight_adapter/loha.py
ADDED
|
@@ -0,0 +1,378 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from functools import cache
|
| 3 |
+
from typing import Optional
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
import torch.nn.functional as F
|
| 7 |
+
import comfy.model_management
|
| 8 |
+
from .base import WeightAdapterBase, WeightAdapterTrainBase, weight_decompose
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
@cache
|
| 12 |
+
def _warn_loha_bypass_inefficient():
|
| 13 |
+
"""One-time warning about LoHa bypass inefficiency."""
|
| 14 |
+
logging.warning(
|
| 15 |
+
"LoHa bypass mode is inefficient: full weight diff is computed each forward pass. "
|
| 16 |
+
"Consider using LoRA or LoKr for training with bypass mode."
|
| 17 |
+
)
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class HadaWeight(torch.autograd.Function):
|
| 21 |
+
@staticmethod
|
| 22 |
+
def forward(ctx, w1u, w1d, w2u, w2d, scale=torch.tensor(1)):
|
| 23 |
+
ctx.save_for_backward(w1d, w1u, w2d, w2u, scale)
|
| 24 |
+
diff_weight = ((w1u @ w1d) * (w2u @ w2d)) * scale
|
| 25 |
+
return diff_weight
|
| 26 |
+
|
| 27 |
+
@staticmethod
|
| 28 |
+
def backward(ctx, grad_out):
|
| 29 |
+
(w1d, w1u, w2d, w2u, scale) = ctx.saved_tensors
|
| 30 |
+
grad_out = grad_out * scale
|
| 31 |
+
temp = grad_out * (w2u @ w2d)
|
| 32 |
+
grad_w1u = temp @ w1d.T
|
| 33 |
+
grad_w1d = w1u.T @ temp
|
| 34 |
+
|
| 35 |
+
temp = grad_out * (w1u @ w1d)
|
| 36 |
+
grad_w2u = temp @ w2d.T
|
| 37 |
+
grad_w2d = w2u.T @ temp
|
| 38 |
+
|
| 39 |
+
del temp
|
| 40 |
+
return grad_w1u, grad_w1d, grad_w2u, grad_w2d, None
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
class HadaWeightTucker(torch.autograd.Function):
|
| 44 |
+
@staticmethod
|
| 45 |
+
def forward(ctx, t1, w1u, w1d, t2, w2u, w2d, scale=torch.tensor(1)):
|
| 46 |
+
ctx.save_for_backward(t1, w1d, w1u, t2, w2d, w2u, scale)
|
| 47 |
+
|
| 48 |
+
rebuild1 = torch.einsum("i j ..., j r, i p -> p r ...", t1, w1d, w1u)
|
| 49 |
+
rebuild2 = torch.einsum("i j ..., j r, i p -> p r ...", t2, w2d, w2u)
|
| 50 |
+
|
| 51 |
+
return rebuild1 * rebuild2 * scale
|
| 52 |
+
|
| 53 |
+
@staticmethod
|
| 54 |
+
def backward(ctx, grad_out):
|
| 55 |
+
(t1, w1d, w1u, t2, w2d, w2u, scale) = ctx.saved_tensors
|
| 56 |
+
grad_out = grad_out * scale
|
| 57 |
+
|
| 58 |
+
temp = torch.einsum("i j ..., j r -> i r ...", t2, w2d)
|
| 59 |
+
rebuild = torch.einsum("i j ..., i r -> r j ...", temp, w2u)
|
| 60 |
+
|
| 61 |
+
grad_w = rebuild * grad_out
|
| 62 |
+
del rebuild
|
| 63 |
+
|
| 64 |
+
grad_w1u = torch.einsum("r j ..., i j ... -> r i", temp, grad_w)
|
| 65 |
+
grad_temp = torch.einsum("i j ..., i r -> r j ...", grad_w, w1u.T)
|
| 66 |
+
del grad_w, temp
|
| 67 |
+
|
| 68 |
+
grad_w1d = torch.einsum("i r ..., i j ... -> r j", t1, grad_temp)
|
| 69 |
+
grad_t1 = torch.einsum("i j ..., j r -> i r ...", grad_temp, w1d.T)
|
| 70 |
+
del grad_temp
|
| 71 |
+
|
| 72 |
+
temp = torch.einsum("i j ..., j r -> i r ...", t1, w1d)
|
| 73 |
+
rebuild = torch.einsum("i j ..., i r -> r j ...", temp, w1u)
|
| 74 |
+
|
| 75 |
+
grad_w = rebuild * grad_out
|
| 76 |
+
del rebuild
|
| 77 |
+
|
| 78 |
+
grad_w2u = torch.einsum("r j ..., i j ... -> r i", temp, grad_w)
|
| 79 |
+
grad_temp = torch.einsum("i j ..., i r -> r j ...", grad_w, w2u.T)
|
| 80 |
+
del grad_w, temp
|
| 81 |
+
|
| 82 |
+
grad_w2d = torch.einsum("i r ..., i j ... -> r j", t2, grad_temp)
|
| 83 |
+
grad_t2 = torch.einsum("i j ..., j r -> i r ...", grad_temp, w2d.T)
|
| 84 |
+
del grad_temp
|
| 85 |
+
return grad_t1, grad_w1u, grad_w1d, grad_t2, grad_w2u, grad_w2d, None
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
class LohaDiff(WeightAdapterTrainBase):
|
| 89 |
+
def __init__(self, weights):
|
| 90 |
+
super().__init__()
|
| 91 |
+
# Unpack weights tuple from LoHaAdapter
|
| 92 |
+
w1a, w1b, alpha, w2a, w2b, t1, t2, _ = weights
|
| 93 |
+
|
| 94 |
+
# Create trainable parameters
|
| 95 |
+
self.hada_w1_a = torch.nn.Parameter(w1a)
|
| 96 |
+
self.hada_w1_b = torch.nn.Parameter(w1b)
|
| 97 |
+
self.hada_w2_a = torch.nn.Parameter(w2a)
|
| 98 |
+
self.hada_w2_b = torch.nn.Parameter(w2b)
|
| 99 |
+
|
| 100 |
+
self.use_tucker = False
|
| 101 |
+
if t1 is not None and t2 is not None:
|
| 102 |
+
self.use_tucker = True
|
| 103 |
+
self.hada_t1 = torch.nn.Parameter(t1)
|
| 104 |
+
self.hada_t2 = torch.nn.Parameter(t2)
|
| 105 |
+
else:
|
| 106 |
+
# Keep the attributes for consistent access
|
| 107 |
+
self.hada_t1 = None
|
| 108 |
+
self.hada_t2 = None
|
| 109 |
+
|
| 110 |
+
# Store rank and non-trainable alpha
|
| 111 |
+
self.rank = w1b.shape[0]
|
| 112 |
+
self.alpha = torch.nn.Parameter(torch.tensor(alpha), requires_grad=False)
|
| 113 |
+
|
| 114 |
+
def __call__(self, w):
|
| 115 |
+
org_dtype = w.dtype
|
| 116 |
+
|
| 117 |
+
scale = self.alpha / self.rank
|
| 118 |
+
if self.use_tucker:
|
| 119 |
+
diff_weight = HadaWeightTucker.apply(
|
| 120 |
+
self.hada_t1,
|
| 121 |
+
self.hada_w1_a,
|
| 122 |
+
self.hada_w1_b,
|
| 123 |
+
self.hada_t2,
|
| 124 |
+
self.hada_w2_a,
|
| 125 |
+
self.hada_w2_b,
|
| 126 |
+
scale,
|
| 127 |
+
)
|
| 128 |
+
else:
|
| 129 |
+
diff_weight = HadaWeight.apply(
|
| 130 |
+
self.hada_w1_a, self.hada_w1_b, self.hada_w2_a, self.hada_w2_b, scale
|
| 131 |
+
)
|
| 132 |
+
|
| 133 |
+
# Add the scaled difference to the original weight
|
| 134 |
+
weight = w.to(diff_weight) + diff_weight.reshape(w.shape)
|
| 135 |
+
|
| 136 |
+
return weight.to(org_dtype)
|
| 137 |
+
|
| 138 |
+
def passive_memory_usage(self):
|
| 139 |
+
"""Calculates memory usage of the trainable parameters."""
|
| 140 |
+
return sum(param.numel() * param.element_size() for param in self.parameters())
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
class LoHaAdapter(WeightAdapterBase):
|
| 144 |
+
name = "loha"
|
| 145 |
+
|
| 146 |
+
def __init__(self, loaded_keys, weights):
|
| 147 |
+
self.loaded_keys = loaded_keys
|
| 148 |
+
self.weights = weights
|
| 149 |
+
|
| 150 |
+
@classmethod
|
| 151 |
+
def create_train(cls, weight, rank=1, alpha=1.0):
|
| 152 |
+
out_dim = weight.shape[0]
|
| 153 |
+
in_dim = weight.shape[1:].numel()
|
| 154 |
+
mat1 = torch.empty(out_dim, rank, device=weight.device, dtype=torch.float32)
|
| 155 |
+
mat2 = torch.empty(rank, in_dim, device=weight.device, dtype=torch.float32)
|
| 156 |
+
torch.nn.init.normal_(mat1, 0.1)
|
| 157 |
+
torch.nn.init.constant_(mat2, 0.0)
|
| 158 |
+
mat3 = torch.empty(out_dim, rank, device=weight.device, dtype=torch.float32)
|
| 159 |
+
mat4 = torch.empty(rank, in_dim, device=weight.device, dtype=torch.float32)
|
| 160 |
+
torch.nn.init.normal_(mat3, 0.1)
|
| 161 |
+
torch.nn.init.normal_(mat4, 0.01)
|
| 162 |
+
return LohaDiff((mat1, mat2, alpha, mat3, mat4, None, None, None))
|
| 163 |
+
|
| 164 |
+
def to_train(self):
|
| 165 |
+
return LohaDiff(self.weights)
|
| 166 |
+
|
| 167 |
+
@classmethod
|
| 168 |
+
def load(
|
| 169 |
+
cls,
|
| 170 |
+
x: str,
|
| 171 |
+
lora: dict[str, torch.Tensor],
|
| 172 |
+
alpha: float,
|
| 173 |
+
dora_scale: torch.Tensor,
|
| 174 |
+
loaded_keys: set[str] = None,
|
| 175 |
+
) -> Optional["LoHaAdapter"]:
|
| 176 |
+
if loaded_keys is None:
|
| 177 |
+
loaded_keys = set()
|
| 178 |
+
|
| 179 |
+
hada_w1_a_name = "{}.hada_w1_a".format(x)
|
| 180 |
+
hada_w1_b_name = "{}.hada_w1_b".format(x)
|
| 181 |
+
hada_w2_a_name = "{}.hada_w2_a".format(x)
|
| 182 |
+
hada_w2_b_name = "{}.hada_w2_b".format(x)
|
| 183 |
+
hada_t1_name = "{}.hada_t1".format(x)
|
| 184 |
+
hada_t2_name = "{}.hada_t2".format(x)
|
| 185 |
+
if hada_w1_a_name in lora.keys():
|
| 186 |
+
hada_t1 = None
|
| 187 |
+
hada_t2 = None
|
| 188 |
+
if hada_t1_name in lora.keys():
|
| 189 |
+
hada_t1 = lora[hada_t1_name]
|
| 190 |
+
hada_t2 = lora[hada_t2_name]
|
| 191 |
+
loaded_keys.add(hada_t1_name)
|
| 192 |
+
loaded_keys.add(hada_t2_name)
|
| 193 |
+
|
| 194 |
+
weights = (
|
| 195 |
+
lora[hada_w1_a_name],
|
| 196 |
+
lora[hada_w1_b_name],
|
| 197 |
+
alpha,
|
| 198 |
+
lora[hada_w2_a_name],
|
| 199 |
+
lora[hada_w2_b_name],
|
| 200 |
+
hada_t1,
|
| 201 |
+
hada_t2,
|
| 202 |
+
dora_scale,
|
| 203 |
+
)
|
| 204 |
+
loaded_keys.add(hada_w1_a_name)
|
| 205 |
+
loaded_keys.add(hada_w1_b_name)
|
| 206 |
+
loaded_keys.add(hada_w2_a_name)
|
| 207 |
+
loaded_keys.add(hada_w2_b_name)
|
| 208 |
+
return cls(loaded_keys, weights)
|
| 209 |
+
else:
|
| 210 |
+
return None
|
| 211 |
+
|
| 212 |
+
def calculate_weight(
|
| 213 |
+
self,
|
| 214 |
+
weight,
|
| 215 |
+
key,
|
| 216 |
+
strength,
|
| 217 |
+
strength_model,
|
| 218 |
+
offset,
|
| 219 |
+
function,
|
| 220 |
+
intermediate_dtype=torch.float32,
|
| 221 |
+
original_weight=None,
|
| 222 |
+
):
|
| 223 |
+
v = self.weights
|
| 224 |
+
w1a = v[0]
|
| 225 |
+
w1b = v[1]
|
| 226 |
+
if v[2] is not None:
|
| 227 |
+
alpha = v[2] / w1b.shape[0]
|
| 228 |
+
else:
|
| 229 |
+
alpha = 1.0
|
| 230 |
+
|
| 231 |
+
w2a = v[3]
|
| 232 |
+
w2b = v[4]
|
| 233 |
+
dora_scale = v[7]
|
| 234 |
+
if v[5] is not None: # cp decomposition
|
| 235 |
+
t1 = v[5]
|
| 236 |
+
t2 = v[6]
|
| 237 |
+
m1 = torch.einsum(
|
| 238 |
+
"i j k l, j r, i p -> p r k l",
|
| 239 |
+
comfy.model_management.cast_to_device(
|
| 240 |
+
t1, weight.device, intermediate_dtype
|
| 241 |
+
),
|
| 242 |
+
comfy.model_management.cast_to_device(
|
| 243 |
+
w1b, weight.device, intermediate_dtype
|
| 244 |
+
),
|
| 245 |
+
comfy.model_management.cast_to_device(
|
| 246 |
+
w1a, weight.device, intermediate_dtype
|
| 247 |
+
),
|
| 248 |
+
)
|
| 249 |
+
|
| 250 |
+
m2 = torch.einsum(
|
| 251 |
+
"i j k l, j r, i p -> p r k l",
|
| 252 |
+
comfy.model_management.cast_to_device(
|
| 253 |
+
t2, weight.device, intermediate_dtype
|
| 254 |
+
),
|
| 255 |
+
comfy.model_management.cast_to_device(
|
| 256 |
+
w2b, weight.device, intermediate_dtype
|
| 257 |
+
),
|
| 258 |
+
comfy.model_management.cast_to_device(
|
| 259 |
+
w2a, weight.device, intermediate_dtype
|
| 260 |
+
),
|
| 261 |
+
)
|
| 262 |
+
else:
|
| 263 |
+
m1 = torch.mm(
|
| 264 |
+
comfy.model_management.cast_to_device(
|
| 265 |
+
w1a, weight.device, intermediate_dtype
|
| 266 |
+
),
|
| 267 |
+
comfy.model_management.cast_to_device(
|
| 268 |
+
w1b, weight.device, intermediate_dtype
|
| 269 |
+
),
|
| 270 |
+
)
|
| 271 |
+
m2 = torch.mm(
|
| 272 |
+
comfy.model_management.cast_to_device(
|
| 273 |
+
w2a, weight.device, intermediate_dtype
|
| 274 |
+
),
|
| 275 |
+
comfy.model_management.cast_to_device(
|
| 276 |
+
w2b, weight.device, intermediate_dtype
|
| 277 |
+
),
|
| 278 |
+
)
|
| 279 |
+
|
| 280 |
+
try:
|
| 281 |
+
lora_diff = (m1 * m2).reshape(weight.shape)
|
| 282 |
+
if dora_scale is not None:
|
| 283 |
+
weight = weight_decompose(
|
| 284 |
+
dora_scale,
|
| 285 |
+
weight,
|
| 286 |
+
lora_diff,
|
| 287 |
+
alpha,
|
| 288 |
+
strength,
|
| 289 |
+
intermediate_dtype,
|
| 290 |
+
function,
|
| 291 |
+
)
|
| 292 |
+
else:
|
| 293 |
+
weight += function(((strength * alpha) * lora_diff).type(weight.dtype))
|
| 294 |
+
except Exception as e:
|
| 295 |
+
logging.error("ERROR {} {} {}".format(self.name, key, e))
|
| 296 |
+
return weight
|
| 297 |
+
|
| 298 |
+
def h(self, x: torch.Tensor, base_out: torch.Tensor) -> torch.Tensor:
|
| 299 |
+
"""
|
| 300 |
+
Additive bypass component for LoHa: h(x) = diff_weight @ x
|
| 301 |
+
|
| 302 |
+
WARNING: Inefficient - computes full Hadamard product each forward.
|
| 303 |
+
|
| 304 |
+
Note:
|
| 305 |
+
Does not access original model weights - bypass mode is designed
|
| 306 |
+
for quantized models where weights may not be accessible.
|
| 307 |
+
|
| 308 |
+
Args:
|
| 309 |
+
x: Input tensor
|
| 310 |
+
base_out: Output from base forward (unused, for API consistency)
|
| 311 |
+
|
| 312 |
+
Reference: LyCORIS functional/loha.py bypass_forward_diff
|
| 313 |
+
"""
|
| 314 |
+
_warn_loha_bypass_inefficient()
|
| 315 |
+
|
| 316 |
+
# FUNC_LIST: [None, None, F.linear, F.conv1d, F.conv2d, F.conv3d]
|
| 317 |
+
FUNC_LIST = [None, None, F.linear, F.conv1d, F.conv2d, F.conv3d]
|
| 318 |
+
|
| 319 |
+
v = self.weights
|
| 320 |
+
# v[0]=w1a, v[1]=w1b, v[2]=alpha, v[3]=w2a, v[4]=w2b, v[5]=t1, v[6]=t2, v[7]=dora
|
| 321 |
+
w1a = v[0]
|
| 322 |
+
w1b = v[1]
|
| 323 |
+
alpha = v[2]
|
| 324 |
+
w2a = v[3]
|
| 325 |
+
w2b = v[4]
|
| 326 |
+
t1 = v[5]
|
| 327 |
+
t2 = v[6]
|
| 328 |
+
|
| 329 |
+
# Compute scale
|
| 330 |
+
rank = w1b.shape[0]
|
| 331 |
+
scale = (alpha / rank if alpha is not None else 1.0) * getattr(
|
| 332 |
+
self, "multiplier", 1.0
|
| 333 |
+
)
|
| 334 |
+
|
| 335 |
+
# Cast dtype
|
| 336 |
+
w1a = w1a.to(dtype=x.dtype)
|
| 337 |
+
w1b = w1b.to(dtype=x.dtype)
|
| 338 |
+
w2a = w2a.to(dtype=x.dtype)
|
| 339 |
+
w2b = w2b.to(dtype=x.dtype)
|
| 340 |
+
|
| 341 |
+
# Use module info from bypass injection, not weight dimension
|
| 342 |
+
is_conv = getattr(self, "is_conv", False)
|
| 343 |
+
conv_dim = getattr(self, "conv_dim", 0)
|
| 344 |
+
kw_dict = getattr(self, "kw_dict", {})
|
| 345 |
+
|
| 346 |
+
# Compute diff weight using Hadamard product
|
| 347 |
+
if t1 is not None and t2 is not None:
|
| 348 |
+
t1 = t1.to(dtype=x.dtype)
|
| 349 |
+
t2 = t2.to(dtype=x.dtype)
|
| 350 |
+
m1 = torch.einsum("i j k l, j r, i p -> p r k l", t1, w1b, w1a)
|
| 351 |
+
m2 = torch.einsum("i j k l, j r, i p -> p r k l", t2, w2b, w2a)
|
| 352 |
+
diff_weight = (m1 * m2) * scale
|
| 353 |
+
else:
|
| 354 |
+
m1 = w1a @ w1b
|
| 355 |
+
m2 = w2a @ w2b
|
| 356 |
+
diff_weight = (m1 * m2) * scale
|
| 357 |
+
|
| 358 |
+
if is_conv:
|
| 359 |
+
op = FUNC_LIST[conv_dim + 2]
|
| 360 |
+
kernel_size = getattr(self, "kernel_size", (1,) * conv_dim)
|
| 361 |
+
in_channels = getattr(self, "in_channels", None)
|
| 362 |
+
|
| 363 |
+
# Reshape 2D diff_weight to conv format using kernel_size
|
| 364 |
+
# diff_weight: [out_channels, in_channels * prod(kernel_size)] -> [out_channels, in_channels, *kernel_size]
|
| 365 |
+
if diff_weight.dim() == 2:
|
| 366 |
+
if in_channels is not None:
|
| 367 |
+
diff_weight = diff_weight.view(
|
| 368 |
+
diff_weight.shape[0], in_channels, *kernel_size
|
| 369 |
+
)
|
| 370 |
+
else:
|
| 371 |
+
diff_weight = diff_weight.view(
|
| 372 |
+
*diff_weight.shape, *([1] * conv_dim)
|
| 373 |
+
)
|
| 374 |
+
else:
|
| 375 |
+
op = F.linear
|
| 376 |
+
kw_dict = {}
|
| 377 |
+
|
| 378 |
+
return op(x, diff_weight, **kw_dict)
|
comfy/weight_adapter/lokr.py
ADDED
|
@@ -0,0 +1,481 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from typing import Optional
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import torch.nn.functional as F
|
| 6 |
+
import comfy.model_management
|
| 7 |
+
from .base import (
|
| 8 |
+
WeightAdapterBase,
|
| 9 |
+
WeightAdapterTrainBase,
|
| 10 |
+
weight_decompose,
|
| 11 |
+
factorization,
|
| 12 |
+
)
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class LokrDiff(WeightAdapterTrainBase):
|
| 16 |
+
def __init__(self, weights):
|
| 17 |
+
super().__init__()
|
| 18 |
+
(
|
| 19 |
+
lokr_w1,
|
| 20 |
+
lokr_w2,
|
| 21 |
+
alpha,
|
| 22 |
+
lokr_w1_a,
|
| 23 |
+
lokr_w1_b,
|
| 24 |
+
lokr_w2_a,
|
| 25 |
+
lokr_w2_b,
|
| 26 |
+
lokr_t2,
|
| 27 |
+
dora_scale,
|
| 28 |
+
) = weights
|
| 29 |
+
self.use_tucker = False
|
| 30 |
+
if lokr_w1_a is not None:
|
| 31 |
+
_, rank_a = lokr_w1_a.shape[0], lokr_w1_a.shape[1]
|
| 32 |
+
rank_a, _ = lokr_w1_b.shape[0], lokr_w1_b.shape[1]
|
| 33 |
+
self.lokr_w1_a = torch.nn.Parameter(lokr_w1_a)
|
| 34 |
+
self.lokr_w1_b = torch.nn.Parameter(lokr_w1_b)
|
| 35 |
+
self.w1_rebuild = True
|
| 36 |
+
self.ranka = rank_a
|
| 37 |
+
|
| 38 |
+
if lokr_w2_a is not None:
|
| 39 |
+
_, rank_b = lokr_w2_a.shape[0], lokr_w2_a.shape[1]
|
| 40 |
+
rank_b, _ = lokr_w2_b.shape[0], lokr_w2_b.shape[1]
|
| 41 |
+
self.lokr_w2_a = torch.nn.Parameter(lokr_w2_a)
|
| 42 |
+
self.lokr_w2_b = torch.nn.Parameter(lokr_w2_b)
|
| 43 |
+
if lokr_t2 is not None:
|
| 44 |
+
self.use_tucker = True
|
| 45 |
+
self.lokr_t2 = torch.nn.Parameter(lokr_t2)
|
| 46 |
+
self.w2_rebuild = True
|
| 47 |
+
self.rankb = rank_b
|
| 48 |
+
|
| 49 |
+
if lokr_w1 is not None:
|
| 50 |
+
self.lokr_w1 = torch.nn.Parameter(lokr_w1)
|
| 51 |
+
self.w1_rebuild = False
|
| 52 |
+
|
| 53 |
+
if lokr_w2 is not None:
|
| 54 |
+
self.lokr_w2 = torch.nn.Parameter(lokr_w2)
|
| 55 |
+
self.w2_rebuild = False
|
| 56 |
+
|
| 57 |
+
self.alpha = torch.nn.Parameter(torch.tensor(alpha), requires_grad=False)
|
| 58 |
+
|
| 59 |
+
@property
|
| 60 |
+
def w1(self):
|
| 61 |
+
if self.w1_rebuild:
|
| 62 |
+
return (self.lokr_w1_a @ self.lokr_w1_b) * (self.alpha / self.ranka)
|
| 63 |
+
else:
|
| 64 |
+
return self.lokr_w1
|
| 65 |
+
|
| 66 |
+
@property
|
| 67 |
+
def w2(self):
|
| 68 |
+
if self.w2_rebuild:
|
| 69 |
+
if self.use_tucker:
|
| 70 |
+
w2 = torch.einsum(
|
| 71 |
+
"i j k l, j r, i p -> p r k l",
|
| 72 |
+
self.lokr_t2,
|
| 73 |
+
self.lokr_w2_b,
|
| 74 |
+
self.lokr_w2_a,
|
| 75 |
+
)
|
| 76 |
+
else:
|
| 77 |
+
w2 = self.lokr_w2_a @ self.lokr_w2_b
|
| 78 |
+
return w2 * (self.alpha / self.rankb)
|
| 79 |
+
else:
|
| 80 |
+
return self.lokr_w2
|
| 81 |
+
|
| 82 |
+
def __call__(self, w):
|
| 83 |
+
w1 = self.w1
|
| 84 |
+
w2 = self.w2
|
| 85 |
+
# Unsqueeze w1 to match w2 dims for proper kron product (like LyCORIS make_kron)
|
| 86 |
+
for _ in range(w2.dim() - w1.dim()):
|
| 87 |
+
w1 = w1.unsqueeze(-1)
|
| 88 |
+
diff = torch.kron(w1, w2)
|
| 89 |
+
return w + diff.reshape(w.shape).to(w)
|
| 90 |
+
|
| 91 |
+
def h(self, x: torch.Tensor, base_out: torch.Tensor) -> torch.Tensor:
|
| 92 |
+
"""
|
| 93 |
+
Additive bypass component for LoKr training: efficient Kronecker product.
|
| 94 |
+
|
| 95 |
+
Uses w1/w2 properties which handle both direct and decomposed cases.
|
| 96 |
+
For create_train (direct w1/w2), no alpha scaling in properties.
|
| 97 |
+
For to_train (decomposed), alpha/rank scaling is in properties.
|
| 98 |
+
|
| 99 |
+
Args:
|
| 100 |
+
x: Input tensor
|
| 101 |
+
base_out: Output from base forward (unused, for API consistency)
|
| 102 |
+
"""
|
| 103 |
+
# Get w1, w2 from properties (handles rebuild vs direct)
|
| 104 |
+
w1 = self.w1
|
| 105 |
+
w2 = self.w2
|
| 106 |
+
|
| 107 |
+
# Multiplier from bypass injection
|
| 108 |
+
multiplier = getattr(self, "multiplier", 1.0)
|
| 109 |
+
|
| 110 |
+
# Get module info from bypass injection
|
| 111 |
+
is_conv = getattr(self, "is_conv", False)
|
| 112 |
+
conv_dim = getattr(self, "conv_dim", 0)
|
| 113 |
+
kw_dict = getattr(self, "kw_dict", {})
|
| 114 |
+
|
| 115 |
+
# Efficient Kronecker application without materializing full weight
|
| 116 |
+
# kron(w1, w2) @ x can be computed as nested operations
|
| 117 |
+
# w1: [out_l, in_m], w2: [out_k, in_n, *k_size]
|
| 118 |
+
# Full weight would be [out_l*out_k, in_m*in_n, *k_size]
|
| 119 |
+
|
| 120 |
+
uq = w1.size(1) # in_m - inner grouping dimension
|
| 121 |
+
|
| 122 |
+
if is_conv:
|
| 123 |
+
conv_fn = (F.conv1d, F.conv2d, F.conv3d)[conv_dim - 1]
|
| 124 |
+
|
| 125 |
+
B, C_in, *spatial = x.shape
|
| 126 |
+
# Reshape input for grouped application: [B * uq, C_in // uq, *spatial]
|
| 127 |
+
h_in_group = x.reshape(B * uq, -1, *spatial)
|
| 128 |
+
|
| 129 |
+
# Ensure w2 has conv dims
|
| 130 |
+
if w2.dim() == 2:
|
| 131 |
+
w2 = w2.view(*w2.shape, *([1] * conv_dim))
|
| 132 |
+
|
| 133 |
+
# Apply w2 path with stride/padding
|
| 134 |
+
hb = conv_fn(h_in_group, w2, **kw_dict)
|
| 135 |
+
|
| 136 |
+
# Reshape for cross-group operation
|
| 137 |
+
hb = hb.view(B, -1, *hb.shape[1:])
|
| 138 |
+
h_cross = hb.transpose(1, -1)
|
| 139 |
+
|
| 140 |
+
# Apply w1 (always 2D, applied as linear on channel dim)
|
| 141 |
+
hc = F.linear(h_cross, w1)
|
| 142 |
+
hc = hc.transpose(1, -1)
|
| 143 |
+
|
| 144 |
+
# Reshape to output
|
| 145 |
+
out = hc.reshape(B, -1, *hc.shape[3:])
|
| 146 |
+
else:
|
| 147 |
+
# Linear case
|
| 148 |
+
# Reshape input: [..., in_m * in_n] -> [..., uq (in_m), in_n]
|
| 149 |
+
h_in_group = x.reshape(*x.shape[:-1], uq, -1)
|
| 150 |
+
|
| 151 |
+
# Apply w2: [..., uq, in_n] @ [out_k, in_n].T -> [..., uq, out_k]
|
| 152 |
+
hb = F.linear(h_in_group, w2)
|
| 153 |
+
|
| 154 |
+
# Transpose for w1: [..., uq, out_k] -> [..., out_k, uq]
|
| 155 |
+
h_cross = hb.transpose(-1, -2)
|
| 156 |
+
|
| 157 |
+
# Apply w1: [..., out_k, uq] @ [out_l, uq].T -> [..., out_k, out_l]
|
| 158 |
+
hc = F.linear(h_cross, w1)
|
| 159 |
+
|
| 160 |
+
# Transpose back and flatten: [..., out_k, out_l] -> [..., out_l * out_k]
|
| 161 |
+
hc = hc.transpose(-1, -2)
|
| 162 |
+
out = hc.reshape(*hc.shape[:-2], -1)
|
| 163 |
+
|
| 164 |
+
return out * multiplier
|
| 165 |
+
|
| 166 |
+
def passive_memory_usage(self):
|
| 167 |
+
return sum(param.numel() * param.element_size() for param in self.parameters())
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
class LoKrAdapter(WeightAdapterBase):
|
| 171 |
+
name = "lokr"
|
| 172 |
+
|
| 173 |
+
def __init__(self, loaded_keys, weights):
|
| 174 |
+
self.loaded_keys = loaded_keys
|
| 175 |
+
self.weights = weights
|
| 176 |
+
|
| 177 |
+
@classmethod
|
| 178 |
+
def create_train(cls, weight, rank=1, alpha=1.0):
|
| 179 |
+
out_dim = weight.shape[0]
|
| 180 |
+
in_dim = weight.shape[1] # Just in_channels, not flattened with kernel
|
| 181 |
+
k_size = weight.shape[2:] if weight.dim() > 2 else ()
|
| 182 |
+
|
| 183 |
+
out_l, out_k = factorization(out_dim, rank)
|
| 184 |
+
in_m, in_n = factorization(in_dim, rank)
|
| 185 |
+
|
| 186 |
+
# w1: [out_l, in_m]
|
| 187 |
+
mat1 = torch.empty(out_l, in_m, device=weight.device, dtype=torch.float32)
|
| 188 |
+
# w2: [out_k, in_n, *k_size] for conv, [out_k, in_n] for linear
|
| 189 |
+
mat2 = torch.empty(
|
| 190 |
+
out_k, in_n, *k_size, device=weight.device, dtype=torch.float32
|
| 191 |
+
)
|
| 192 |
+
|
| 193 |
+
torch.nn.init.kaiming_uniform_(mat2, a=5**0.5)
|
| 194 |
+
torch.nn.init.constant_(mat1, 0.0)
|
| 195 |
+
return LokrDiff((mat1, mat2, alpha, None, None, None, None, None, None))
|
| 196 |
+
|
| 197 |
+
def to_train(self):
|
| 198 |
+
return LokrDiff(self.weights)
|
| 199 |
+
|
| 200 |
+
@classmethod
|
| 201 |
+
def load(
|
| 202 |
+
cls,
|
| 203 |
+
x: str,
|
| 204 |
+
lora: dict[str, torch.Tensor],
|
| 205 |
+
alpha: float,
|
| 206 |
+
dora_scale: torch.Tensor,
|
| 207 |
+
loaded_keys: set[str] = None,
|
| 208 |
+
) -> Optional["LoKrAdapter"]:
|
| 209 |
+
if loaded_keys is None:
|
| 210 |
+
loaded_keys = set()
|
| 211 |
+
lokr_w1_name = "{}.lokr_w1".format(x)
|
| 212 |
+
lokr_w2_name = "{}.lokr_w2".format(x)
|
| 213 |
+
lokr_w1_a_name = "{}.lokr_w1_a".format(x)
|
| 214 |
+
lokr_w1_b_name = "{}.lokr_w1_b".format(x)
|
| 215 |
+
lokr_t2_name = "{}.lokr_t2".format(x)
|
| 216 |
+
lokr_w2_a_name = "{}.lokr_w2_a".format(x)
|
| 217 |
+
lokr_w2_b_name = "{}.lokr_w2_b".format(x)
|
| 218 |
+
|
| 219 |
+
lokr_w1 = None
|
| 220 |
+
if lokr_w1_name in lora.keys():
|
| 221 |
+
lokr_w1 = lora[lokr_w1_name]
|
| 222 |
+
loaded_keys.add(lokr_w1_name)
|
| 223 |
+
|
| 224 |
+
lokr_w2 = None
|
| 225 |
+
if lokr_w2_name in lora.keys():
|
| 226 |
+
lokr_w2 = lora[lokr_w2_name]
|
| 227 |
+
loaded_keys.add(lokr_w2_name)
|
| 228 |
+
|
| 229 |
+
lokr_w1_a = None
|
| 230 |
+
if lokr_w1_a_name in lora.keys():
|
| 231 |
+
lokr_w1_a = lora[lokr_w1_a_name]
|
| 232 |
+
loaded_keys.add(lokr_w1_a_name)
|
| 233 |
+
|
| 234 |
+
lokr_w1_b = None
|
| 235 |
+
if lokr_w1_b_name in lora.keys():
|
| 236 |
+
lokr_w1_b = lora[lokr_w1_b_name]
|
| 237 |
+
loaded_keys.add(lokr_w1_b_name)
|
| 238 |
+
|
| 239 |
+
lokr_w2_a = None
|
| 240 |
+
if lokr_w2_a_name in lora.keys():
|
| 241 |
+
lokr_w2_a = lora[lokr_w2_a_name]
|
| 242 |
+
loaded_keys.add(lokr_w2_a_name)
|
| 243 |
+
|
| 244 |
+
lokr_w2_b = None
|
| 245 |
+
if lokr_w2_b_name in lora.keys():
|
| 246 |
+
lokr_w2_b = lora[lokr_w2_b_name]
|
| 247 |
+
loaded_keys.add(lokr_w2_b_name)
|
| 248 |
+
|
| 249 |
+
lokr_t2 = None
|
| 250 |
+
if lokr_t2_name in lora.keys():
|
| 251 |
+
lokr_t2 = lora[lokr_t2_name]
|
| 252 |
+
loaded_keys.add(lokr_t2_name)
|
| 253 |
+
|
| 254 |
+
if (
|
| 255 |
+
(lokr_w1 is not None)
|
| 256 |
+
or (lokr_w2 is not None)
|
| 257 |
+
or (lokr_w1_a is not None)
|
| 258 |
+
or (lokr_w2_a is not None)
|
| 259 |
+
):
|
| 260 |
+
weights = (
|
| 261 |
+
lokr_w1,
|
| 262 |
+
lokr_w2,
|
| 263 |
+
alpha,
|
| 264 |
+
lokr_w1_a,
|
| 265 |
+
lokr_w1_b,
|
| 266 |
+
lokr_w2_a,
|
| 267 |
+
lokr_w2_b,
|
| 268 |
+
lokr_t2,
|
| 269 |
+
dora_scale,
|
| 270 |
+
)
|
| 271 |
+
return cls(loaded_keys, weights)
|
| 272 |
+
else:
|
| 273 |
+
return None
|
| 274 |
+
|
| 275 |
+
def calculate_weight(
|
| 276 |
+
self,
|
| 277 |
+
weight,
|
| 278 |
+
key,
|
| 279 |
+
strength,
|
| 280 |
+
strength_model,
|
| 281 |
+
offset,
|
| 282 |
+
function,
|
| 283 |
+
intermediate_dtype=torch.float32,
|
| 284 |
+
original_weight=None,
|
| 285 |
+
):
|
| 286 |
+
v = self.weights
|
| 287 |
+
w1 = v[0]
|
| 288 |
+
w2 = v[1]
|
| 289 |
+
w1_a = v[3]
|
| 290 |
+
w1_b = v[4]
|
| 291 |
+
w2_a = v[5]
|
| 292 |
+
w2_b = v[6]
|
| 293 |
+
t2 = v[7]
|
| 294 |
+
dora_scale = v[8]
|
| 295 |
+
dim = None
|
| 296 |
+
|
| 297 |
+
if w1 is None:
|
| 298 |
+
dim = w1_b.shape[0]
|
| 299 |
+
w1 = torch.mm(
|
| 300 |
+
comfy.model_management.cast_to_device(
|
| 301 |
+
w1_a, weight.device, intermediate_dtype
|
| 302 |
+
),
|
| 303 |
+
comfy.model_management.cast_to_device(
|
| 304 |
+
w1_b, weight.device, intermediate_dtype
|
| 305 |
+
),
|
| 306 |
+
)
|
| 307 |
+
else:
|
| 308 |
+
w1 = comfy.model_management.cast_to_device(
|
| 309 |
+
w1, weight.device, intermediate_dtype
|
| 310 |
+
)
|
| 311 |
+
|
| 312 |
+
if w2 is None:
|
| 313 |
+
dim = w2_b.shape[0]
|
| 314 |
+
if t2 is None:
|
| 315 |
+
w2 = torch.mm(
|
| 316 |
+
comfy.model_management.cast_to_device(
|
| 317 |
+
w2_a, weight.device, intermediate_dtype
|
| 318 |
+
),
|
| 319 |
+
comfy.model_management.cast_to_device(
|
| 320 |
+
w2_b, weight.device, intermediate_dtype
|
| 321 |
+
),
|
| 322 |
+
)
|
| 323 |
+
else:
|
| 324 |
+
w2 = torch.einsum(
|
| 325 |
+
"i j k l, j r, i p -> p r k l",
|
| 326 |
+
comfy.model_management.cast_to_device(
|
| 327 |
+
t2, weight.device, intermediate_dtype
|
| 328 |
+
),
|
| 329 |
+
comfy.model_management.cast_to_device(
|
| 330 |
+
w2_b, weight.device, intermediate_dtype
|
| 331 |
+
),
|
| 332 |
+
comfy.model_management.cast_to_device(
|
| 333 |
+
w2_a, weight.device, intermediate_dtype
|
| 334 |
+
),
|
| 335 |
+
)
|
| 336 |
+
else:
|
| 337 |
+
w2 = comfy.model_management.cast_to_device(
|
| 338 |
+
w2, weight.device, intermediate_dtype
|
| 339 |
+
)
|
| 340 |
+
|
| 341 |
+
if len(w2.shape) == 4:
|
| 342 |
+
w1 = w1.unsqueeze(2).unsqueeze(2)
|
| 343 |
+
if v[2] is not None and dim is not None:
|
| 344 |
+
alpha = v[2] / dim
|
| 345 |
+
else:
|
| 346 |
+
alpha = 1.0
|
| 347 |
+
|
| 348 |
+
try:
|
| 349 |
+
lora_diff = torch.kron(w1, w2).reshape(weight.shape)
|
| 350 |
+
if dora_scale is not None:
|
| 351 |
+
weight = weight_decompose(
|
| 352 |
+
dora_scale,
|
| 353 |
+
weight,
|
| 354 |
+
lora_diff,
|
| 355 |
+
alpha,
|
| 356 |
+
strength,
|
| 357 |
+
intermediate_dtype,
|
| 358 |
+
function,
|
| 359 |
+
)
|
| 360 |
+
else:
|
| 361 |
+
weight += function(((strength * alpha) * lora_diff).type(weight.dtype))
|
| 362 |
+
except Exception as e:
|
| 363 |
+
logging.error("ERROR {} {} {}".format(self.name, key, e))
|
| 364 |
+
return weight
|
| 365 |
+
|
| 366 |
+
def h(self, x: torch.Tensor, base_out: torch.Tensor) -> torch.Tensor:
|
| 367 |
+
"""
|
| 368 |
+
Additive bypass component for LoKr: efficient Kronecker product application.
|
| 369 |
+
|
| 370 |
+
Note:
|
| 371 |
+
Does not access original model weights - bypass mode is designed
|
| 372 |
+
for quantized models where weights may not be accessible.
|
| 373 |
+
|
| 374 |
+
Args:
|
| 375 |
+
x: Input tensor
|
| 376 |
+
base_out: Output from base forward (unused, for API consistency)
|
| 377 |
+
|
| 378 |
+
Reference: LyCORIS functional/lokr.py bypass_forward_diff
|
| 379 |
+
"""
|
| 380 |
+
# FUNC_LIST: [None, None, F.linear, F.conv1d, F.conv2d, F.conv3d]
|
| 381 |
+
FUNC_LIST = [None, None, F.linear, F.conv1d, F.conv2d, F.conv3d]
|
| 382 |
+
|
| 383 |
+
v = self.weights
|
| 384 |
+
# v[0]=w1, v[1]=w2, v[2]=alpha, v[3]=w1_a, v[4]=w1_b, v[5]=w2_a, v[6]=w2_b, v[7]=t2, v[8]=dora
|
| 385 |
+
w1 = v[0]
|
| 386 |
+
w2 = v[1]
|
| 387 |
+
alpha = v[2]
|
| 388 |
+
w1_a = v[3]
|
| 389 |
+
w1_b = v[4]
|
| 390 |
+
w2_a = v[5]
|
| 391 |
+
w2_b = v[6]
|
| 392 |
+
t2 = v[7]
|
| 393 |
+
|
| 394 |
+
use_w1 = w1 is not None
|
| 395 |
+
use_w2 = w2 is not None
|
| 396 |
+
tucker = t2 is not None
|
| 397 |
+
|
| 398 |
+
# Use module info from bypass injection, not weight dimension
|
| 399 |
+
is_conv = getattr(self, "is_conv", False)
|
| 400 |
+
conv_dim = getattr(self, "conv_dim", 0)
|
| 401 |
+
kw_dict = getattr(self, "kw_dict", {}) if is_conv else {}
|
| 402 |
+
|
| 403 |
+
if is_conv:
|
| 404 |
+
op = FUNC_LIST[conv_dim + 2]
|
| 405 |
+
else:
|
| 406 |
+
op = F.linear
|
| 407 |
+
|
| 408 |
+
# Determine rank and scale
|
| 409 |
+
rank = w1_b.size(0) if not use_w1 else w2_b.size(0) if not use_w2 else alpha
|
| 410 |
+
scale = (alpha / rank if alpha is not None else 1.0) * getattr(
|
| 411 |
+
self, "multiplier", 1.0
|
| 412 |
+
)
|
| 413 |
+
|
| 414 |
+
# Build c (w1)
|
| 415 |
+
if use_w1:
|
| 416 |
+
c = w1.to(dtype=x.dtype)
|
| 417 |
+
else:
|
| 418 |
+
c = w1_a.to(dtype=x.dtype) @ w1_b.to(dtype=x.dtype)
|
| 419 |
+
uq = c.size(1)
|
| 420 |
+
|
| 421 |
+
# Build w2 components
|
| 422 |
+
if use_w2:
|
| 423 |
+
ba = w2.to(dtype=x.dtype)
|
| 424 |
+
else:
|
| 425 |
+
a = w2_b.to(dtype=x.dtype)
|
| 426 |
+
b = w2_a.to(dtype=x.dtype)
|
| 427 |
+
if is_conv:
|
| 428 |
+
if tucker:
|
| 429 |
+
# Tucker: a, b get 1s appended (kernel is in t2)
|
| 430 |
+
if a.dim() == 2:
|
| 431 |
+
a = a.view(*a.shape, *([1] * conv_dim))
|
| 432 |
+
if b.dim() == 2:
|
| 433 |
+
b = b.view(*b.shape, *([1] * conv_dim))
|
| 434 |
+
else:
|
| 435 |
+
# Non-tucker conv: b may need 1s appended
|
| 436 |
+
if b.dim() == 2:
|
| 437 |
+
b = b.view(*b.shape, *([1] * conv_dim))
|
| 438 |
+
|
| 439 |
+
# Reshape input by uq groups
|
| 440 |
+
if is_conv:
|
| 441 |
+
B, _, *rest = x.shape
|
| 442 |
+
h_in_group = x.reshape(B * uq, -1, *rest)
|
| 443 |
+
else:
|
| 444 |
+
h_in_group = x.reshape(*x.shape[:-1], uq, -1)
|
| 445 |
+
|
| 446 |
+
# Apply w2 path
|
| 447 |
+
if use_w2:
|
| 448 |
+
hb = op(h_in_group, ba, **kw_dict)
|
| 449 |
+
else:
|
| 450 |
+
if is_conv:
|
| 451 |
+
if tucker:
|
| 452 |
+
t = t2.to(dtype=x.dtype)
|
| 453 |
+
if t.dim() == 2:
|
| 454 |
+
t = t.view(*t.shape, *([1] * conv_dim))
|
| 455 |
+
ha = op(h_in_group, a)
|
| 456 |
+
ht = op(ha, t, **kw_dict)
|
| 457 |
+
hb = op(ht, b)
|
| 458 |
+
else:
|
| 459 |
+
ha = op(h_in_group, a, **kw_dict)
|
| 460 |
+
hb = op(ha, b)
|
| 461 |
+
else:
|
| 462 |
+
ha = op(h_in_group, a)
|
| 463 |
+
hb = op(ha, b)
|
| 464 |
+
|
| 465 |
+
# Reshape and apply c (w1)
|
| 466 |
+
if is_conv:
|
| 467 |
+
hb = hb.view(B, -1, *hb.shape[1:])
|
| 468 |
+
h_cross_group = hb.transpose(1, -1)
|
| 469 |
+
else:
|
| 470 |
+
h_cross_group = hb.transpose(-1, -2)
|
| 471 |
+
|
| 472 |
+
hc = F.linear(h_cross_group, c)
|
| 473 |
+
|
| 474 |
+
if is_conv:
|
| 475 |
+
hc = hc.transpose(1, -1)
|
| 476 |
+
out = hc.reshape(B, -1, *hc.shape[3:])
|
| 477 |
+
else:
|
| 478 |
+
hc = hc.transpose(-1, -2)
|
| 479 |
+
out = hc.reshape(*hc.shape[:-2], -1)
|
| 480 |
+
|
| 481 |
+
return out * scale
|
comfy/weight_adapter/lora.py
ADDED
|
@@ -0,0 +1,368 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from typing import Optional
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import torch.nn.functional as F
|
| 6 |
+
import comfy.model_management
|
| 7 |
+
from .base import (
|
| 8 |
+
WeightAdapterBase,
|
| 9 |
+
WeightAdapterTrainBase,
|
| 10 |
+
weight_decompose,
|
| 11 |
+
pad_tensor_to_shape,
|
| 12 |
+
tucker_weight_from_conv,
|
| 13 |
+
)
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
class LoraDiff(WeightAdapterTrainBase):
|
| 17 |
+
def __init__(self, weights):
|
| 18 |
+
super().__init__()
|
| 19 |
+
mat1, mat2, alpha, mid, dora_scale, reshape = weights
|
| 20 |
+
out_dim, rank = mat1.shape[0], mat1.shape[1]
|
| 21 |
+
rank, in_dim = mat2.shape[0], mat2.shape[1]
|
| 22 |
+
if mid is not None:
|
| 23 |
+
convdim = mid.ndim - 2
|
| 24 |
+
layer = (torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d)[convdim]
|
| 25 |
+
else:
|
| 26 |
+
layer = torch.nn.Linear
|
| 27 |
+
self.lora_up = layer(rank, out_dim, bias=False)
|
| 28 |
+
self.lora_down = layer(in_dim, rank, bias=False)
|
| 29 |
+
self.lora_up.weight.data.copy_(mat1)
|
| 30 |
+
self.lora_down.weight.data.copy_(mat2)
|
| 31 |
+
if mid is not None:
|
| 32 |
+
self.lora_mid = layer(mid, rank, bias=False)
|
| 33 |
+
self.lora_mid.weight.data.copy_(mid)
|
| 34 |
+
else:
|
| 35 |
+
self.lora_mid = None
|
| 36 |
+
self.rank = rank
|
| 37 |
+
self.alpha = torch.nn.Parameter(torch.tensor(alpha), requires_grad=False)
|
| 38 |
+
|
| 39 |
+
def __call__(self, w):
|
| 40 |
+
org_dtype = w.dtype
|
| 41 |
+
if self.lora_mid is None:
|
| 42 |
+
diff = self.lora_up.weight @ self.lora_down.weight
|
| 43 |
+
else:
|
| 44 |
+
diff = tucker_weight_from_conv(
|
| 45 |
+
self.lora_up.weight, self.lora_down.weight, self.lora_mid.weight
|
| 46 |
+
)
|
| 47 |
+
scale = self.alpha / self.rank
|
| 48 |
+
weight = w + scale * diff.reshape(w.shape)
|
| 49 |
+
return weight.to(org_dtype)
|
| 50 |
+
|
| 51 |
+
def h(self, x: torch.Tensor, base_out: torch.Tensor) -> torch.Tensor:
|
| 52 |
+
"""
|
| 53 |
+
Additive bypass component for LoRA training: h(x) = up(down(x)) * scale
|
| 54 |
+
|
| 55 |
+
Simple implementation using the nn.Module weights directly.
|
| 56 |
+
No mid/dora/reshape branches (create_train doesn't create them).
|
| 57 |
+
|
| 58 |
+
Args:
|
| 59 |
+
x: Input tensor
|
| 60 |
+
base_out: Output from base forward (unused, for API consistency)
|
| 61 |
+
"""
|
| 62 |
+
# Compute scale = alpha / rank * multiplier
|
| 63 |
+
scale = (self.alpha / self.rank) * getattr(self, "multiplier", 1.0)
|
| 64 |
+
|
| 65 |
+
# Get module info from bypass injection
|
| 66 |
+
is_conv = getattr(self, "is_conv", False)
|
| 67 |
+
conv_dim = getattr(self, "conv_dim", 0)
|
| 68 |
+
kw_dict = getattr(self, "kw_dict", {})
|
| 69 |
+
|
| 70 |
+
# Get weights (keep in original dtype for numerical stability)
|
| 71 |
+
down_weight = self.lora_down.weight
|
| 72 |
+
up_weight = self.lora_up.weight
|
| 73 |
+
|
| 74 |
+
if is_conv:
|
| 75 |
+
# Conv path: use functional conv
|
| 76 |
+
# conv_dim: 1=conv1d, 2=conv2d, 3=conv3d
|
| 77 |
+
conv_fn = (F.conv1d, F.conv2d, F.conv3d)[conv_dim - 1]
|
| 78 |
+
|
| 79 |
+
# Reshape 2D weights to conv format if needed
|
| 80 |
+
# down: [rank, in_features] -> [rank, in_channels, *kernel_size]
|
| 81 |
+
# up: [out_features, rank] -> [out_features, rank, 1, 1, ...]
|
| 82 |
+
if down_weight.dim() == 2:
|
| 83 |
+
kernel_size = getattr(self, "kernel_size", (1,) * conv_dim)
|
| 84 |
+
in_channels = getattr(self, "in_channels", None)
|
| 85 |
+
if in_channels is not None:
|
| 86 |
+
down_weight = down_weight.view(
|
| 87 |
+
down_weight.shape[0], in_channels, *kernel_size
|
| 88 |
+
)
|
| 89 |
+
else:
|
| 90 |
+
# Fallback: assume 1x1 kernel
|
| 91 |
+
down_weight = down_weight.view(
|
| 92 |
+
*down_weight.shape, *([1] * conv_dim)
|
| 93 |
+
)
|
| 94 |
+
if up_weight.dim() == 2:
|
| 95 |
+
# up always uses 1x1 kernel
|
| 96 |
+
up_weight = up_weight.view(*up_weight.shape, *([1] * conv_dim))
|
| 97 |
+
|
| 98 |
+
# down conv uses stride/padding from module, up is 1x1
|
| 99 |
+
hidden = conv_fn(x, down_weight, **kw_dict)
|
| 100 |
+
|
| 101 |
+
# mid layer if exists (tucker decomposition)
|
| 102 |
+
if self.lora_mid is not None:
|
| 103 |
+
mid_weight = self.lora_mid.weight
|
| 104 |
+
if mid_weight.dim() == 2:
|
| 105 |
+
mid_weight = mid_weight.view(*mid_weight.shape, *([1] * conv_dim))
|
| 106 |
+
hidden = conv_fn(hidden, mid_weight)
|
| 107 |
+
|
| 108 |
+
# up conv is always 1x1 (no stride/padding)
|
| 109 |
+
out = conv_fn(hidden, up_weight)
|
| 110 |
+
else:
|
| 111 |
+
# Linear path: simple matmul chain
|
| 112 |
+
hidden = F.linear(x, down_weight)
|
| 113 |
+
|
| 114 |
+
# mid layer if exists
|
| 115 |
+
if self.lora_mid is not None:
|
| 116 |
+
mid_weight = self.lora_mid.weight
|
| 117 |
+
hidden = F.linear(hidden, mid_weight)
|
| 118 |
+
|
| 119 |
+
out = F.linear(hidden, up_weight)
|
| 120 |
+
|
| 121 |
+
return out * scale
|
| 122 |
+
|
| 123 |
+
def passive_memory_usage(self):
|
| 124 |
+
return sum(param.numel() * param.element_size() for param in self.parameters())
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
class LoRAAdapter(WeightAdapterBase):
|
| 128 |
+
name = "lora"
|
| 129 |
+
|
| 130 |
+
def __init__(self, loaded_keys, weights):
|
| 131 |
+
self.loaded_keys = loaded_keys
|
| 132 |
+
self.weights = weights
|
| 133 |
+
|
| 134 |
+
@classmethod
|
| 135 |
+
def create_train(cls, weight, rank=1, alpha=1.0):
|
| 136 |
+
out_dim = weight.shape[0]
|
| 137 |
+
in_dim = weight.shape[1:].numel()
|
| 138 |
+
mat1 = torch.empty(out_dim, rank, device=weight.device, dtype=torch.float32)
|
| 139 |
+
mat2 = torch.empty(rank, in_dim, device=weight.device, dtype=torch.float32)
|
| 140 |
+
torch.nn.init.kaiming_uniform_(mat1, a=5**0.5)
|
| 141 |
+
torch.nn.init.constant_(mat2, 0.0)
|
| 142 |
+
return LoraDiff((mat1, mat2, alpha, None, None, None))
|
| 143 |
+
|
| 144 |
+
def to_train(self):
|
| 145 |
+
return LoraDiff(self.weights)
|
| 146 |
+
|
| 147 |
+
@classmethod
|
| 148 |
+
def load(
|
| 149 |
+
cls,
|
| 150 |
+
x: str,
|
| 151 |
+
lora: dict[str, torch.Tensor],
|
| 152 |
+
alpha: float,
|
| 153 |
+
dora_scale: torch.Tensor,
|
| 154 |
+
loaded_keys: set[str] = None,
|
| 155 |
+
) -> Optional["LoRAAdapter"]:
|
| 156 |
+
if loaded_keys is None:
|
| 157 |
+
loaded_keys = set()
|
| 158 |
+
|
| 159 |
+
reshape_name = "{}.reshape_weight".format(x)
|
| 160 |
+
regular_lora = "{}.lora_up.weight".format(x)
|
| 161 |
+
diffusers_lora = "{}_lora.up.weight".format(x)
|
| 162 |
+
diffusers2_lora = "{}.lora_B.weight".format(x)
|
| 163 |
+
diffusers3_lora = "{}.lora.up.weight".format(x)
|
| 164 |
+
mochi_lora = "{}.lora_B".format(x)
|
| 165 |
+
transformers_lora = "{}.lora_linear_layer.up.weight".format(x)
|
| 166 |
+
qwen_default_lora = "{}.lora_B.default.weight".format(x)
|
| 167 |
+
A_name = None
|
| 168 |
+
|
| 169 |
+
if regular_lora in lora.keys():
|
| 170 |
+
A_name = regular_lora
|
| 171 |
+
B_name = "{}.lora_down.weight".format(x)
|
| 172 |
+
mid_name = "{}.lora_mid.weight".format(x)
|
| 173 |
+
elif diffusers_lora in lora.keys():
|
| 174 |
+
A_name = diffusers_lora
|
| 175 |
+
B_name = "{}_lora.down.weight".format(x)
|
| 176 |
+
mid_name = None
|
| 177 |
+
elif diffusers2_lora in lora.keys():
|
| 178 |
+
A_name = diffusers2_lora
|
| 179 |
+
B_name = "{}.lora_A.weight".format(x)
|
| 180 |
+
mid_name = None
|
| 181 |
+
elif diffusers3_lora in lora.keys():
|
| 182 |
+
A_name = diffusers3_lora
|
| 183 |
+
B_name = "{}.lora.down.weight".format(x)
|
| 184 |
+
mid_name = None
|
| 185 |
+
elif mochi_lora in lora.keys():
|
| 186 |
+
A_name = mochi_lora
|
| 187 |
+
B_name = "{}.lora_A".format(x)
|
| 188 |
+
mid_name = None
|
| 189 |
+
elif transformers_lora in lora.keys():
|
| 190 |
+
A_name = transformers_lora
|
| 191 |
+
B_name = "{}.lora_linear_layer.down.weight".format(x)
|
| 192 |
+
mid_name = None
|
| 193 |
+
elif qwen_default_lora in lora.keys():
|
| 194 |
+
A_name = qwen_default_lora
|
| 195 |
+
B_name = "{}.lora_A.default.weight".format(x)
|
| 196 |
+
mid_name = None
|
| 197 |
+
|
| 198 |
+
if A_name is not None:
|
| 199 |
+
mid = None
|
| 200 |
+
if mid_name is not None and mid_name in lora.keys():
|
| 201 |
+
mid = lora[mid_name]
|
| 202 |
+
loaded_keys.add(mid_name)
|
| 203 |
+
reshape = None
|
| 204 |
+
if reshape_name in lora.keys():
|
| 205 |
+
try:
|
| 206 |
+
reshape = lora[reshape_name].tolist()
|
| 207 |
+
loaded_keys.add(reshape_name)
|
| 208 |
+
except:
|
| 209 |
+
pass
|
| 210 |
+
weights = (lora[A_name], lora[B_name], alpha, mid, dora_scale, reshape)
|
| 211 |
+
loaded_keys.add(A_name)
|
| 212 |
+
loaded_keys.add(B_name)
|
| 213 |
+
return cls(loaded_keys, weights)
|
| 214 |
+
else:
|
| 215 |
+
return None
|
| 216 |
+
|
| 217 |
+
def calculate_shape(
|
| 218 |
+
self,
|
| 219 |
+
key
|
| 220 |
+
):
|
| 221 |
+
reshape = self.weights[5]
|
| 222 |
+
return tuple(reshape) if reshape is not None else None
|
| 223 |
+
|
| 224 |
+
def calculate_weight(
|
| 225 |
+
self,
|
| 226 |
+
weight,
|
| 227 |
+
key,
|
| 228 |
+
strength,
|
| 229 |
+
strength_model,
|
| 230 |
+
offset,
|
| 231 |
+
function,
|
| 232 |
+
intermediate_dtype=torch.float32,
|
| 233 |
+
original_weight=None,
|
| 234 |
+
):
|
| 235 |
+
v = self.weights
|
| 236 |
+
mat1 = comfy.model_management.cast_to_device(
|
| 237 |
+
v[0], weight.device, intermediate_dtype
|
| 238 |
+
)
|
| 239 |
+
mat2 = comfy.model_management.cast_to_device(
|
| 240 |
+
v[1], weight.device, intermediate_dtype
|
| 241 |
+
)
|
| 242 |
+
dora_scale = v[4]
|
| 243 |
+
reshape = v[5]
|
| 244 |
+
|
| 245 |
+
if reshape is not None:
|
| 246 |
+
weight = pad_tensor_to_shape(weight, reshape)
|
| 247 |
+
|
| 248 |
+
if v[2] is not None:
|
| 249 |
+
alpha = v[2] / mat2.shape[0]
|
| 250 |
+
else:
|
| 251 |
+
alpha = 1.0
|
| 252 |
+
|
| 253 |
+
if v[3] is not None:
|
| 254 |
+
# locon mid weights, hopefully the math is fine because I didn't properly test it
|
| 255 |
+
mat3 = comfy.model_management.cast_to_device(
|
| 256 |
+
v[3], weight.device, intermediate_dtype
|
| 257 |
+
)
|
| 258 |
+
final_shape = [mat2.shape[1], mat2.shape[0], mat3.shape[2], mat3.shape[3]]
|
| 259 |
+
mat2 = (
|
| 260 |
+
torch.mm(
|
| 261 |
+
mat2.transpose(0, 1).flatten(start_dim=1),
|
| 262 |
+
mat3.transpose(0, 1).flatten(start_dim=1),
|
| 263 |
+
)
|
| 264 |
+
.reshape(final_shape)
|
| 265 |
+
.transpose(0, 1)
|
| 266 |
+
)
|
| 267 |
+
try:
|
| 268 |
+
lora_diff = torch.mm(
|
| 269 |
+
mat1.flatten(start_dim=1), mat2.flatten(start_dim=1)
|
| 270 |
+
).reshape(weight.shape)
|
| 271 |
+
del mat1, mat2
|
| 272 |
+
if dora_scale is not None:
|
| 273 |
+
weight = weight_decompose(
|
| 274 |
+
dora_scale,
|
| 275 |
+
weight,
|
| 276 |
+
lora_diff,
|
| 277 |
+
alpha,
|
| 278 |
+
strength,
|
| 279 |
+
intermediate_dtype,
|
| 280 |
+
function,
|
| 281 |
+
)
|
| 282 |
+
else:
|
| 283 |
+
weight += function(((strength * alpha) * lora_diff).type(weight.dtype))
|
| 284 |
+
except Exception as e:
|
| 285 |
+
logging.error("ERROR {} {} {}".format(self.name, key, e))
|
| 286 |
+
return weight
|
| 287 |
+
|
| 288 |
+
def h(self, x: torch.Tensor, base_out: torch.Tensor) -> torch.Tensor:
|
| 289 |
+
"""
|
| 290 |
+
Additive bypass component for LoRA: h(x) = up(down(x)) * scale
|
| 291 |
+
|
| 292 |
+
Note:
|
| 293 |
+
Does not access original model weights - bypass mode is designed
|
| 294 |
+
for quantized models where weights may not be accessible.
|
| 295 |
+
|
| 296 |
+
Args:
|
| 297 |
+
x: Input tensor
|
| 298 |
+
base_out: Output from base forward (unused, for API consistency)
|
| 299 |
+
|
| 300 |
+
Reference: LyCORIS functional/locon.py bypass_forward_diff
|
| 301 |
+
"""
|
| 302 |
+
# FUNC_LIST: [None, None, F.linear, F.conv1d, F.conv2d, F.conv3d]
|
| 303 |
+
FUNC_LIST = [None, None, F.linear, F.conv1d, F.conv2d, F.conv3d]
|
| 304 |
+
|
| 305 |
+
v = self.weights
|
| 306 |
+
# v[0]=up, v[1]=down, v[2]=alpha, v[3]=mid, v[4]=dora_scale, v[5]=reshape
|
| 307 |
+
up = v[0]
|
| 308 |
+
down = v[1]
|
| 309 |
+
alpha = v[2]
|
| 310 |
+
mid = v[3]
|
| 311 |
+
|
| 312 |
+
# Compute scale = alpha / rank
|
| 313 |
+
rank = down.shape[0]
|
| 314 |
+
if alpha is not None:
|
| 315 |
+
scale = alpha / rank
|
| 316 |
+
else:
|
| 317 |
+
scale = 1.0
|
| 318 |
+
scale = scale * getattr(self, "multiplier", 1.0)
|
| 319 |
+
|
| 320 |
+
# Cast dtype
|
| 321 |
+
up = up.to(dtype=x.dtype)
|
| 322 |
+
down = down.to(dtype=x.dtype)
|
| 323 |
+
|
| 324 |
+
# Use module info from bypass injection, not weight dimension
|
| 325 |
+
is_conv = getattr(self, "is_conv", False)
|
| 326 |
+
conv_dim = getattr(self, "conv_dim", 0)
|
| 327 |
+
kw_dict = getattr(self, "kw_dict", {})
|
| 328 |
+
|
| 329 |
+
if is_conv:
|
| 330 |
+
op = FUNC_LIST[
|
| 331 |
+
conv_dim + 2
|
| 332 |
+
] # conv_dim 1->conv1d(3), 2->conv2d(4), 3->conv3d(5)
|
| 333 |
+
kernel_size = getattr(self, "kernel_size", (1,) * conv_dim)
|
| 334 |
+
in_channels = getattr(self, "in_channels", None)
|
| 335 |
+
|
| 336 |
+
# Reshape 2D weights to conv format using kernel_size
|
| 337 |
+
# down: [rank, in_channels * prod(kernel_size)] -> [rank, in_channels, *kernel_size]
|
| 338 |
+
# up: [out_channels, rank] -> [out_channels, rank, 1, 1, ...] (1x1 kernel)
|
| 339 |
+
if down.dim() == 2:
|
| 340 |
+
# down.shape[1] = in_channels * prod(kernel_size)
|
| 341 |
+
if in_channels is not None:
|
| 342 |
+
down = down.view(down.shape[0], in_channels, *kernel_size)
|
| 343 |
+
else:
|
| 344 |
+
# Fallback: assume 1x1 kernel if in_channels unknown
|
| 345 |
+
down = down.view(*down.shape, *([1] * conv_dim))
|
| 346 |
+
if up.dim() == 2:
|
| 347 |
+
# up always uses 1x1 kernel
|
| 348 |
+
up = up.view(*up.shape, *([1] * conv_dim))
|
| 349 |
+
if mid is not None:
|
| 350 |
+
mid = mid.to(dtype=x.dtype)
|
| 351 |
+
if mid.dim() == 2:
|
| 352 |
+
mid = mid.view(*mid.shape, *([1] * conv_dim))
|
| 353 |
+
else:
|
| 354 |
+
op = F.linear
|
| 355 |
+
kw_dict = {} # linear doesn't take stride/padding
|
| 356 |
+
|
| 357 |
+
# Simple chain: down -> mid (if tucker) -> up
|
| 358 |
+
if mid is not None:
|
| 359 |
+
if not is_conv:
|
| 360 |
+
mid = mid.to(dtype=x.dtype)
|
| 361 |
+
hidden = op(x, down)
|
| 362 |
+
hidden = op(hidden, mid, **kw_dict)
|
| 363 |
+
out = op(hidden, up)
|
| 364 |
+
else:
|
| 365 |
+
hidden = op(x, down, **kw_dict)
|
| 366 |
+
out = op(hidden, up)
|
| 367 |
+
|
| 368 |
+
return out * scale
|
comfy/weight_adapter/oft.py
ADDED
|
@@ -0,0 +1,327 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from typing import Optional
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import comfy.model_management
|
| 6 |
+
from .base import (
|
| 7 |
+
WeightAdapterBase,
|
| 8 |
+
WeightAdapterTrainBase,
|
| 9 |
+
weight_decompose,
|
| 10 |
+
factorization,
|
| 11 |
+
)
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class OFTDiff(WeightAdapterTrainBase):
|
| 15 |
+
def __init__(self, weights):
|
| 16 |
+
super().__init__()
|
| 17 |
+
# Unpack weights tuple from OFTAdapter
|
| 18 |
+
blocks, rescale, alpha, _ = weights
|
| 19 |
+
|
| 20 |
+
# Create trainable parameters
|
| 21 |
+
self.oft_blocks = torch.nn.Parameter(blocks)
|
| 22 |
+
if rescale is not None:
|
| 23 |
+
self.rescale = torch.nn.Parameter(rescale)
|
| 24 |
+
self.rescaled = True
|
| 25 |
+
else:
|
| 26 |
+
self.rescaled = False
|
| 27 |
+
self.block_num, self.block_size, _ = blocks.shape
|
| 28 |
+
self.constraint = float(alpha)
|
| 29 |
+
self.alpha = torch.nn.Parameter(torch.tensor(alpha), requires_grad=False)
|
| 30 |
+
|
| 31 |
+
def __call__(self, w):
|
| 32 |
+
org_dtype = w.dtype
|
| 33 |
+
I = torch.eye(self.block_size, device=self.oft_blocks.device)
|
| 34 |
+
|
| 35 |
+
## generate r
|
| 36 |
+
# for Q = -Q^T
|
| 37 |
+
q = self.oft_blocks - self.oft_blocks.transpose(1, 2)
|
| 38 |
+
normed_q = q
|
| 39 |
+
if self.constraint:
|
| 40 |
+
q_norm = torch.norm(q) + 1e-8
|
| 41 |
+
if q_norm > self.constraint:
|
| 42 |
+
normed_q = q * self.constraint / q_norm
|
| 43 |
+
# use float() to prevent unsupported type
|
| 44 |
+
r = (I + normed_q) @ (I - normed_q).float().inverse()
|
| 45 |
+
|
| 46 |
+
## Apply chunked matmul on weight
|
| 47 |
+
_, *shape = w.shape
|
| 48 |
+
org_weight = w.to(dtype=r.dtype)
|
| 49 |
+
org_weight = org_weight.unflatten(0, (self.block_num, self.block_size))
|
| 50 |
+
# Init R=0, so add I on it to ensure the output of step0 is original model output
|
| 51 |
+
weight = torch.einsum(
|
| 52 |
+
"k n m, k n ... -> k m ...",
|
| 53 |
+
r,
|
| 54 |
+
org_weight,
|
| 55 |
+
).flatten(0, 1)
|
| 56 |
+
if self.rescaled:
|
| 57 |
+
weight = self.rescale * weight
|
| 58 |
+
return weight.to(org_dtype)
|
| 59 |
+
|
| 60 |
+
def _get_orthogonal_matrix(self, device, dtype):
|
| 61 |
+
"""Compute the orthogonal rotation matrix R from OFT blocks."""
|
| 62 |
+
blocks = self.oft_blocks.to(device=device, dtype=dtype)
|
| 63 |
+
I = torch.eye(self.block_size, device=device, dtype=dtype)
|
| 64 |
+
|
| 65 |
+
# Q = blocks - blocks^T (skew-symmetric)
|
| 66 |
+
q = blocks - blocks.transpose(1, 2)
|
| 67 |
+
normed_q = q
|
| 68 |
+
|
| 69 |
+
# Apply constraint if set
|
| 70 |
+
if self.constraint:
|
| 71 |
+
q_norm = torch.norm(q) + 1e-8
|
| 72 |
+
if q_norm > self.constraint:
|
| 73 |
+
normed_q = q * self.constraint / q_norm
|
| 74 |
+
|
| 75 |
+
# Cayley transform: R = (I + Q)(I - Q)^-1
|
| 76 |
+
r = (I + normed_q) @ (I - normed_q).float().inverse()
|
| 77 |
+
return r.to(dtype)
|
| 78 |
+
|
| 79 |
+
def h(self, x: torch.Tensor, base_out: torch.Tensor) -> torch.Tensor:
|
| 80 |
+
"""
|
| 81 |
+
OFT has no additive component - returns zeros matching base_out shape.
|
| 82 |
+
|
| 83 |
+
OFT only transforms the output via g(), it doesn't add to it.
|
| 84 |
+
"""
|
| 85 |
+
return torch.zeros_like(base_out)
|
| 86 |
+
|
| 87 |
+
def g(self, y: torch.Tensor) -> torch.Tensor:
|
| 88 |
+
"""
|
| 89 |
+
Output transformation for OFT: applies orthogonal rotation.
|
| 90 |
+
|
| 91 |
+
OFT transforms output channels using block-diagonal orthogonal matrices.
|
| 92 |
+
"""
|
| 93 |
+
r = self._get_orthogonal_matrix(y.device, y.dtype)
|
| 94 |
+
|
| 95 |
+
# Apply multiplier to interpolate between identity and full transform
|
| 96 |
+
multiplier = getattr(self, "multiplier", 1.0)
|
| 97 |
+
I = torch.eye(self.block_size, device=y.device, dtype=y.dtype)
|
| 98 |
+
r = r * multiplier + (1 - multiplier) * I
|
| 99 |
+
|
| 100 |
+
# Use module info from bypass injection
|
| 101 |
+
is_conv = getattr(self, "is_conv", y.dim() > 2)
|
| 102 |
+
|
| 103 |
+
if is_conv:
|
| 104 |
+
# Conv output: (N, C, H, W, ...) -> transpose to (N, H, W, ..., C)
|
| 105 |
+
y = y.transpose(1, -1)
|
| 106 |
+
|
| 107 |
+
# y now has channels in last dim
|
| 108 |
+
*batch_shape, out_features = y.shape
|
| 109 |
+
|
| 110 |
+
# Reshape to apply block-diagonal transform
|
| 111 |
+
# (*, out_features) -> (*, block_num, block_size)
|
| 112 |
+
y_blocked = y.reshape(*batch_shape, self.block_num, self.block_size)
|
| 113 |
+
|
| 114 |
+
# Apply orthogonal transform: R @ y for each block
|
| 115 |
+
# r: (block_num, block_size, block_size), y_blocked: (*, block_num, block_size)
|
| 116 |
+
out_blocked = torch.einsum("k n m, ... k n -> ... k m", r, y_blocked)
|
| 117 |
+
|
| 118 |
+
# Reshape back: (*, block_num, block_size) -> (*, out_features)
|
| 119 |
+
out = out_blocked.reshape(*batch_shape, out_features)
|
| 120 |
+
|
| 121 |
+
# Apply rescale if present
|
| 122 |
+
if self.rescaled:
|
| 123 |
+
rescale = self.rescale.to(device=y.device, dtype=y.dtype)
|
| 124 |
+
out = out * rescale.view(-1)
|
| 125 |
+
|
| 126 |
+
if is_conv:
|
| 127 |
+
# Transpose back: (N, H, W, ..., C) -> (N, C, H, W, ...)
|
| 128 |
+
out = out.transpose(1, -1)
|
| 129 |
+
|
| 130 |
+
return out
|
| 131 |
+
|
| 132 |
+
def passive_memory_usage(self):
|
| 133 |
+
"""Calculates memory usage of the trainable parameters."""
|
| 134 |
+
return sum(param.numel() * param.element_size() for param in self.parameters())
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
class OFTAdapter(WeightAdapterBase):
|
| 138 |
+
name = "oft"
|
| 139 |
+
|
| 140 |
+
def __init__(self, loaded_keys, weights):
|
| 141 |
+
self.loaded_keys = loaded_keys
|
| 142 |
+
self.weights = weights
|
| 143 |
+
|
| 144 |
+
@classmethod
|
| 145 |
+
def create_train(cls, weight, rank=1, alpha=1.0):
|
| 146 |
+
out_dim = weight.shape[0]
|
| 147 |
+
block_size, block_num = factorization(out_dim, rank)
|
| 148 |
+
block = torch.zeros(
|
| 149 |
+
block_num, block_size, block_size, device=weight.device, dtype=torch.float32
|
| 150 |
+
)
|
| 151 |
+
return OFTDiff((block, None, alpha, None))
|
| 152 |
+
|
| 153 |
+
def to_train(self):
|
| 154 |
+
return OFTDiff(self.weights)
|
| 155 |
+
|
| 156 |
+
@classmethod
|
| 157 |
+
def load(
|
| 158 |
+
cls,
|
| 159 |
+
x: str,
|
| 160 |
+
lora: dict[str, torch.Tensor],
|
| 161 |
+
alpha: float,
|
| 162 |
+
dora_scale: torch.Tensor,
|
| 163 |
+
loaded_keys: set[str] = None,
|
| 164 |
+
) -> Optional["OFTAdapter"]:
|
| 165 |
+
if loaded_keys is None:
|
| 166 |
+
loaded_keys = set()
|
| 167 |
+
blocks_name = "{}.oft_blocks".format(x)
|
| 168 |
+
rescale_name = "{}.rescale".format(x)
|
| 169 |
+
|
| 170 |
+
blocks = None
|
| 171 |
+
if blocks_name in lora.keys():
|
| 172 |
+
blocks = lora[blocks_name]
|
| 173 |
+
if blocks.ndim == 3:
|
| 174 |
+
loaded_keys.add(blocks_name)
|
| 175 |
+
else:
|
| 176 |
+
blocks = None
|
| 177 |
+
if blocks is None:
|
| 178 |
+
return None
|
| 179 |
+
|
| 180 |
+
rescale = None
|
| 181 |
+
if rescale_name in lora.keys():
|
| 182 |
+
rescale = lora[rescale_name]
|
| 183 |
+
loaded_keys.add(rescale_name)
|
| 184 |
+
|
| 185 |
+
weights = (blocks, rescale, alpha, dora_scale)
|
| 186 |
+
return cls(loaded_keys, weights)
|
| 187 |
+
|
| 188 |
+
def calculate_weight(
|
| 189 |
+
self,
|
| 190 |
+
weight,
|
| 191 |
+
key,
|
| 192 |
+
strength,
|
| 193 |
+
strength_model,
|
| 194 |
+
offset,
|
| 195 |
+
function,
|
| 196 |
+
intermediate_dtype=torch.float32,
|
| 197 |
+
original_weight=None,
|
| 198 |
+
):
|
| 199 |
+
v = self.weights
|
| 200 |
+
blocks = v[0]
|
| 201 |
+
rescale = v[1]
|
| 202 |
+
alpha = v[2]
|
| 203 |
+
if alpha is None:
|
| 204 |
+
alpha = 0
|
| 205 |
+
dora_scale = v[3]
|
| 206 |
+
|
| 207 |
+
blocks = comfy.model_management.cast_to_device(
|
| 208 |
+
blocks, weight.device, intermediate_dtype
|
| 209 |
+
)
|
| 210 |
+
if rescale is not None:
|
| 211 |
+
rescale = comfy.model_management.cast_to_device(
|
| 212 |
+
rescale, weight.device, intermediate_dtype
|
| 213 |
+
)
|
| 214 |
+
|
| 215 |
+
block_num, block_size, *_ = blocks.shape
|
| 216 |
+
|
| 217 |
+
try:
|
| 218 |
+
# Get r
|
| 219 |
+
I = torch.eye(block_size, device=blocks.device, dtype=blocks.dtype)
|
| 220 |
+
# for Q = -Q^T
|
| 221 |
+
q = blocks - blocks.transpose(1, 2)
|
| 222 |
+
normed_q = q
|
| 223 |
+
if alpha > 0: # alpha in oft/boft is for constraint
|
| 224 |
+
q_norm = torch.norm(q) + 1e-8
|
| 225 |
+
if q_norm > alpha:
|
| 226 |
+
normed_q = q * alpha / q_norm
|
| 227 |
+
# use float() to prevent unsupported type in .inverse()
|
| 228 |
+
r = (I + normed_q) @ (I - normed_q).float().inverse()
|
| 229 |
+
r = r.to(weight)
|
| 230 |
+
# Create I in weight's dtype for the einsum
|
| 231 |
+
I_w = torch.eye(block_size, device=weight.device, dtype=weight.dtype)
|
| 232 |
+
_, *shape = weight.shape
|
| 233 |
+
lora_diff = torch.einsum(
|
| 234 |
+
"k n m, k n ... -> k m ...",
|
| 235 |
+
(r * strength) - strength * I_w,
|
| 236 |
+
weight.view(block_num, block_size, *shape),
|
| 237 |
+
).view(-1, *shape)
|
| 238 |
+
if dora_scale is not None:
|
| 239 |
+
weight = weight_decompose(
|
| 240 |
+
dora_scale,
|
| 241 |
+
weight,
|
| 242 |
+
lora_diff,
|
| 243 |
+
alpha,
|
| 244 |
+
strength,
|
| 245 |
+
intermediate_dtype,
|
| 246 |
+
function,
|
| 247 |
+
)
|
| 248 |
+
else:
|
| 249 |
+
weight += function((strength * lora_diff).type(weight.dtype))
|
| 250 |
+
except Exception as e:
|
| 251 |
+
logging.error("ERROR {} {} {}".format(self.name, key, e))
|
| 252 |
+
return weight
|
| 253 |
+
|
| 254 |
+
def _get_orthogonal_matrix(self, device, dtype):
|
| 255 |
+
"""Compute the orthogonal rotation matrix R from OFT blocks."""
|
| 256 |
+
v = self.weights
|
| 257 |
+
blocks = v[0].to(device=device, dtype=dtype)
|
| 258 |
+
alpha = v[2]
|
| 259 |
+
if alpha is None:
|
| 260 |
+
alpha = 0
|
| 261 |
+
|
| 262 |
+
block_num, block_size, _ = blocks.shape
|
| 263 |
+
I = torch.eye(block_size, device=device, dtype=dtype)
|
| 264 |
+
|
| 265 |
+
# Q = blocks - blocks^T (skew-symmetric)
|
| 266 |
+
q = blocks - blocks.transpose(1, 2)
|
| 267 |
+
normed_q = q
|
| 268 |
+
|
| 269 |
+
# Apply constraint if alpha > 0
|
| 270 |
+
if alpha > 0:
|
| 271 |
+
q_norm = torch.norm(q) + 1e-8
|
| 272 |
+
if q_norm > alpha:
|
| 273 |
+
normed_q = q * alpha / q_norm
|
| 274 |
+
|
| 275 |
+
# Cayley transform: R = (I + Q)(I - Q)^-1
|
| 276 |
+
r = (I + normed_q) @ (I - normed_q).float().inverse()
|
| 277 |
+
return r, block_num, block_size
|
| 278 |
+
|
| 279 |
+
def g(self, y: torch.Tensor) -> torch.Tensor:
|
| 280 |
+
"""
|
| 281 |
+
Output transformation for OFT: applies orthogonal rotation to output.
|
| 282 |
+
|
| 283 |
+
OFT transforms the output channels using block-diagonal orthogonal matrices.
|
| 284 |
+
|
| 285 |
+
Reference: LyCORIS DiagOFTModule._bypass_forward
|
| 286 |
+
"""
|
| 287 |
+
v = self.weights
|
| 288 |
+
rescale = v[1]
|
| 289 |
+
|
| 290 |
+
r, block_num, block_size = self._get_orthogonal_matrix(y.device, y.dtype)
|
| 291 |
+
|
| 292 |
+
# Apply multiplier to interpolate between identity and full transform
|
| 293 |
+
multiplier = getattr(self, "multiplier", 1.0)
|
| 294 |
+
I = torch.eye(block_size, device=y.device, dtype=y.dtype)
|
| 295 |
+
r = r * multiplier + (1 - multiplier) * I
|
| 296 |
+
|
| 297 |
+
# Use module info from bypass injection to determine conv vs linear
|
| 298 |
+
is_conv = getattr(self, "is_conv", y.dim() > 2)
|
| 299 |
+
|
| 300 |
+
if is_conv:
|
| 301 |
+
# Conv output: (N, C, H, W, ...) -> transpose to (N, H, W, ..., C)
|
| 302 |
+
y = y.transpose(1, -1)
|
| 303 |
+
|
| 304 |
+
# y now has channels in last dim
|
| 305 |
+
*batch_shape, out_features = y.shape
|
| 306 |
+
|
| 307 |
+
# Reshape to apply block-diagonal transform
|
| 308 |
+
# (*, out_features) -> (*, block_num, block_size)
|
| 309 |
+
y_blocked = y.view(*batch_shape, block_num, block_size)
|
| 310 |
+
|
| 311 |
+
# Apply orthogonal transform: R @ y for each block
|
| 312 |
+
# r: (block_num, block_size, block_size), y_blocked: (*, block_num, block_size)
|
| 313 |
+
out_blocked = torch.einsum("k n m, ... k n -> ... k m", r, y_blocked)
|
| 314 |
+
|
| 315 |
+
# Reshape back: (*, block_num, block_size) -> (*, out_features)
|
| 316 |
+
out = out_blocked.view(*batch_shape, out_features)
|
| 317 |
+
|
| 318 |
+
# Apply rescale if present
|
| 319 |
+
if rescale is not None:
|
| 320 |
+
rescale = rescale.to(device=y.device, dtype=y.dtype)
|
| 321 |
+
out = out * rescale.view(-1)
|
| 322 |
+
|
| 323 |
+
if is_conv:
|
| 324 |
+
# Transpose back: (N, H, W, ..., C) -> (N, C, H, W, ...)
|
| 325 |
+
out = out.transpose(1, -1)
|
| 326 |
+
|
| 327 |
+
return out
|
comfy_api/feature_flags.py
ADDED
|
@@ -0,0 +1,166 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Feature flags module for ComfyUI WebSocket protocol negotiation.
|
| 3 |
+
|
| 4 |
+
This module handles capability negotiation between frontend and backend,
|
| 5 |
+
allowing graceful protocol evolution while maintaining backward compatibility.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
import logging
|
| 9 |
+
from typing import Any, TypedDict
|
| 10 |
+
|
| 11 |
+
from comfy.cli_args import args
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class FeatureFlagInfo(TypedDict):
|
| 15 |
+
type: str
|
| 16 |
+
default: Any
|
| 17 |
+
description: str
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
# Registry of known CLI-settable feature flags.
|
| 21 |
+
# Launchers can query this via --list-feature-flags to discover valid flags.
|
| 22 |
+
CLI_FEATURE_FLAG_REGISTRY: dict[str, FeatureFlagInfo] = {
|
| 23 |
+
"show_signin_button": {
|
| 24 |
+
"type": "bool",
|
| 25 |
+
"default": False,
|
| 26 |
+
"description": "Show the sign-in button in the frontend even when not signed in",
|
| 27 |
+
},
|
| 28 |
+
"enable_telemetry": {
|
| 29 |
+
"type": "bool",
|
| 30 |
+
"default": False,
|
| 31 |
+
"description": "Signal the frontend that telemetry collection is enabled",
|
| 32 |
+
},
|
| 33 |
+
}
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def _coerce_bool(v: str) -> bool:
|
| 37 |
+
"""Strict bool coercion: only 'true'/'false' (case-insensitive).
|
| 38 |
+
|
| 39 |
+
Anything else raises ValueError so the caller can warn and drop the flag,
|
| 40 |
+
rather than silently treating typos like 'ture' or 'yes' as False.
|
| 41 |
+
"""
|
| 42 |
+
lower = v.lower()
|
| 43 |
+
if lower == "true":
|
| 44 |
+
return True
|
| 45 |
+
if lower == "false":
|
| 46 |
+
return False
|
| 47 |
+
raise ValueError(f"expected 'true' or 'false', got {v!r}")
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
_COERCE_FNS: dict[str, Any] = {
|
| 51 |
+
"bool": _coerce_bool,
|
| 52 |
+
"int": lambda v: int(v),
|
| 53 |
+
"float": lambda v: float(v),
|
| 54 |
+
}
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def _coerce_flag_value(key: str, raw_value: str) -> Any:
|
| 58 |
+
"""Coerce a raw string value using the registry type, or keep as string.
|
| 59 |
+
|
| 60 |
+
Returns the raw string if the key is unregistered or the type is unknown.
|
| 61 |
+
Raises ValueError/TypeError if the key is registered with a known type but
|
| 62 |
+
the value cannot be coerced; callers are expected to warn and drop the flag.
|
| 63 |
+
"""
|
| 64 |
+
info = CLI_FEATURE_FLAG_REGISTRY.get(key)
|
| 65 |
+
if info is None:
|
| 66 |
+
return raw_value
|
| 67 |
+
coerce = _COERCE_FNS.get(info["type"])
|
| 68 |
+
if coerce is None:
|
| 69 |
+
return raw_value
|
| 70 |
+
return coerce(raw_value)
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
def _parse_cli_feature_flags() -> dict[str, Any]:
|
| 74 |
+
"""Parse --feature-flag key=value pairs from CLI args into a dict.
|
| 75 |
+
|
| 76 |
+
Items without '=' default to the value 'true' (bare flag form).
|
| 77 |
+
Flags whose value cannot be coerced to the registered type are dropped
|
| 78 |
+
with a warning, so a typo like '--feature-flag some_bool=ture' does not
|
| 79 |
+
silently take effect as the wrong value.
|
| 80 |
+
"""
|
| 81 |
+
result: dict[str, Any] = {}
|
| 82 |
+
for item in getattr(args, "feature_flag", []):
|
| 83 |
+
key, sep, raw_value = item.partition("=")
|
| 84 |
+
key = key.strip()
|
| 85 |
+
if not key:
|
| 86 |
+
continue
|
| 87 |
+
if not sep:
|
| 88 |
+
raw_value = "true"
|
| 89 |
+
try:
|
| 90 |
+
result[key] = _coerce_flag_value(key, raw_value.strip())
|
| 91 |
+
except (ValueError, TypeError) as e:
|
| 92 |
+
info = CLI_FEATURE_FLAG_REGISTRY.get(key, {})
|
| 93 |
+
logging.warning(
|
| 94 |
+
"Could not coerce --feature-flag %s=%r to %s (%s); dropping flag.",
|
| 95 |
+
key, raw_value.strip(), info.get("type", "?"), e,
|
| 96 |
+
)
|
| 97 |
+
return result
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
# Default server capabilities
|
| 101 |
+
_CORE_FEATURE_FLAGS: dict[str, Any] = {
|
| 102 |
+
"supports_preview_metadata": True,
|
| 103 |
+
"supports_model_type_tags": True,
|
| 104 |
+
"max_upload_size": args.max_upload_size * 1024 * 1024, # Convert MB to bytes
|
| 105 |
+
"extension": {"manager": {"supports_v4": True}},
|
| 106 |
+
"node_replacements": True,
|
| 107 |
+
"assets": args.enable_assets,
|
| 108 |
+
}
|
| 109 |
+
|
| 110 |
+
# CLI-provided flags cannot overwrite core flags
|
| 111 |
+
_cli_flags = {k: v for k, v in _parse_cli_feature_flags().items() if k not in _CORE_FEATURE_FLAGS}
|
| 112 |
+
|
| 113 |
+
SERVER_FEATURE_FLAGS: dict[str, Any] = {**_CORE_FEATURE_FLAGS, **_cli_flags}
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
def get_connection_feature(
|
| 117 |
+
sockets_metadata: dict[str, dict[str, Any]],
|
| 118 |
+
sid: str,
|
| 119 |
+
feature_name: str,
|
| 120 |
+
default: Any = False
|
| 121 |
+
) -> Any:
|
| 122 |
+
"""
|
| 123 |
+
Get a feature flag value for a specific connection.
|
| 124 |
+
|
| 125 |
+
Args:
|
| 126 |
+
sockets_metadata: Dictionary of socket metadata
|
| 127 |
+
sid: Session ID of the connection
|
| 128 |
+
feature_name: Name of the feature to check
|
| 129 |
+
default: Default value if feature not found
|
| 130 |
+
|
| 131 |
+
Returns:
|
| 132 |
+
Feature value or default if not found
|
| 133 |
+
"""
|
| 134 |
+
if sid not in sockets_metadata:
|
| 135 |
+
return default
|
| 136 |
+
|
| 137 |
+
return sockets_metadata[sid].get("feature_flags", {}).get(feature_name, default)
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
def supports_feature(
|
| 141 |
+
sockets_metadata: dict[str, dict[str, Any]],
|
| 142 |
+
sid: str,
|
| 143 |
+
feature_name: str
|
| 144 |
+
) -> bool:
|
| 145 |
+
"""
|
| 146 |
+
Check if a connection supports a specific feature.
|
| 147 |
+
|
| 148 |
+
Args:
|
| 149 |
+
sockets_metadata: Dictionary of socket metadata
|
| 150 |
+
sid: Session ID of the connection
|
| 151 |
+
feature_name: Name of the feature to check
|
| 152 |
+
|
| 153 |
+
Returns:
|
| 154 |
+
Boolean indicating if feature is supported
|
| 155 |
+
"""
|
| 156 |
+
return get_connection_feature(sockets_metadata, sid, feature_name, False) is True
|
| 157 |
+
|
| 158 |
+
|
| 159 |
+
def get_server_features() -> dict[str, Any]:
|
| 160 |
+
"""
|
| 161 |
+
Get the server's feature flags.
|
| 162 |
+
|
| 163 |
+
Returns:
|
| 164 |
+
Dictionary of server feature flags
|
| 165 |
+
"""
|
| 166 |
+
return SERVER_FEATURE_FLAGS.copy()
|
comfy_api/generate_api_stubs.py
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""
|
| 3 |
+
Script to generate .pyi stub files for the synchronous API wrappers.
|
| 4 |
+
This allows generating stubs without running the full ComfyUI application.
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
import os
|
| 8 |
+
import sys
|
| 9 |
+
import logging
|
| 10 |
+
import importlib
|
| 11 |
+
|
| 12 |
+
# Add ComfyUI to path so we can import modules
|
| 13 |
+
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
| 14 |
+
|
| 15 |
+
from comfy_api.internal.async_to_sync import AsyncToSyncConverter
|
| 16 |
+
from comfy_api.version_list import supported_versions
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def generate_stubs_for_module(module_name: str) -> None:
|
| 20 |
+
"""Generate stub files for a specific module that exports ComfyAPI and ComfyAPISync."""
|
| 21 |
+
try:
|
| 22 |
+
# Import the module
|
| 23 |
+
module = importlib.import_module(module_name)
|
| 24 |
+
|
| 25 |
+
# Check if module has ComfyAPISync (the sync wrapper)
|
| 26 |
+
if hasattr(module, "ComfyAPISync"):
|
| 27 |
+
# Module already has a sync class
|
| 28 |
+
api_class = getattr(module, "ComfyAPI", None)
|
| 29 |
+
sync_class = getattr(module, "ComfyAPISync")
|
| 30 |
+
|
| 31 |
+
if api_class:
|
| 32 |
+
# Generate the stub file
|
| 33 |
+
AsyncToSyncConverter.generate_stub_file(api_class, sync_class)
|
| 34 |
+
logging.info(f"Generated stub file for {module_name}")
|
| 35 |
+
else:
|
| 36 |
+
logging.warning(
|
| 37 |
+
f"Module {module_name} has ComfyAPISync but no ComfyAPI"
|
| 38 |
+
)
|
| 39 |
+
|
| 40 |
+
elif hasattr(module, "ComfyAPI"):
|
| 41 |
+
# Module only has async API, need to create sync wrapper first
|
| 42 |
+
from comfy_api.internal.async_to_sync import create_sync_class
|
| 43 |
+
|
| 44 |
+
api_class = getattr(module, "ComfyAPI")
|
| 45 |
+
sync_class = create_sync_class(api_class)
|
| 46 |
+
|
| 47 |
+
# Generate the stub file
|
| 48 |
+
AsyncToSyncConverter.generate_stub_file(api_class, sync_class)
|
| 49 |
+
logging.info(f"Generated stub file for {module_name}")
|
| 50 |
+
else:
|
| 51 |
+
logging.warning(
|
| 52 |
+
f"Module {module_name} does not export ComfyAPI or ComfyAPISync"
|
| 53 |
+
)
|
| 54 |
+
|
| 55 |
+
except Exception as e:
|
| 56 |
+
logging.error(f"Failed to generate stub for {module_name}: {e}")
|
| 57 |
+
import traceback
|
| 58 |
+
|
| 59 |
+
traceback.print_exc()
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def main():
|
| 63 |
+
"""Main function to generate all API stub files."""
|
| 64 |
+
logging.basicConfig(level=logging.INFO)
|
| 65 |
+
|
| 66 |
+
logging.info("Starting stub generation...")
|
| 67 |
+
|
| 68 |
+
# Dynamically get module names from supported_versions
|
| 69 |
+
api_modules = []
|
| 70 |
+
for api_class in supported_versions:
|
| 71 |
+
# Extract module name from the class
|
| 72 |
+
module_name = api_class.__module__
|
| 73 |
+
if module_name not in api_modules:
|
| 74 |
+
api_modules.append(module_name)
|
| 75 |
+
|
| 76 |
+
logging.info(f"Found {len(api_modules)} API modules: {api_modules}")
|
| 77 |
+
|
| 78 |
+
# Generate stubs for each module
|
| 79 |
+
for module_name in api_modules:
|
| 80 |
+
generate_stubs_for_module(module_name)
|
| 81 |
+
|
| 82 |
+
logging.info("Stub generation complete!")
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
if __name__ == "__main__":
|
| 86 |
+
main()
|
comfy_api/input/__init__.py
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# This file only exists for backwards compatibility.
|
| 2 |
+
from comfy_api.latest._input import (
|
| 3 |
+
ImageInput,
|
| 4 |
+
AudioInput,
|
| 5 |
+
MaskInput,
|
| 6 |
+
LatentInput,
|
| 7 |
+
VideoInput,
|
| 8 |
+
CurvePoint,
|
| 9 |
+
CurveInput,
|
| 10 |
+
MonotoneCubicCurve,
|
| 11 |
+
LinearCurve,
|
| 12 |
+
RangeInput,
|
| 13 |
+
)
|
| 14 |
+
|
| 15 |
+
__all__ = [
|
| 16 |
+
"ImageInput",
|
| 17 |
+
"AudioInput",
|
| 18 |
+
"MaskInput",
|
| 19 |
+
"LatentInput",
|
| 20 |
+
"VideoInput",
|
| 21 |
+
"CurvePoint",
|
| 22 |
+
"CurveInput",
|
| 23 |
+
"MonotoneCubicCurve",
|
| 24 |
+
"LinearCurve",
|
| 25 |
+
"RangeInput",
|
| 26 |
+
]
|
comfy_api/input/basic_types.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# This file only exists for backwards compatibility.
|
| 2 |
+
from comfy_api.latest._input.basic_types import (
|
| 3 |
+
ImageInput,
|
| 4 |
+
AudioInput,
|
| 5 |
+
MaskInput,
|
| 6 |
+
LatentInput,
|
| 7 |
+
)
|
| 8 |
+
|
| 9 |
+
__all__ = [
|
| 10 |
+
"ImageInput",
|
| 11 |
+
"AudioInput",
|
| 12 |
+
"MaskInput",
|
| 13 |
+
"LatentInput",
|
| 14 |
+
]
|
comfy_api/input/video_types.py
ADDED
|
@@ -0,0 +1,6 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# This file only exists for backwards compatibility.
|
| 2 |
+
from comfy_api.latest._input.video_types import VideoInput
|
| 3 |
+
|
| 4 |
+
__all__ = [
|
| 5 |
+
"VideoInput",
|
| 6 |
+
]
|
comfy_api/input_impl/__init__.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# This file only exists for backwards compatibility.
|
| 2 |
+
from comfy_api.latest._input_impl import VideoFromFile, VideoFromComponents
|
| 3 |
+
|
| 4 |
+
__all__ = [
|
| 5 |
+
"VideoFromFile",
|
| 6 |
+
"VideoFromComponents",
|
| 7 |
+
]
|
comfy_api/input_impl/video_types.py
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# This file only exists for backwards compatibility.
|
| 2 |
+
from comfy_api.latest._input_impl.video_types import * # noqa: F403
|
comfy_api/internal/__init__.py
ADDED
|
@@ -0,0 +1,150 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Internal infrastructure for ComfyAPI
|
| 2 |
+
from .api_registry import (
|
| 3 |
+
ComfyAPIBase as ComfyAPIBase,
|
| 4 |
+
ComfyAPIWithVersion as ComfyAPIWithVersion,
|
| 5 |
+
register_versions as register_versions,
|
| 6 |
+
get_all_versions as get_all_versions,
|
| 7 |
+
)
|
| 8 |
+
|
| 9 |
+
import asyncio
|
| 10 |
+
from dataclasses import asdict
|
| 11 |
+
from typing import Callable, Optional
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def first_real_override(cls: type, name: str, *, base: type=None) -> Optional[Callable]:
|
| 15 |
+
"""Return the *callable* override of `name` visible on `cls`, or None if every
|
| 16 |
+
implementation up to (and including) `base` is the placeholder defined on `base`.
|
| 17 |
+
|
| 18 |
+
If base is not provided, it will assume cls has a GET_BASE_CLASS
|
| 19 |
+
"""
|
| 20 |
+
if base is None:
|
| 21 |
+
if not hasattr(cls, "GET_BASE_CLASS"):
|
| 22 |
+
raise ValueError("base is required if cls does not have a GET_BASE_CLASS; is this a valid ComfyNode subclass?")
|
| 23 |
+
base = cls.GET_BASE_CLASS()
|
| 24 |
+
base_attr = getattr(base, name, None)
|
| 25 |
+
if base_attr is None:
|
| 26 |
+
return None
|
| 27 |
+
base_func = base_attr.__func__
|
| 28 |
+
for c in cls.mro(): # NodeB, NodeA, ComfyNode, object …
|
| 29 |
+
if c is base: # reached the placeholder – we're done
|
| 30 |
+
break
|
| 31 |
+
if name in c.__dict__: # first class that *defines* the attr
|
| 32 |
+
func = getattr(c, name).__func__
|
| 33 |
+
if func is not base_func: # real override
|
| 34 |
+
return getattr(cls, name) # bound to *cls*
|
| 35 |
+
return None
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
class _ComfyNodeInternal:
|
| 39 |
+
"""Class that all V3-based APIs inherit from for ComfyNode.
|
| 40 |
+
|
| 41 |
+
This is intended to only be referenced within execution.py, as it has to handle all V3 APIs going forward."""
|
| 42 |
+
@classmethod
|
| 43 |
+
def GET_NODE_INFO_V1(cls):
|
| 44 |
+
...
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
class _NodeOutputInternal:
|
| 48 |
+
"""Class that all V3-based APIs inherit from for NodeOutput.
|
| 49 |
+
|
| 50 |
+
This is intended to only be referenced within execution.py, as it has to handle all V3 APIs going forward."""
|
| 51 |
+
...
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def as_pruned_dict(dataclass_obj):
|
| 55 |
+
'''Return dict of dataclass object with pruned None values.'''
|
| 56 |
+
return prune_dict(asdict(dataclass_obj))
|
| 57 |
+
|
| 58 |
+
def prune_dict(d: dict):
|
| 59 |
+
return {k: v for k,v in d.items() if v is not None}
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def is_class(obj):
|
| 63 |
+
'''
|
| 64 |
+
Returns True if is a class type.
|
| 65 |
+
Returns False if is a class instance.
|
| 66 |
+
'''
|
| 67 |
+
return isinstance(obj, type)
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def copy_class(cls: type) -> type:
|
| 71 |
+
'''
|
| 72 |
+
Copy a class and its attributes.
|
| 73 |
+
'''
|
| 74 |
+
if cls is None:
|
| 75 |
+
return None
|
| 76 |
+
cls_dict = {
|
| 77 |
+
k: v for k, v in cls.__dict__.items()
|
| 78 |
+
if k not in ('__dict__', '__weakref__', '__module__', '__doc__')
|
| 79 |
+
}
|
| 80 |
+
# new class
|
| 81 |
+
new_cls = type(
|
| 82 |
+
cls.__name__,
|
| 83 |
+
(cls,),
|
| 84 |
+
cls_dict
|
| 85 |
+
)
|
| 86 |
+
# metadata preservation
|
| 87 |
+
new_cls.__module__ = cls.__module__
|
| 88 |
+
new_cls.__doc__ = cls.__doc__
|
| 89 |
+
return new_cls
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
class classproperty(object):
|
| 93 |
+
def __init__(self, f):
|
| 94 |
+
self.f = f
|
| 95 |
+
def __get__(self, obj, owner):
|
| 96 |
+
return self.f(owner)
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
# NOTE: this was ai generated and validated by hand
|
| 100 |
+
def shallow_clone_class(cls, new_name=None):
|
| 101 |
+
'''
|
| 102 |
+
Shallow clone a class while preserving super() functionality.
|
| 103 |
+
'''
|
| 104 |
+
new_name = new_name or f"{cls.__name__}Clone"
|
| 105 |
+
# Include the original class in the bases to maintain proper inheritance
|
| 106 |
+
new_bases = (cls,) + cls.__bases__
|
| 107 |
+
return type(new_name, new_bases, dict(cls.__dict__))
|
| 108 |
+
|
| 109 |
+
# NOTE: this was ai generated and validated by hand
|
| 110 |
+
def lock_class(cls):
|
| 111 |
+
'''
|
| 112 |
+
Lock a class so that its top-levelattributes cannot be modified.
|
| 113 |
+
'''
|
| 114 |
+
# Locked instance __setattr__
|
| 115 |
+
def locked_instance_setattr(self, name, value):
|
| 116 |
+
raise AttributeError(
|
| 117 |
+
f"Cannot set attribute '{name}' on immutable instance of {type(self).__name__}"
|
| 118 |
+
)
|
| 119 |
+
# Locked metaclass
|
| 120 |
+
class LockedMeta(type(cls)):
|
| 121 |
+
def __setattr__(cls_, name, value):
|
| 122 |
+
raise AttributeError(
|
| 123 |
+
f"Cannot modify class attribute '{name}' on locked class '{cls_.__name__}'"
|
| 124 |
+
)
|
| 125 |
+
# Rebuild class with locked behavior
|
| 126 |
+
locked_dict = dict(cls.__dict__)
|
| 127 |
+
locked_dict['__setattr__'] = locked_instance_setattr
|
| 128 |
+
|
| 129 |
+
return LockedMeta(cls.__name__, cls.__bases__, locked_dict)
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
def make_locked_method_func(type_obj, func, class_clone):
|
| 133 |
+
"""
|
| 134 |
+
Returns a function that, when called with **inputs, will execute:
|
| 135 |
+
getattr(type_obj, func).__func__(lock_class(class_clone), **inputs)
|
| 136 |
+
|
| 137 |
+
Supports both synchronous and asynchronous methods.
|
| 138 |
+
"""
|
| 139 |
+
locked_class = lock_class(class_clone)
|
| 140 |
+
method = getattr(type_obj, func).__func__
|
| 141 |
+
|
| 142 |
+
# Check if the original method is async
|
| 143 |
+
if asyncio.iscoroutinefunction(method):
|
| 144 |
+
async def wrapped_async_func(**inputs):
|
| 145 |
+
return await method(locked_class, **inputs)
|
| 146 |
+
return wrapped_async_func
|
| 147 |
+
else:
|
| 148 |
+
def wrapped_func(**inputs):
|
| 149 |
+
return method(locked_class, **inputs)
|
| 150 |
+
return wrapped_func
|
comfy_api/internal/api_registry.py
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import NamedTuple
|
| 2 |
+
from comfy_api.internal.singleton import ProxiedSingleton
|
| 3 |
+
from packaging import version as packaging_version
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class ComfyAPIBase(ProxiedSingleton):
|
| 7 |
+
def __init__(self):
|
| 8 |
+
pass
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class ComfyAPIWithVersion(NamedTuple):
|
| 12 |
+
version: str
|
| 13 |
+
api_class: type[ComfyAPIBase]
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def parse_version(version_str: str) -> packaging_version.Version:
|
| 17 |
+
"""
|
| 18 |
+
Parses a version string into a packaging_version.Version object.
|
| 19 |
+
Raises ValueError if the version string is invalid.
|
| 20 |
+
"""
|
| 21 |
+
if version_str == "latest":
|
| 22 |
+
return packaging_version.parse("9999999.9999999.9999999")
|
| 23 |
+
return packaging_version.parse(version_str)
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
registered_versions: list[ComfyAPIWithVersion] = []
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def register_versions(versions: list[ComfyAPIWithVersion]):
|
| 30 |
+
versions.sort(key=lambda x: parse_version(x.version))
|
| 31 |
+
global registered_versions
|
| 32 |
+
registered_versions = versions
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def get_all_versions() -> list[ComfyAPIWithVersion]:
|
| 36 |
+
"""
|
| 37 |
+
Returns a list of all registered ComfyAPI versions.
|
| 38 |
+
"""
|
| 39 |
+
return registered_versions
|
comfy_api/internal/async_to_sync.py
ADDED
|
@@ -0,0 +1,1002 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import asyncio
|
| 2 |
+
import concurrent.futures
|
| 3 |
+
import contextvars
|
| 4 |
+
import functools
|
| 5 |
+
import inspect
|
| 6 |
+
import logging
|
| 7 |
+
import os
|
| 8 |
+
import textwrap
|
| 9 |
+
import threading
|
| 10 |
+
from enum import Enum
|
| 11 |
+
from typing import Optional, get_origin, get_args, get_type_hints
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class TypeTracker:
|
| 15 |
+
"""Tracks types discovered during stub generation for automatic import generation."""
|
| 16 |
+
|
| 17 |
+
def __init__(self):
|
| 18 |
+
self.discovered_types = {} # type_name -> (module, qualname)
|
| 19 |
+
self.builtin_types = {
|
| 20 |
+
"Any",
|
| 21 |
+
"Dict",
|
| 22 |
+
"List",
|
| 23 |
+
"Optional",
|
| 24 |
+
"Tuple",
|
| 25 |
+
"Union",
|
| 26 |
+
"Set",
|
| 27 |
+
"Sequence",
|
| 28 |
+
"cast",
|
| 29 |
+
"NamedTuple",
|
| 30 |
+
"str",
|
| 31 |
+
"int",
|
| 32 |
+
"float",
|
| 33 |
+
"bool",
|
| 34 |
+
"None",
|
| 35 |
+
"bytes",
|
| 36 |
+
"object",
|
| 37 |
+
"type",
|
| 38 |
+
"dict",
|
| 39 |
+
"list",
|
| 40 |
+
"tuple",
|
| 41 |
+
"set",
|
| 42 |
+
}
|
| 43 |
+
self.already_imported = (
|
| 44 |
+
set()
|
| 45 |
+
) # Track types already imported to avoid duplicates
|
| 46 |
+
|
| 47 |
+
def track_type(self, annotation):
|
| 48 |
+
"""Track a type annotation and record its module/import info."""
|
| 49 |
+
if annotation is None or annotation is type(None):
|
| 50 |
+
return
|
| 51 |
+
|
| 52 |
+
# Skip builtins and typing module types we already import
|
| 53 |
+
type_name = getattr(annotation, "__name__", None)
|
| 54 |
+
if type_name and (
|
| 55 |
+
type_name in self.builtin_types or type_name in self.already_imported
|
| 56 |
+
):
|
| 57 |
+
return
|
| 58 |
+
|
| 59 |
+
# Get module and qualname
|
| 60 |
+
module = getattr(annotation, "__module__", None)
|
| 61 |
+
qualname = getattr(annotation, "__qualname__", type_name or "")
|
| 62 |
+
|
| 63 |
+
# Skip types from typing module (they're already imported)
|
| 64 |
+
if module == "typing":
|
| 65 |
+
return
|
| 66 |
+
|
| 67 |
+
# Skip UnionType and GenericAlias from types module as they're handled specially
|
| 68 |
+
if module == "types" and type_name in ("UnionType", "GenericAlias"):
|
| 69 |
+
return
|
| 70 |
+
|
| 71 |
+
if module and module not in ["builtins", "__main__"]:
|
| 72 |
+
# Store the type info
|
| 73 |
+
if type_name:
|
| 74 |
+
self.discovered_types[type_name] = (module, qualname)
|
| 75 |
+
|
| 76 |
+
def get_imports(self, main_module_name: str) -> list[str]:
|
| 77 |
+
"""Generate import statements for all discovered types."""
|
| 78 |
+
imports = []
|
| 79 |
+
imports_by_module = {}
|
| 80 |
+
|
| 81 |
+
for type_name, (module, qualname) in sorted(self.discovered_types.items()):
|
| 82 |
+
# Skip types from the main module (they're already imported)
|
| 83 |
+
if main_module_name and module == main_module_name:
|
| 84 |
+
continue
|
| 85 |
+
|
| 86 |
+
if module not in imports_by_module:
|
| 87 |
+
imports_by_module[module] = []
|
| 88 |
+
if type_name not in imports_by_module[module]: # Avoid duplicates
|
| 89 |
+
imports_by_module[module].append(type_name)
|
| 90 |
+
|
| 91 |
+
# Generate import statements
|
| 92 |
+
for module, types in sorted(imports_by_module.items()):
|
| 93 |
+
if len(types) == 1:
|
| 94 |
+
imports.append(f"from {module} import {types[0]}")
|
| 95 |
+
else:
|
| 96 |
+
imports.append(f"from {module} import {', '.join(sorted(set(types)))}")
|
| 97 |
+
|
| 98 |
+
return imports
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
class AsyncToSyncConverter:
|
| 102 |
+
"""
|
| 103 |
+
Provides utilities to convert async classes to sync classes with proper type hints.
|
| 104 |
+
"""
|
| 105 |
+
|
| 106 |
+
_thread_pool: Optional[concurrent.futures.ThreadPoolExecutor] = None
|
| 107 |
+
_thread_pool_lock = threading.Lock()
|
| 108 |
+
_thread_pool_initialized = False
|
| 109 |
+
|
| 110 |
+
@classmethod
|
| 111 |
+
def get_thread_pool(cls, max_workers=None) -> concurrent.futures.ThreadPoolExecutor:
|
| 112 |
+
"""Get or create the shared thread pool with proper thread-safe initialization."""
|
| 113 |
+
# Fast path - check if already initialized without acquiring lock
|
| 114 |
+
if cls._thread_pool_initialized:
|
| 115 |
+
assert cls._thread_pool is not None, "Thread pool should be initialized"
|
| 116 |
+
return cls._thread_pool
|
| 117 |
+
|
| 118 |
+
# Slow path - acquire lock and create pool if needed
|
| 119 |
+
with cls._thread_pool_lock:
|
| 120 |
+
if not cls._thread_pool_initialized:
|
| 121 |
+
cls._thread_pool = concurrent.futures.ThreadPoolExecutor(
|
| 122 |
+
max_workers=max_workers, thread_name_prefix="async_to_sync_"
|
| 123 |
+
)
|
| 124 |
+
cls._thread_pool_initialized = True
|
| 125 |
+
|
| 126 |
+
# This should never be None at this point, but add assertion for type checker
|
| 127 |
+
assert cls._thread_pool is not None
|
| 128 |
+
return cls._thread_pool
|
| 129 |
+
|
| 130 |
+
@classmethod
|
| 131 |
+
def run_async_in_thread(cls, coro_func, *args, **kwargs):
|
| 132 |
+
"""
|
| 133 |
+
Run an async function in a separate thread from the thread pool.
|
| 134 |
+
Blocks until the async function completes.
|
| 135 |
+
Properly propagates contextvars between threads and manages event loops.
|
| 136 |
+
"""
|
| 137 |
+
# Capture current context - this includes all context variables
|
| 138 |
+
context = contextvars.copy_context()
|
| 139 |
+
|
| 140 |
+
# Store the result and any exception that occurs
|
| 141 |
+
result_container: dict = {"result": None, "exception": None}
|
| 142 |
+
|
| 143 |
+
# Function that runs in the thread pool
|
| 144 |
+
def run_in_thread():
|
| 145 |
+
# Create new event loop for this thread
|
| 146 |
+
loop = asyncio.new_event_loop()
|
| 147 |
+
asyncio.set_event_loop(loop)
|
| 148 |
+
|
| 149 |
+
try:
|
| 150 |
+
# Create the coroutine within the context
|
| 151 |
+
async def run_with_context():
|
| 152 |
+
# The coroutine function might access context variables
|
| 153 |
+
return await coro_func(*args, **kwargs)
|
| 154 |
+
|
| 155 |
+
# Run the coroutine with the captured context
|
| 156 |
+
# This ensures all context variables are available in the async function
|
| 157 |
+
result = context.run(loop.run_until_complete, run_with_context())
|
| 158 |
+
result_container["result"] = result
|
| 159 |
+
except Exception as e:
|
| 160 |
+
# Store the exception to re-raise in the calling thread
|
| 161 |
+
result_container["exception"] = e
|
| 162 |
+
finally:
|
| 163 |
+
# Ensure event loop is properly closed to prevent warnings
|
| 164 |
+
try:
|
| 165 |
+
# Cancel any remaining tasks
|
| 166 |
+
pending = asyncio.all_tasks(loop)
|
| 167 |
+
for task in pending:
|
| 168 |
+
task.cancel()
|
| 169 |
+
|
| 170 |
+
# Run the loop briefly to handle cancellations
|
| 171 |
+
if pending:
|
| 172 |
+
loop.run_until_complete(
|
| 173 |
+
asyncio.gather(*pending, return_exceptions=True)
|
| 174 |
+
)
|
| 175 |
+
except Exception:
|
| 176 |
+
pass # Ignore errors during cleanup
|
| 177 |
+
|
| 178 |
+
# Close the event loop
|
| 179 |
+
loop.close()
|
| 180 |
+
|
| 181 |
+
# Clear the event loop from the thread
|
| 182 |
+
asyncio.set_event_loop(None)
|
| 183 |
+
|
| 184 |
+
# Submit to thread pool and wait for result
|
| 185 |
+
thread_pool = cls.get_thread_pool()
|
| 186 |
+
future = thread_pool.submit(run_in_thread)
|
| 187 |
+
future.result() # Wait for completion
|
| 188 |
+
|
| 189 |
+
# Re-raise any exception that occurred in the thread
|
| 190 |
+
if result_container["exception"] is not None:
|
| 191 |
+
raise result_container["exception"]
|
| 192 |
+
|
| 193 |
+
return result_container["result"]
|
| 194 |
+
|
| 195 |
+
@classmethod
|
| 196 |
+
def create_sync_class(cls, async_class: type, thread_pool_size=10) -> type:
|
| 197 |
+
"""
|
| 198 |
+
Creates a new class with synchronous versions of all async methods.
|
| 199 |
+
|
| 200 |
+
Args:
|
| 201 |
+
async_class: The async class to convert
|
| 202 |
+
thread_pool_size: Size of thread pool to use
|
| 203 |
+
|
| 204 |
+
Returns:
|
| 205 |
+
A new class with sync versions of all async methods
|
| 206 |
+
"""
|
| 207 |
+
sync_class_name = "ComfyAPISyncStub"
|
| 208 |
+
cls.get_thread_pool(thread_pool_size)
|
| 209 |
+
|
| 210 |
+
# Create a proper class with docstrings and proper base classes
|
| 211 |
+
sync_class_dict = {
|
| 212 |
+
"__doc__": async_class.__doc__,
|
| 213 |
+
"__module__": async_class.__module__,
|
| 214 |
+
"__qualname__": sync_class_name,
|
| 215 |
+
"__orig_class__": async_class, # Store original class for typing references
|
| 216 |
+
}
|
| 217 |
+
|
| 218 |
+
# Create __init__ method
|
| 219 |
+
def __init__(self, *args, **kwargs):
|
| 220 |
+
self._async_instance = async_class(*args, **kwargs)
|
| 221 |
+
|
| 222 |
+
# Handle annotated class attributes (like execution: Execution)
|
| 223 |
+
# Get all annotations from the class hierarchy and resolve string annotations
|
| 224 |
+
try:
|
| 225 |
+
# get_type_hints resolves string annotations to actual type objects
|
| 226 |
+
# This handles classes using 'from __future__ import annotations'
|
| 227 |
+
all_annotations = get_type_hints(async_class)
|
| 228 |
+
except Exception:
|
| 229 |
+
# Fallback to raw annotations if get_type_hints fails
|
| 230 |
+
# (e.g., for undefined forward references)
|
| 231 |
+
all_annotations = {}
|
| 232 |
+
for base_class in reversed(inspect.getmro(async_class)):
|
| 233 |
+
if hasattr(base_class, "__annotations__"):
|
| 234 |
+
all_annotations.update(base_class.__annotations__)
|
| 235 |
+
|
| 236 |
+
# For each annotated attribute, check if it needs to be created or wrapped
|
| 237 |
+
for attr_name, attr_type in all_annotations.items():
|
| 238 |
+
if hasattr(self._async_instance, attr_name):
|
| 239 |
+
# Attribute exists on the instance
|
| 240 |
+
attr = getattr(self._async_instance, attr_name)
|
| 241 |
+
# Check if this attribute needs a sync wrapper
|
| 242 |
+
if hasattr(attr, "__class__"):
|
| 243 |
+
from comfy_api.internal.singleton import ProxiedSingleton
|
| 244 |
+
|
| 245 |
+
if isinstance(attr, ProxiedSingleton):
|
| 246 |
+
# Create a sync version of this attribute
|
| 247 |
+
try:
|
| 248 |
+
sync_attr_class = cls.create_sync_class(attr.__class__)
|
| 249 |
+
# Create instance of the sync wrapper with the async instance
|
| 250 |
+
sync_attr = object.__new__(sync_attr_class) # type: ignore
|
| 251 |
+
sync_attr._async_instance = attr
|
| 252 |
+
setattr(self, attr_name, sync_attr)
|
| 253 |
+
except Exception:
|
| 254 |
+
# If we can't create a sync version, keep the original
|
| 255 |
+
setattr(self, attr_name, attr)
|
| 256 |
+
else:
|
| 257 |
+
# Not async, just copy the reference
|
| 258 |
+
setattr(self, attr_name, attr)
|
| 259 |
+
else:
|
| 260 |
+
# Attribute doesn't exist, but is annotated - create it
|
| 261 |
+
# This handles cases like execution: Execution
|
| 262 |
+
if isinstance(attr_type, type):
|
| 263 |
+
# Check if the type is defined as an inner class
|
| 264 |
+
if hasattr(async_class, attr_type.__name__):
|
| 265 |
+
inner_class = getattr(async_class, attr_type.__name__)
|
| 266 |
+
from comfy_api.internal.singleton import ProxiedSingleton
|
| 267 |
+
|
| 268 |
+
# Create an instance of the inner class
|
| 269 |
+
try:
|
| 270 |
+
# For ProxiedSingleton classes, get or create the singleton instance
|
| 271 |
+
if issubclass(inner_class, ProxiedSingleton):
|
| 272 |
+
async_instance = inner_class.get_instance()
|
| 273 |
+
else:
|
| 274 |
+
async_instance = inner_class()
|
| 275 |
+
|
| 276 |
+
# Create sync wrapper
|
| 277 |
+
sync_attr_class = cls.create_sync_class(inner_class)
|
| 278 |
+
sync_attr = object.__new__(sync_attr_class) # type: ignore
|
| 279 |
+
sync_attr._async_instance = async_instance
|
| 280 |
+
setattr(self, attr_name, sync_attr)
|
| 281 |
+
# Also set on the async instance for consistency
|
| 282 |
+
setattr(self._async_instance, attr_name, async_instance)
|
| 283 |
+
except Exception as e:
|
| 284 |
+
logging.warning(
|
| 285 |
+
f"Failed to create instance for {attr_name}: {e}"
|
| 286 |
+
)
|
| 287 |
+
|
| 288 |
+
# Handle other instance attributes that might not be annotated
|
| 289 |
+
for name, attr in inspect.getmembers(self._async_instance):
|
| 290 |
+
if name.startswith("_") or hasattr(self, name):
|
| 291 |
+
continue
|
| 292 |
+
|
| 293 |
+
# If attribute is an instance of a class, and that class is defined in the original class
|
| 294 |
+
# we need to check if it needs a sync wrapper
|
| 295 |
+
if isinstance(attr, object) and not isinstance(
|
| 296 |
+
attr, (str, int, float, bool, list, dict, tuple)
|
| 297 |
+
):
|
| 298 |
+
from comfy_api.internal.singleton import ProxiedSingleton
|
| 299 |
+
|
| 300 |
+
if isinstance(attr, ProxiedSingleton):
|
| 301 |
+
# Create a sync version of this nested class
|
| 302 |
+
try:
|
| 303 |
+
sync_attr_class = cls.create_sync_class(attr.__class__)
|
| 304 |
+
# Create instance of the sync wrapper with the async instance
|
| 305 |
+
sync_attr = object.__new__(sync_attr_class) # type: ignore
|
| 306 |
+
sync_attr._async_instance = attr
|
| 307 |
+
setattr(self, name, sync_attr)
|
| 308 |
+
except Exception:
|
| 309 |
+
# If we can't create a sync version, keep the original
|
| 310 |
+
setattr(self, name, attr)
|
| 311 |
+
|
| 312 |
+
sync_class_dict["__init__"] = __init__
|
| 313 |
+
|
| 314 |
+
# Process methods from the async class
|
| 315 |
+
for name, method in inspect.getmembers(
|
| 316 |
+
async_class, predicate=inspect.isfunction
|
| 317 |
+
):
|
| 318 |
+
if name.startswith("_"):
|
| 319 |
+
continue
|
| 320 |
+
|
| 321 |
+
# Extract the actual return type from a coroutine
|
| 322 |
+
if inspect.iscoroutinefunction(method):
|
| 323 |
+
# Create sync version of async method with proper signature
|
| 324 |
+
@functools.wraps(method)
|
| 325 |
+
def sync_method(self, *args, _method_name=name, **kwargs):
|
| 326 |
+
async_method = getattr(self._async_instance, _method_name)
|
| 327 |
+
return AsyncToSyncConverter.run_async_in_thread(
|
| 328 |
+
async_method, *args, **kwargs
|
| 329 |
+
)
|
| 330 |
+
|
| 331 |
+
# Add to the class dict
|
| 332 |
+
sync_class_dict[name] = sync_method
|
| 333 |
+
else:
|
| 334 |
+
# For regular methods, create a proxy method
|
| 335 |
+
@functools.wraps(method)
|
| 336 |
+
def proxy_method(self, *args, _method_name=name, **kwargs):
|
| 337 |
+
method = getattr(self._async_instance, _method_name)
|
| 338 |
+
return method(*args, **kwargs)
|
| 339 |
+
|
| 340 |
+
# Add to the class dict
|
| 341 |
+
sync_class_dict[name] = proxy_method
|
| 342 |
+
|
| 343 |
+
# Handle property access
|
| 344 |
+
for name, prop in inspect.getmembers(
|
| 345 |
+
async_class, lambda x: isinstance(x, property)
|
| 346 |
+
):
|
| 347 |
+
|
| 348 |
+
def make_property(name, prop_obj):
|
| 349 |
+
def getter(self):
|
| 350 |
+
value = getattr(self._async_instance, name)
|
| 351 |
+
if inspect.iscoroutinefunction(value):
|
| 352 |
+
|
| 353 |
+
def sync_fn(*args, **kwargs):
|
| 354 |
+
return AsyncToSyncConverter.run_async_in_thread(
|
| 355 |
+
value, *args, **kwargs
|
| 356 |
+
)
|
| 357 |
+
|
| 358 |
+
return sync_fn
|
| 359 |
+
return value
|
| 360 |
+
|
| 361 |
+
def setter(self, value):
|
| 362 |
+
setattr(self._async_instance, name, value)
|
| 363 |
+
|
| 364 |
+
return property(getter, setter if prop_obj.fset else None)
|
| 365 |
+
|
| 366 |
+
sync_class_dict[name] = make_property(name, prop)
|
| 367 |
+
|
| 368 |
+
# Create the class
|
| 369 |
+
sync_class = type(sync_class_name, (object,), sync_class_dict)
|
| 370 |
+
|
| 371 |
+
return sync_class
|
| 372 |
+
|
| 373 |
+
@classmethod
|
| 374 |
+
def _format_type_annotation(
|
| 375 |
+
cls, annotation, type_tracker: Optional[TypeTracker] = None
|
| 376 |
+
) -> str:
|
| 377 |
+
"""Convert a type annotation to its string representation for stub files."""
|
| 378 |
+
if (
|
| 379 |
+
annotation is inspect.Parameter.empty
|
| 380 |
+
or annotation is inspect.Signature.empty
|
| 381 |
+
):
|
| 382 |
+
return "Any"
|
| 383 |
+
|
| 384 |
+
# Handle None type
|
| 385 |
+
if annotation is type(None):
|
| 386 |
+
return "None"
|
| 387 |
+
|
| 388 |
+
# Track the type if we have a tracker
|
| 389 |
+
if type_tracker:
|
| 390 |
+
type_tracker.track_type(annotation)
|
| 391 |
+
|
| 392 |
+
# Try using typing.get_origin/get_args for Python 3.8+
|
| 393 |
+
try:
|
| 394 |
+
origin = get_origin(annotation)
|
| 395 |
+
args = get_args(annotation)
|
| 396 |
+
|
| 397 |
+
if origin is not None:
|
| 398 |
+
# Track the origin type
|
| 399 |
+
if type_tracker:
|
| 400 |
+
type_tracker.track_type(origin)
|
| 401 |
+
|
| 402 |
+
# Get the origin name
|
| 403 |
+
origin_name = getattr(origin, "__name__", str(origin))
|
| 404 |
+
if "." in origin_name:
|
| 405 |
+
origin_name = origin_name.split(".")[-1]
|
| 406 |
+
|
| 407 |
+
# Special handling for types.UnionType (Python 3.10+ pipe operator)
|
| 408 |
+
# Convert to old-style Union for compatibility
|
| 409 |
+
if str(origin) == "<class 'types.UnionType'>" or origin_name == "UnionType":
|
| 410 |
+
origin_name = "Union"
|
| 411 |
+
|
| 412 |
+
# Format arguments recursively
|
| 413 |
+
if args:
|
| 414 |
+
formatted_args = []
|
| 415 |
+
for arg in args:
|
| 416 |
+
# Track each type in the union
|
| 417 |
+
if type_tracker:
|
| 418 |
+
type_tracker.track_type(arg)
|
| 419 |
+
formatted_args.append(cls._format_type_annotation(arg, type_tracker))
|
| 420 |
+
return f"{origin_name}[{', '.join(formatted_args)}]"
|
| 421 |
+
else:
|
| 422 |
+
return origin_name
|
| 423 |
+
except (AttributeError, TypeError):
|
| 424 |
+
# Fallback for older Python versions or non-generic types
|
| 425 |
+
pass
|
| 426 |
+
|
| 427 |
+
# Handle generic types the old way for compatibility
|
| 428 |
+
if hasattr(annotation, "__origin__") and hasattr(annotation, "__args__"):
|
| 429 |
+
origin = annotation.__origin__
|
| 430 |
+
origin_name = (
|
| 431 |
+
origin.__name__
|
| 432 |
+
if hasattr(origin, "__name__")
|
| 433 |
+
else str(origin).split("'")[1]
|
| 434 |
+
)
|
| 435 |
+
|
| 436 |
+
# Format each type argument
|
| 437 |
+
args = []
|
| 438 |
+
for arg in annotation.__args__:
|
| 439 |
+
args.append(cls._format_type_annotation(arg, type_tracker))
|
| 440 |
+
|
| 441 |
+
return f"{origin_name}[{', '.join(args)}]"
|
| 442 |
+
|
| 443 |
+
# Handle regular types with __name__
|
| 444 |
+
if hasattr(annotation, "__name__"):
|
| 445 |
+
return annotation.__name__
|
| 446 |
+
|
| 447 |
+
# Handle special module types (like types from typing module)
|
| 448 |
+
if hasattr(annotation, "__module__") and hasattr(annotation, "__qualname__"):
|
| 449 |
+
# For types like typing.Literal, typing.TypedDict, etc.
|
| 450 |
+
return annotation.__qualname__
|
| 451 |
+
|
| 452 |
+
# Last resort: string conversion with cleanup
|
| 453 |
+
type_str = str(annotation)
|
| 454 |
+
|
| 455 |
+
# Clean up common patterns more robustly
|
| 456 |
+
if type_str.startswith("<class '") and type_str.endswith("'>"):
|
| 457 |
+
type_str = type_str[8:-2] # Remove "<class '" and "'>"
|
| 458 |
+
|
| 459 |
+
# Remove module prefixes for common modules
|
| 460 |
+
for prefix in ["typing.", "builtins.", "types."]:
|
| 461 |
+
if type_str.startswith(prefix):
|
| 462 |
+
type_str = type_str[len(prefix) :]
|
| 463 |
+
|
| 464 |
+
# Handle special cases
|
| 465 |
+
if type_str in ("_empty", "inspect._empty"):
|
| 466 |
+
return "None"
|
| 467 |
+
|
| 468 |
+
# Fix NoneType (this should rarely be needed now)
|
| 469 |
+
if type_str == "NoneType":
|
| 470 |
+
return "None"
|
| 471 |
+
|
| 472 |
+
return type_str
|
| 473 |
+
|
| 474 |
+
@classmethod
|
| 475 |
+
def _extract_coroutine_return_type(cls, annotation):
|
| 476 |
+
"""Extract the actual return type from a Coroutine annotation."""
|
| 477 |
+
if hasattr(annotation, "__args__") and len(annotation.__args__) > 2:
|
| 478 |
+
# Coroutine[Any, Any, ReturnType] -> extract ReturnType
|
| 479 |
+
return annotation.__args__[2]
|
| 480 |
+
return annotation
|
| 481 |
+
|
| 482 |
+
@classmethod
|
| 483 |
+
def _format_parameter_default(cls, default_value) -> str:
|
| 484 |
+
"""Format a parameter's default value for stub files."""
|
| 485 |
+
if default_value is inspect.Parameter.empty:
|
| 486 |
+
return ""
|
| 487 |
+
elif default_value is None:
|
| 488 |
+
return " = None"
|
| 489 |
+
elif isinstance(default_value, bool):
|
| 490 |
+
return f" = {default_value}"
|
| 491 |
+
elif default_value == {}:
|
| 492 |
+
return " = {}"
|
| 493 |
+
elif default_value == []:
|
| 494 |
+
return " = []"
|
| 495 |
+
else:
|
| 496 |
+
return f" = {default_value}"
|
| 497 |
+
|
| 498 |
+
@classmethod
|
| 499 |
+
def _format_method_parameters(
|
| 500 |
+
cls,
|
| 501 |
+
sig: inspect.Signature,
|
| 502 |
+
skip_self: bool = True,
|
| 503 |
+
type_hints: Optional[dict] = None,
|
| 504 |
+
type_tracker: Optional[TypeTracker] = None,
|
| 505 |
+
) -> str:
|
| 506 |
+
"""Format method parameters for stub files."""
|
| 507 |
+
params = []
|
| 508 |
+
if type_hints is None:
|
| 509 |
+
type_hints = {}
|
| 510 |
+
|
| 511 |
+
for i, (param_name, param) in enumerate(sig.parameters.items()):
|
| 512 |
+
if i == 0 and param_name == "self" and skip_self:
|
| 513 |
+
params.append("self")
|
| 514 |
+
else:
|
| 515 |
+
# Get type annotation from type hints if available, otherwise from signature
|
| 516 |
+
annotation = type_hints.get(param_name, param.annotation)
|
| 517 |
+
type_str = cls._format_type_annotation(annotation, type_tracker)
|
| 518 |
+
|
| 519 |
+
# Get default value
|
| 520 |
+
default_str = cls._format_parameter_default(param.default)
|
| 521 |
+
|
| 522 |
+
# Combine parameter parts
|
| 523 |
+
if annotation is inspect.Parameter.empty:
|
| 524 |
+
params.append(f"{param_name}: Any{default_str}")
|
| 525 |
+
else:
|
| 526 |
+
params.append(f"{param_name}: {type_str}{default_str}")
|
| 527 |
+
|
| 528 |
+
return ", ".join(params)
|
| 529 |
+
|
| 530 |
+
@classmethod
|
| 531 |
+
def _generate_method_signature(
|
| 532 |
+
cls,
|
| 533 |
+
method_name: str,
|
| 534 |
+
method,
|
| 535 |
+
is_async: bool = False,
|
| 536 |
+
type_tracker: Optional[TypeTracker] = None,
|
| 537 |
+
) -> str:
|
| 538 |
+
"""Generate a complete method signature for stub files."""
|
| 539 |
+
sig = inspect.signature(method)
|
| 540 |
+
|
| 541 |
+
# Try to get evaluated type hints to resolve string annotations
|
| 542 |
+
try:
|
| 543 |
+
from typing import get_type_hints
|
| 544 |
+
type_hints = get_type_hints(method)
|
| 545 |
+
except Exception:
|
| 546 |
+
# Fallback to empty dict if we can't get type hints
|
| 547 |
+
type_hints = {}
|
| 548 |
+
|
| 549 |
+
# For async methods, extract the actual return type
|
| 550 |
+
return_annotation = type_hints.get('return', sig.return_annotation)
|
| 551 |
+
if is_async and inspect.iscoroutinefunction(method):
|
| 552 |
+
return_annotation = cls._extract_coroutine_return_type(return_annotation)
|
| 553 |
+
|
| 554 |
+
# Format parameters with type hints
|
| 555 |
+
params_str = cls._format_method_parameters(sig, type_hints=type_hints, type_tracker=type_tracker)
|
| 556 |
+
|
| 557 |
+
# Format return type
|
| 558 |
+
return_type = cls._format_type_annotation(return_annotation, type_tracker)
|
| 559 |
+
if return_annotation is inspect.Signature.empty:
|
| 560 |
+
return_type = "None"
|
| 561 |
+
|
| 562 |
+
return f"def {method_name}({params_str}) -> {return_type}: ..."
|
| 563 |
+
|
| 564 |
+
@classmethod
|
| 565 |
+
def _generate_imports(
|
| 566 |
+
cls, async_class: type, type_tracker: TypeTracker
|
| 567 |
+
) -> list[str]:
|
| 568 |
+
"""Generate import statements for the stub file."""
|
| 569 |
+
imports = []
|
| 570 |
+
|
| 571 |
+
# Add standard typing imports
|
| 572 |
+
imports.append(
|
| 573 |
+
"from typing import Any, Dict, List, Optional, Tuple, Union, Set, Sequence, cast, NamedTuple"
|
| 574 |
+
)
|
| 575 |
+
|
| 576 |
+
# Add imports from the original module
|
| 577 |
+
if async_class.__module__ != "builtins":
|
| 578 |
+
module = inspect.getmodule(async_class)
|
| 579 |
+
additional_types = []
|
| 580 |
+
|
| 581 |
+
if module:
|
| 582 |
+
# Check if module has __all__ defined
|
| 583 |
+
module_all = getattr(module, "__all__", None)
|
| 584 |
+
|
| 585 |
+
for name, obj in sorted(inspect.getmembers(module)):
|
| 586 |
+
if isinstance(obj, type):
|
| 587 |
+
# Skip if __all__ is defined and this name isn't in it
|
| 588 |
+
# unless it's already been tracked as used in type annotations
|
| 589 |
+
if module_all is not None and name not in module_all:
|
| 590 |
+
# Check if this type was actually used in annotations
|
| 591 |
+
if name not in type_tracker.discovered_types:
|
| 592 |
+
continue
|
| 593 |
+
|
| 594 |
+
# Check for NamedTuple
|
| 595 |
+
if issubclass(obj, tuple) and hasattr(obj, "_fields"):
|
| 596 |
+
additional_types.append(name)
|
| 597 |
+
# Mark as already imported
|
| 598 |
+
type_tracker.already_imported.add(name)
|
| 599 |
+
# Check for Enum
|
| 600 |
+
elif issubclass(obj, Enum) and name != "Enum":
|
| 601 |
+
additional_types.append(name)
|
| 602 |
+
# Mark as already imported
|
| 603 |
+
type_tracker.already_imported.add(name)
|
| 604 |
+
|
| 605 |
+
if additional_types:
|
| 606 |
+
type_imports = ", ".join([async_class.__name__] + additional_types)
|
| 607 |
+
imports.append(f"from {async_class.__module__} import {type_imports}")
|
| 608 |
+
else:
|
| 609 |
+
imports.append(
|
| 610 |
+
f"from {async_class.__module__} import {async_class.__name__}"
|
| 611 |
+
)
|
| 612 |
+
|
| 613 |
+
# Add imports for all discovered types
|
| 614 |
+
# Pass the main module name to avoid duplicate imports
|
| 615 |
+
imports.extend(
|
| 616 |
+
type_tracker.get_imports(main_module_name=async_class.__module__)
|
| 617 |
+
)
|
| 618 |
+
|
| 619 |
+
# Add base module import if needed
|
| 620 |
+
if hasattr(inspect.getmodule(async_class), "__name__"):
|
| 621 |
+
module_name = inspect.getmodule(async_class).__name__
|
| 622 |
+
if "." in module_name:
|
| 623 |
+
base_module = module_name.split(".")[0]
|
| 624 |
+
# Only add if not already importing from it
|
| 625 |
+
if not any(imp.startswith(f"from {base_module}") for imp in imports):
|
| 626 |
+
imports.append(f"import {base_module}")
|
| 627 |
+
|
| 628 |
+
return imports
|
| 629 |
+
|
| 630 |
+
@classmethod
|
| 631 |
+
def _get_class_attributes(cls, async_class: type) -> list[tuple[str, type]]:
|
| 632 |
+
"""Extract class attributes that are classes themselves."""
|
| 633 |
+
class_attributes = []
|
| 634 |
+
|
| 635 |
+
# Get resolved type hints to handle string annotations
|
| 636 |
+
try:
|
| 637 |
+
type_hints = get_type_hints(async_class)
|
| 638 |
+
except Exception:
|
| 639 |
+
type_hints = {}
|
| 640 |
+
|
| 641 |
+
# Look for class attributes that are classes
|
| 642 |
+
for name, attr in sorted(inspect.getmembers(async_class)):
|
| 643 |
+
if isinstance(attr, type) and not name.startswith("_"):
|
| 644 |
+
class_attributes.append((name, attr))
|
| 645 |
+
elif name in type_hints:
|
| 646 |
+
# Use resolved type hint instead of raw annotation
|
| 647 |
+
annotation = type_hints[name]
|
| 648 |
+
if isinstance(annotation, type):
|
| 649 |
+
class_attributes.append((name, annotation))
|
| 650 |
+
|
| 651 |
+
return class_attributes
|
| 652 |
+
|
| 653 |
+
@classmethod
|
| 654 |
+
def _generate_inner_class_stub(
|
| 655 |
+
cls,
|
| 656 |
+
name: str,
|
| 657 |
+
attr: type,
|
| 658 |
+
indent: str = " ",
|
| 659 |
+
type_tracker: Optional[TypeTracker] = None,
|
| 660 |
+
) -> list[str]:
|
| 661 |
+
"""Generate stub for an inner class."""
|
| 662 |
+
stub_lines = []
|
| 663 |
+
stub_lines.append(f"{indent}class {name}Sync:")
|
| 664 |
+
|
| 665 |
+
# Add docstring if available
|
| 666 |
+
if hasattr(attr, "__doc__") and attr.__doc__:
|
| 667 |
+
stub_lines.extend(
|
| 668 |
+
cls._format_docstring_for_stub(attr.__doc__, f"{indent} ")
|
| 669 |
+
)
|
| 670 |
+
|
| 671 |
+
# Add __init__ if it exists
|
| 672 |
+
if hasattr(attr, "__init__"):
|
| 673 |
+
try:
|
| 674 |
+
init_method = getattr(attr, "__init__")
|
| 675 |
+
init_sig = inspect.signature(init_method)
|
| 676 |
+
|
| 677 |
+
# Try to get type hints
|
| 678 |
+
try:
|
| 679 |
+
from typing import get_type_hints
|
| 680 |
+
init_hints = get_type_hints(init_method)
|
| 681 |
+
except Exception:
|
| 682 |
+
init_hints = {}
|
| 683 |
+
|
| 684 |
+
# Format parameters
|
| 685 |
+
params_str = cls._format_method_parameters(
|
| 686 |
+
init_sig, type_hints=init_hints, type_tracker=type_tracker
|
| 687 |
+
)
|
| 688 |
+
# Add __init__ docstring if available (before the method)
|
| 689 |
+
if hasattr(init_method, "__doc__") and init_method.__doc__:
|
| 690 |
+
stub_lines.extend(
|
| 691 |
+
cls._format_docstring_for_stub(
|
| 692 |
+
init_method.__doc__, f"{indent} "
|
| 693 |
+
)
|
| 694 |
+
)
|
| 695 |
+
stub_lines.append(
|
| 696 |
+
f"{indent} def __init__({params_str}) -> None: ..."
|
| 697 |
+
)
|
| 698 |
+
except (ValueError, TypeError):
|
| 699 |
+
stub_lines.append(
|
| 700 |
+
f"{indent} def __init__(self, *args, **kwargs) -> None: ..."
|
| 701 |
+
)
|
| 702 |
+
|
| 703 |
+
# Add methods to the inner class
|
| 704 |
+
has_methods = False
|
| 705 |
+
for method_name, method in sorted(
|
| 706 |
+
inspect.getmembers(attr, predicate=inspect.isfunction)
|
| 707 |
+
):
|
| 708 |
+
if method_name.startswith("_"):
|
| 709 |
+
continue
|
| 710 |
+
|
| 711 |
+
has_methods = True
|
| 712 |
+
try:
|
| 713 |
+
# Add method docstring if available (before the method signature)
|
| 714 |
+
if method.__doc__:
|
| 715 |
+
stub_lines.extend(
|
| 716 |
+
cls._format_docstring_for_stub(method.__doc__, f"{indent} ")
|
| 717 |
+
)
|
| 718 |
+
|
| 719 |
+
method_sig = cls._generate_method_signature(
|
| 720 |
+
method_name, method, is_async=True, type_tracker=type_tracker
|
| 721 |
+
)
|
| 722 |
+
stub_lines.append(f"{indent} {method_sig}")
|
| 723 |
+
except (ValueError, TypeError):
|
| 724 |
+
stub_lines.append(
|
| 725 |
+
f"{indent} def {method_name}(self, *args, **kwargs): ..."
|
| 726 |
+
)
|
| 727 |
+
|
| 728 |
+
if not has_methods:
|
| 729 |
+
stub_lines.append(f"{indent} pass")
|
| 730 |
+
|
| 731 |
+
return stub_lines
|
| 732 |
+
|
| 733 |
+
@classmethod
|
| 734 |
+
def _format_docstring_for_stub(
|
| 735 |
+
cls, docstring: str, indent: str = " "
|
| 736 |
+
) -> list[str]:
|
| 737 |
+
"""Format a docstring for inclusion in a stub file with proper indentation."""
|
| 738 |
+
if not docstring:
|
| 739 |
+
return []
|
| 740 |
+
|
| 741 |
+
# First, dedent the docstring to remove any existing indentation
|
| 742 |
+
dedented = textwrap.dedent(docstring).strip()
|
| 743 |
+
|
| 744 |
+
# Split into lines
|
| 745 |
+
lines = dedented.split("\n")
|
| 746 |
+
|
| 747 |
+
# Build the properly indented docstring
|
| 748 |
+
result = []
|
| 749 |
+
result.append(f'{indent}"""')
|
| 750 |
+
|
| 751 |
+
for line in lines:
|
| 752 |
+
if line.strip(): # Non-empty line
|
| 753 |
+
result.append(f"{indent}{line}")
|
| 754 |
+
else: # Empty line
|
| 755 |
+
result.append("")
|
| 756 |
+
|
| 757 |
+
result.append(f'{indent}"""')
|
| 758 |
+
return result
|
| 759 |
+
|
| 760 |
+
@classmethod
|
| 761 |
+
def _post_process_stub_content(cls, stub_content: list[str]) -> list[str]:
|
| 762 |
+
"""Post-process stub content to fix any remaining issues."""
|
| 763 |
+
processed = []
|
| 764 |
+
|
| 765 |
+
for line in stub_content:
|
| 766 |
+
# Skip processing imports
|
| 767 |
+
if line.startswith(("from ", "import ")):
|
| 768 |
+
processed.append(line)
|
| 769 |
+
continue
|
| 770 |
+
|
| 771 |
+
# Fix method signatures missing return types
|
| 772 |
+
if (
|
| 773 |
+
line.strip().startswith("def ")
|
| 774 |
+
and line.strip().endswith(": ...")
|
| 775 |
+
and ") -> " not in line
|
| 776 |
+
):
|
| 777 |
+
# Add -> None for methods without return annotation
|
| 778 |
+
line = line.replace(": ...", " -> None: ...")
|
| 779 |
+
|
| 780 |
+
processed.append(line)
|
| 781 |
+
|
| 782 |
+
return processed
|
| 783 |
+
|
| 784 |
+
@classmethod
|
| 785 |
+
def generate_stub_file(cls, async_class: type, sync_class: type) -> None:
|
| 786 |
+
"""
|
| 787 |
+
Generate a .pyi stub file for the sync class to help IDEs with type checking.
|
| 788 |
+
"""
|
| 789 |
+
try:
|
| 790 |
+
# Only generate stub if we can determine module path
|
| 791 |
+
if async_class.__module__ == "__main__":
|
| 792 |
+
return
|
| 793 |
+
|
| 794 |
+
module = inspect.getmodule(async_class)
|
| 795 |
+
if not module:
|
| 796 |
+
return
|
| 797 |
+
|
| 798 |
+
module_path = module.__file__
|
| 799 |
+
if not module_path:
|
| 800 |
+
return
|
| 801 |
+
|
| 802 |
+
# Create stub file path in a 'generated' subdirectory
|
| 803 |
+
module_dir = os.path.dirname(module_path)
|
| 804 |
+
stub_dir = os.path.join(module_dir, "generated")
|
| 805 |
+
|
| 806 |
+
# Ensure the generated directory exists
|
| 807 |
+
os.makedirs(stub_dir, exist_ok=True)
|
| 808 |
+
|
| 809 |
+
module_name = os.path.basename(module_path)
|
| 810 |
+
if module_name.endswith(".py"):
|
| 811 |
+
module_name = module_name[:-3]
|
| 812 |
+
|
| 813 |
+
sync_stub_path = os.path.join(stub_dir, f"{sync_class.__name__}.pyi")
|
| 814 |
+
|
| 815 |
+
# Create a type tracker for this stub generation
|
| 816 |
+
type_tracker = TypeTracker()
|
| 817 |
+
|
| 818 |
+
stub_content = []
|
| 819 |
+
|
| 820 |
+
# We'll generate imports after processing all methods to capture all types
|
| 821 |
+
# Leave a placeholder for imports
|
| 822 |
+
imports_placeholder_index = len(stub_content)
|
| 823 |
+
stub_content.append("") # Will be replaced with imports later
|
| 824 |
+
|
| 825 |
+
# Class definition
|
| 826 |
+
stub_content.append(f"class {sync_class.__name__}:")
|
| 827 |
+
|
| 828 |
+
# Docstring
|
| 829 |
+
if async_class.__doc__:
|
| 830 |
+
stub_content.extend(
|
| 831 |
+
cls._format_docstring_for_stub(async_class.__doc__, " ")
|
| 832 |
+
)
|
| 833 |
+
|
| 834 |
+
# Generate __init__
|
| 835 |
+
try:
|
| 836 |
+
init_method = async_class.__init__
|
| 837 |
+
init_signature = inspect.signature(init_method)
|
| 838 |
+
|
| 839 |
+
# Try to get type hints for __init__
|
| 840 |
+
try:
|
| 841 |
+
from typing import get_type_hints
|
| 842 |
+
init_hints = get_type_hints(init_method)
|
| 843 |
+
except Exception:
|
| 844 |
+
init_hints = {}
|
| 845 |
+
|
| 846 |
+
# Format parameters
|
| 847 |
+
params_str = cls._format_method_parameters(
|
| 848 |
+
init_signature, type_hints=init_hints, type_tracker=type_tracker
|
| 849 |
+
)
|
| 850 |
+
# Add __init__ docstring if available (before the method)
|
| 851 |
+
if hasattr(init_method, "__doc__") and init_method.__doc__:
|
| 852 |
+
stub_content.extend(
|
| 853 |
+
cls._format_docstring_for_stub(init_method.__doc__, " ")
|
| 854 |
+
)
|
| 855 |
+
stub_content.append(f" def __init__({params_str}) -> None: ...")
|
| 856 |
+
except (ValueError, TypeError):
|
| 857 |
+
stub_content.append(
|
| 858 |
+
" def __init__(self, *args, **kwargs) -> None: ..."
|
| 859 |
+
)
|
| 860 |
+
|
| 861 |
+
stub_content.append("") # Add newline after __init__
|
| 862 |
+
|
| 863 |
+
# Get class attributes
|
| 864 |
+
class_attributes = cls._get_class_attributes(async_class)
|
| 865 |
+
|
| 866 |
+
# Generate inner classes
|
| 867 |
+
for name, attr in class_attributes:
|
| 868 |
+
inner_class_stub = cls._generate_inner_class_stub(
|
| 869 |
+
name, attr, type_tracker=type_tracker
|
| 870 |
+
)
|
| 871 |
+
stub_content.extend(inner_class_stub)
|
| 872 |
+
stub_content.append("") # Add newline after the inner class
|
| 873 |
+
|
| 874 |
+
# Add methods to the main class
|
| 875 |
+
processed_methods = set() # Keep track of methods we've processed
|
| 876 |
+
for name, method in sorted(
|
| 877 |
+
inspect.getmembers(async_class, predicate=inspect.isfunction)
|
| 878 |
+
):
|
| 879 |
+
if name.startswith("_") or name in processed_methods:
|
| 880 |
+
continue
|
| 881 |
+
|
| 882 |
+
processed_methods.add(name)
|
| 883 |
+
|
| 884 |
+
try:
|
| 885 |
+
method_sig = cls._generate_method_signature(
|
| 886 |
+
name, method, is_async=True, type_tracker=type_tracker
|
| 887 |
+
)
|
| 888 |
+
|
| 889 |
+
# Add docstring if available (before the method signature for proper formatting)
|
| 890 |
+
if method.__doc__:
|
| 891 |
+
stub_content.extend(
|
| 892 |
+
cls._format_docstring_for_stub(method.__doc__, " ")
|
| 893 |
+
)
|
| 894 |
+
|
| 895 |
+
stub_content.append(f" {method_sig}")
|
| 896 |
+
|
| 897 |
+
stub_content.append("") # Add newline after each method
|
| 898 |
+
|
| 899 |
+
except (ValueError, TypeError):
|
| 900 |
+
# If we can't get the signature, just add a simple stub
|
| 901 |
+
stub_content.append(f" def {name}(self, *args, **kwargs): ...")
|
| 902 |
+
stub_content.append("") # Add newline
|
| 903 |
+
|
| 904 |
+
# Add properties
|
| 905 |
+
for name, prop in sorted(
|
| 906 |
+
inspect.getmembers(async_class, lambda x: isinstance(x, property))
|
| 907 |
+
):
|
| 908 |
+
stub_content.append(" @property")
|
| 909 |
+
stub_content.append(f" def {name}(self) -> Any: ...")
|
| 910 |
+
if prop.fset:
|
| 911 |
+
stub_content.append(f" @{name}.setter")
|
| 912 |
+
stub_content.append(
|
| 913 |
+
f" def {name}(self, value: Any) -> None: ..."
|
| 914 |
+
)
|
| 915 |
+
stub_content.append("") # Add newline after each property
|
| 916 |
+
|
| 917 |
+
# Add placeholders for the nested class instances
|
| 918 |
+
# Check the actual attribute names from class annotations and attributes
|
| 919 |
+
attribute_mappings = {}
|
| 920 |
+
|
| 921 |
+
# First check annotations for typed attributes (including from parent classes)
|
| 922 |
+
# Resolve string annotations to actual types
|
| 923 |
+
try:
|
| 924 |
+
all_annotations = get_type_hints(async_class)
|
| 925 |
+
except Exception:
|
| 926 |
+
# Fallback to raw annotations
|
| 927 |
+
all_annotations = {}
|
| 928 |
+
for base_class in reversed(inspect.getmro(async_class)):
|
| 929 |
+
if hasattr(base_class, "__annotations__"):
|
| 930 |
+
all_annotations.update(base_class.__annotations__)
|
| 931 |
+
|
| 932 |
+
for attr_name, attr_type in sorted(all_annotations.items()):
|
| 933 |
+
for class_name, class_type in class_attributes:
|
| 934 |
+
# If the class type matches the annotated type
|
| 935 |
+
if (
|
| 936 |
+
attr_type == class_type
|
| 937 |
+
or (hasattr(attr_type, "__name__") and attr_type.__name__ == class_name)
|
| 938 |
+
or (isinstance(attr_type, str) and attr_type == class_name)
|
| 939 |
+
):
|
| 940 |
+
attribute_mappings[class_name] = attr_name
|
| 941 |
+
|
| 942 |
+
# Remove the extra checking - annotations should be sufficient
|
| 943 |
+
|
| 944 |
+
# Add the attribute declarations with proper names
|
| 945 |
+
for class_name, class_type in class_attributes:
|
| 946 |
+
# Check if there's a mapping from annotation
|
| 947 |
+
attr_name = attribute_mappings.get(class_name, class_name)
|
| 948 |
+
# Use the annotation name if it exists, even if the attribute doesn't exist yet
|
| 949 |
+
# This is because the attribute might be created at runtime
|
| 950 |
+
stub_content.append(f" {attr_name}: {class_name}Sync")
|
| 951 |
+
|
| 952 |
+
stub_content.append("") # Add a final newline
|
| 953 |
+
|
| 954 |
+
# Now generate imports with all discovered types
|
| 955 |
+
imports = cls._generate_imports(async_class, type_tracker)
|
| 956 |
+
|
| 957 |
+
# Deduplicate imports while preserving order
|
| 958 |
+
seen = set()
|
| 959 |
+
unique_imports = []
|
| 960 |
+
for imp in imports:
|
| 961 |
+
if imp not in seen:
|
| 962 |
+
seen.add(imp)
|
| 963 |
+
unique_imports.append(imp)
|
| 964 |
+
else:
|
| 965 |
+
logging.warning(f"Duplicate import detected: {imp}")
|
| 966 |
+
|
| 967 |
+
# Replace the placeholder with actual imports
|
| 968 |
+
stub_content[imports_placeholder_index : imports_placeholder_index + 1] = (
|
| 969 |
+
unique_imports
|
| 970 |
+
)
|
| 971 |
+
|
| 972 |
+
# Post-process stub content
|
| 973 |
+
stub_content = cls._post_process_stub_content(stub_content)
|
| 974 |
+
|
| 975 |
+
# Write stub file
|
| 976 |
+
with open(sync_stub_path, "w") as f:
|
| 977 |
+
f.write("\n".join(stub_content))
|
| 978 |
+
|
| 979 |
+
logging.info(f"Generated stub file: {sync_stub_path}")
|
| 980 |
+
|
| 981 |
+
except Exception as e:
|
| 982 |
+
# If stub generation fails, log the error but don't break the main functionality
|
| 983 |
+
logging.error(
|
| 984 |
+
f"Error generating stub file for {sync_class.__name__}: {str(e)}"
|
| 985 |
+
)
|
| 986 |
+
import traceback
|
| 987 |
+
|
| 988 |
+
logging.error(traceback.format_exc())
|
| 989 |
+
|
| 990 |
+
|
| 991 |
+
def create_sync_class(async_class: type, thread_pool_size=10) -> type:
|
| 992 |
+
"""
|
| 993 |
+
Creates a sync version of an async class
|
| 994 |
+
|
| 995 |
+
Args:
|
| 996 |
+
async_class: The async class to convert
|
| 997 |
+
thread_pool_size: Size of thread pool to use
|
| 998 |
+
|
| 999 |
+
Returns:
|
| 1000 |
+
A new class with sync versions of all async methods
|
| 1001 |
+
"""
|
| 1002 |
+
return AsyncToSyncConverter.create_sync_class(async_class, thread_pool_size)
|
comfy_api/internal/singleton.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import TypeVar
|
| 2 |
+
|
| 3 |
+
class SingletonMetaclass(type):
|
| 4 |
+
T = TypeVar("T", bound="SingletonMetaclass")
|
| 5 |
+
_instances = {}
|
| 6 |
+
|
| 7 |
+
def __call__(cls, *args, **kwargs):
|
| 8 |
+
if cls not in cls._instances:
|
| 9 |
+
cls._instances[cls] = super(SingletonMetaclass, cls).__call__(
|
| 10 |
+
*args, **kwargs
|
| 11 |
+
)
|
| 12 |
+
return cls._instances[cls]
|
| 13 |
+
|
| 14 |
+
def inject_instance(cls: type[T], instance: T) -> None:
|
| 15 |
+
assert cls not in SingletonMetaclass._instances, (
|
| 16 |
+
"Cannot inject instance after first instantiation"
|
| 17 |
+
)
|
| 18 |
+
SingletonMetaclass._instances[cls] = instance
|
| 19 |
+
|
| 20 |
+
def get_instance(cls: type[T], *args, **kwargs) -> T:
|
| 21 |
+
"""
|
| 22 |
+
Gets the singleton instance of the class, creating it if it doesn't exist.
|
| 23 |
+
"""
|
| 24 |
+
if cls not in SingletonMetaclass._instances:
|
| 25 |
+
SingletonMetaclass._instances[cls] = super(
|
| 26 |
+
SingletonMetaclass, cls
|
| 27 |
+
).__call__(*args, **kwargs)
|
| 28 |
+
return cls._instances[cls]
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
class ProxiedSingleton(object, metaclass=SingletonMetaclass):
|
| 32 |
+
def __init__(self):
|
| 33 |
+
super().__init__()
|
comfy_api/latest/__init__.py
ADDED
|
@@ -0,0 +1,177 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from abc import ABC, abstractmethod
|
| 2 |
+
from typing import TYPE_CHECKING
|
| 3 |
+
from comfy_api.internal import ComfyAPIBase
|
| 4 |
+
from comfy_api.internal.singleton import ProxiedSingleton
|
| 5 |
+
from comfy_api.internal.async_to_sync import create_sync_class
|
| 6 |
+
from ._input import ImageInput, AudioInput, MaskInput, LatentInput, VideoInput
|
| 7 |
+
from ._input_impl import VideoFromFile, VideoFromComponents
|
| 8 |
+
from ._util import VideoCodec, VideoContainer, VideoComponents, MESH, VOXEL, SPLAT, File3D
|
| 9 |
+
from . import _io_public as io
|
| 10 |
+
from . import _ui_public as ui
|
| 11 |
+
from comfy_execution.utils import get_executing_context
|
| 12 |
+
from comfy_execution.progress import get_progress_state, PreviewImageTuple
|
| 13 |
+
from PIL import Image
|
| 14 |
+
from comfy.cli_args import args
|
| 15 |
+
import numpy as np
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class ComfyAPI_latest(ComfyAPIBase):
|
| 19 |
+
VERSION = "latest"
|
| 20 |
+
STABLE = False
|
| 21 |
+
|
| 22 |
+
def __init__(self):
|
| 23 |
+
super().__init__()
|
| 24 |
+
self.node_replacement = self.NodeReplacement()
|
| 25 |
+
self.execution = self.Execution()
|
| 26 |
+
self.caching = self.Caching()
|
| 27 |
+
|
| 28 |
+
class NodeReplacement(ProxiedSingleton):
|
| 29 |
+
async def register(self, node_replace: io.NodeReplace) -> None:
|
| 30 |
+
"""Register a node replacement mapping."""
|
| 31 |
+
from server import PromptServer
|
| 32 |
+
PromptServer.instance.node_replace_manager.register(node_replace)
|
| 33 |
+
|
| 34 |
+
class Execution(ProxiedSingleton):
|
| 35 |
+
async def set_progress(
|
| 36 |
+
self,
|
| 37 |
+
value: float,
|
| 38 |
+
max_value: float,
|
| 39 |
+
node_id: str | None = None,
|
| 40 |
+
preview_image: Image.Image | ImageInput | None = None,
|
| 41 |
+
ignore_size_limit: bool = False,
|
| 42 |
+
) -> None:
|
| 43 |
+
"""
|
| 44 |
+
Update the progress bar displayed in the ComfyUI interface.
|
| 45 |
+
|
| 46 |
+
This function allows custom nodes and API calls to report their progress
|
| 47 |
+
back to the user interface, providing visual feedback during long operations.
|
| 48 |
+
|
| 49 |
+
Migration from previous API: comfy.utils.PROGRESS_BAR_HOOK
|
| 50 |
+
"""
|
| 51 |
+
executing_context = get_executing_context()
|
| 52 |
+
if node_id is None and executing_context is not None:
|
| 53 |
+
node_id = executing_context.node_id
|
| 54 |
+
if node_id is None:
|
| 55 |
+
raise ValueError("node_id must be provided if not in executing context")
|
| 56 |
+
|
| 57 |
+
# Convert preview_image to PreviewImageTuple if needed
|
| 58 |
+
to_display: PreviewImageTuple | Image.Image | ImageInput | None = preview_image
|
| 59 |
+
if to_display is not None:
|
| 60 |
+
# First convert to PIL Image if needed
|
| 61 |
+
if isinstance(to_display, ImageInput):
|
| 62 |
+
# Convert ImageInput (torch.Tensor) to PIL Image
|
| 63 |
+
# Handle tensor shape [B, H, W, C] -> get first image if batch
|
| 64 |
+
tensor = to_display
|
| 65 |
+
if len(tensor.shape) == 4:
|
| 66 |
+
tensor = tensor[0]
|
| 67 |
+
|
| 68 |
+
# Convert to numpy array and scale to 0-255
|
| 69 |
+
image_np = (tensor.cpu().numpy() * 255).astype(np.uint8)
|
| 70 |
+
to_display = Image.fromarray(image_np)
|
| 71 |
+
|
| 72 |
+
if isinstance(to_display, Image.Image):
|
| 73 |
+
# Detect image format from PIL Image
|
| 74 |
+
image_format = to_display.format if to_display.format else "JPEG"
|
| 75 |
+
# Use None for preview_size if ignore_size_limit is True
|
| 76 |
+
preview_size = None if ignore_size_limit else args.preview_size
|
| 77 |
+
to_display = (image_format, to_display, preview_size)
|
| 78 |
+
|
| 79 |
+
get_progress_state().update_progress(
|
| 80 |
+
node_id=node_id,
|
| 81 |
+
value=value,
|
| 82 |
+
max_value=max_value,
|
| 83 |
+
image=to_display,
|
| 84 |
+
)
|
| 85 |
+
|
| 86 |
+
class Caching(ProxiedSingleton):
|
| 87 |
+
"""
|
| 88 |
+
External cache provider API for sharing cached node outputs
|
| 89 |
+
across ComfyUI instances.
|
| 90 |
+
|
| 91 |
+
Example::
|
| 92 |
+
|
| 93 |
+
from comfy_api.latest import Caching
|
| 94 |
+
|
| 95 |
+
class MyCacheProvider(Caching.CacheProvider):
|
| 96 |
+
async def on_lookup(self, context):
|
| 97 |
+
... # check external storage
|
| 98 |
+
|
| 99 |
+
async def on_store(self, context, value):
|
| 100 |
+
... # store to external storage
|
| 101 |
+
|
| 102 |
+
Caching.register_provider(MyCacheProvider())
|
| 103 |
+
"""
|
| 104 |
+
from ._caching import CacheProvider, CacheContext, CacheValue
|
| 105 |
+
|
| 106 |
+
async def register_provider(self, provider: "ComfyAPI_latest.Caching.CacheProvider") -> None:
|
| 107 |
+
"""Register an external cache provider. Providers are called in registration order."""
|
| 108 |
+
from comfy_execution.cache_provider import register_cache_provider
|
| 109 |
+
register_cache_provider(provider)
|
| 110 |
+
|
| 111 |
+
async def unregister_provider(self, provider: "ComfyAPI_latest.Caching.CacheProvider") -> None:
|
| 112 |
+
"""Unregister a previously registered cache provider."""
|
| 113 |
+
from comfy_execution.cache_provider import unregister_cache_provider
|
| 114 |
+
unregister_cache_provider(provider)
|
| 115 |
+
|
| 116 |
+
class ComfyExtension(ABC):
|
| 117 |
+
async def on_load(self) -> None:
|
| 118 |
+
"""
|
| 119 |
+
Called when an extension is loaded.
|
| 120 |
+
This should be used to initialize any global resources needed by the extension.
|
| 121 |
+
"""
|
| 122 |
+
|
| 123 |
+
@abstractmethod
|
| 124 |
+
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
| 125 |
+
"""
|
| 126 |
+
Returns a list of nodes that this extension provides.
|
| 127 |
+
"""
|
| 128 |
+
|
| 129 |
+
class Input:
|
| 130 |
+
Image = ImageInput
|
| 131 |
+
Audio = AudioInput
|
| 132 |
+
Mask = MaskInput
|
| 133 |
+
Latent = LatentInput
|
| 134 |
+
Video = VideoInput
|
| 135 |
+
|
| 136 |
+
class InputImpl:
|
| 137 |
+
VideoFromFile = VideoFromFile
|
| 138 |
+
VideoFromComponents = VideoFromComponents
|
| 139 |
+
|
| 140 |
+
class Types:
|
| 141 |
+
VideoCodec = VideoCodec
|
| 142 |
+
VideoContainer = VideoContainer
|
| 143 |
+
VideoComponents = VideoComponents
|
| 144 |
+
MESH = MESH
|
| 145 |
+
VOXEL = VOXEL
|
| 146 |
+
SPLAT = SPLAT
|
| 147 |
+
File3D = File3D
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
Caching = ComfyAPI_latest.Caching
|
| 151 |
+
|
| 152 |
+
ComfyAPI = ComfyAPI_latest
|
| 153 |
+
|
| 154 |
+
# Create a synchronous version of the API
|
| 155 |
+
if TYPE_CHECKING:
|
| 156 |
+
import comfy_api.latest.generated.ComfyAPISyncStub # type: ignore
|
| 157 |
+
|
| 158 |
+
ComfyAPISync: type[comfy_api.latest.generated.ComfyAPISyncStub.ComfyAPISyncStub]
|
| 159 |
+
ComfyAPISync = create_sync_class(ComfyAPI_latest)
|
| 160 |
+
|
| 161 |
+
# create new aliases for io and ui
|
| 162 |
+
IO = io
|
| 163 |
+
UI = ui
|
| 164 |
+
|
| 165 |
+
__all__ = [
|
| 166 |
+
"ComfyAPI",
|
| 167 |
+
"ComfyAPISync",
|
| 168 |
+
"Input",
|
| 169 |
+
"InputImpl",
|
| 170 |
+
"Types",
|
| 171 |
+
"Caching",
|
| 172 |
+
"ComfyExtension",
|
| 173 |
+
"io",
|
| 174 |
+
"IO",
|
| 175 |
+
"ui",
|
| 176 |
+
"UI",
|
| 177 |
+
]
|
comfy_api/latest/_caching.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from abc import ABC, abstractmethod
|
| 2 |
+
from typing import Optional
|
| 3 |
+
from dataclasses import dataclass
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
@dataclass
|
| 7 |
+
class CacheContext:
|
| 8 |
+
node_id: str
|
| 9 |
+
class_type: str
|
| 10 |
+
cache_key_hash: str # SHA256 hex digest
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
@dataclass
|
| 14 |
+
class CacheValue:
|
| 15 |
+
outputs: list
|
| 16 |
+
ui: dict = None
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class CacheProvider(ABC):
|
| 20 |
+
"""Abstract base class for external cache providers.
|
| 21 |
+
Exceptions from provider methods are caught by the caller and never break execution.
|
| 22 |
+
"""
|
| 23 |
+
|
| 24 |
+
@abstractmethod
|
| 25 |
+
async def on_lookup(self, context: CacheContext) -> Optional[CacheValue]:
|
| 26 |
+
"""Called on local cache miss. Return CacheValue if found, None otherwise."""
|
| 27 |
+
pass
|
| 28 |
+
|
| 29 |
+
@abstractmethod
|
| 30 |
+
async def on_store(self, context: CacheContext, value: CacheValue) -> None:
|
| 31 |
+
"""Called after local store. Dispatched via asyncio.create_task."""
|
| 32 |
+
pass
|
| 33 |
+
|
| 34 |
+
def should_cache(self, context: CacheContext, value: Optional[CacheValue] = None) -> bool:
|
| 35 |
+
"""Return False to skip external caching for this node. Default: True."""
|
| 36 |
+
return True
|
| 37 |
+
|
| 38 |
+
def on_prompt_start(self, prompt_id: str) -> None:
|
| 39 |
+
pass
|
| 40 |
+
|
| 41 |
+
def on_prompt_end(self, prompt_id: str) -> None:
|
| 42 |
+
pass
|
comfy_api/latest/_input/__init__.py
ADDED
|
@@ -0,0 +1,17 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .basic_types import ImageInput, AudioInput, MaskInput, LatentInput
|
| 2 |
+
from .curve_types import CurvePoint, CurveInput, MonotoneCubicCurve, LinearCurve
|
| 3 |
+
from .range_types import RangeInput
|
| 4 |
+
from .video_types import VideoInput
|
| 5 |
+
|
| 6 |
+
__all__ = [
|
| 7 |
+
"ImageInput",
|
| 8 |
+
"AudioInput",
|
| 9 |
+
"VideoInput",
|
| 10 |
+
"MaskInput",
|
| 11 |
+
"LatentInput",
|
| 12 |
+
"CurvePoint",
|
| 13 |
+
"CurveInput",
|
| 14 |
+
"MonotoneCubicCurve",
|
| 15 |
+
"LinearCurve",
|
| 16 |
+
"RangeInput",
|
| 17 |
+
]
|
comfy_api/latest/_input/basic_types.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from typing import TypedDict, Optional
|
| 3 |
+
|
| 4 |
+
ImageInput = torch.Tensor
|
| 5 |
+
"""
|
| 6 |
+
An image in format [B, H, W, C] where B is the batch size, C is the number of channels,
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
MaskInput = torch.Tensor
|
| 10 |
+
"""
|
| 11 |
+
A mask in format [B, H, W] where B is the batch size
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
class AudioInput(TypedDict):
|
| 15 |
+
"""
|
| 16 |
+
TypedDict representing audio input.
|
| 17 |
+
"""
|
| 18 |
+
|
| 19 |
+
waveform: torch.Tensor
|
| 20 |
+
"""
|
| 21 |
+
Tensor in the format [B, C, T] where B is the batch size, C is the number of channels,
|
| 22 |
+
"""
|
| 23 |
+
|
| 24 |
+
sample_rate: int
|
| 25 |
+
|
| 26 |
+
class LatentInput(TypedDict):
|
| 27 |
+
"""
|
| 28 |
+
TypedDict representing latent input.
|
| 29 |
+
"""
|
| 30 |
+
|
| 31 |
+
samples: torch.Tensor
|
| 32 |
+
"""
|
| 33 |
+
Tensor in the format [B, C, H, W] where B is the batch size, C is the number of channels,
|
| 34 |
+
H is the height, and W is the width.
|
| 35 |
+
"""
|
| 36 |
+
|
| 37 |
+
noise_mask: Optional[MaskInput]
|
| 38 |
+
"""
|
| 39 |
+
Optional noise mask tensor in the same format as samples.
|
| 40 |
+
"""
|
| 41 |
+
|
| 42 |
+
batch_index: Optional[list[int]]
|
comfy_api/latest/_input/curve_types.py
ADDED
|
@@ -0,0 +1,219 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import logging
|
| 4 |
+
import math
|
| 5 |
+
from abc import ABC, abstractmethod
|
| 6 |
+
import numpy as np
|
| 7 |
+
|
| 8 |
+
logger = logging.getLogger(__name__)
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
CurvePoint = tuple[float, float]
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class CurveInput(ABC):
|
| 15 |
+
"""Abstract base class for curve inputs.
|
| 16 |
+
|
| 17 |
+
Subclasses represent different curve representations (control-point
|
| 18 |
+
interpolation, analytical functions, LUT-based, etc.) while exposing a
|
| 19 |
+
uniform evaluation interface to downstream nodes.
|
| 20 |
+
"""
|
| 21 |
+
|
| 22 |
+
@property
|
| 23 |
+
@abstractmethod
|
| 24 |
+
def points(self) -> list[CurvePoint]:
|
| 25 |
+
"""The control points that define this curve."""
|
| 26 |
+
|
| 27 |
+
@abstractmethod
|
| 28 |
+
def interp(self, x: float) -> float:
|
| 29 |
+
"""Evaluate the curve at a single *x* value in [0, 1]."""
|
| 30 |
+
|
| 31 |
+
def interp_array(self, xs: np.ndarray) -> np.ndarray:
|
| 32 |
+
"""Vectorised evaluation over a numpy array of x values.
|
| 33 |
+
|
| 34 |
+
Subclasses should override this for better performance. The default
|
| 35 |
+
falls back to scalar ``interp`` calls.
|
| 36 |
+
"""
|
| 37 |
+
return np.fromiter((self.interp(float(x)) for x in xs), dtype=np.float64, count=len(xs))
|
| 38 |
+
|
| 39 |
+
def to_lut(self, size: int = 256) -> np.ndarray:
|
| 40 |
+
"""Generate a float64 lookup table of *size* evenly-spaced samples in [0, 1]."""
|
| 41 |
+
return self.interp_array(np.linspace(0.0, 1.0, size))
|
| 42 |
+
|
| 43 |
+
@staticmethod
|
| 44 |
+
def from_raw(data) -> CurveInput:
|
| 45 |
+
"""Convert raw curve data (dict or point list) to a CurveInput instance.
|
| 46 |
+
|
| 47 |
+
Accepts:
|
| 48 |
+
- A ``CurveInput`` instance (returned as-is).
|
| 49 |
+
- A dict with ``"points"`` and optional ``"interpolation"`` keys.
|
| 50 |
+
- A bare list/sequence of ``(x, y)`` pairs (defaults to monotone cubic).
|
| 51 |
+
"""
|
| 52 |
+
if isinstance(data, CurveInput):
|
| 53 |
+
return data
|
| 54 |
+
if isinstance(data, dict):
|
| 55 |
+
raw_points = data["points"]
|
| 56 |
+
interpolation = data.get("interpolation", "monotone_cubic")
|
| 57 |
+
else:
|
| 58 |
+
raw_points = data
|
| 59 |
+
interpolation = "monotone_cubic"
|
| 60 |
+
points = [(float(x), float(y)) for x, y in raw_points]
|
| 61 |
+
if interpolation == "linear":
|
| 62 |
+
return LinearCurve(points)
|
| 63 |
+
if interpolation != "monotone_cubic":
|
| 64 |
+
logger.warning("Unknown curve interpolation %r, falling back to monotone_cubic", interpolation)
|
| 65 |
+
return MonotoneCubicCurve(points)
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
class MonotoneCubicCurve(CurveInput):
|
| 69 |
+
"""Monotone cubic Hermite interpolation over control points.
|
| 70 |
+
|
| 71 |
+
Mirrors the frontend ``createMonotoneInterpolator`` in
|
| 72 |
+
``ComfyUI_frontend/src/components/curve/curveUtils.ts`` so that
|
| 73 |
+
backend evaluation matches the editor preview exactly.
|
| 74 |
+
|
| 75 |
+
All heavy work (sorting, slope computation) happens once at construction.
|
| 76 |
+
``interp_array`` is fully vectorised with numpy.
|
| 77 |
+
"""
|
| 78 |
+
|
| 79 |
+
def __init__(self, control_points: list[CurvePoint]):
|
| 80 |
+
sorted_pts = sorted(control_points, key=lambda p: p[0])
|
| 81 |
+
self._points = [(float(x), float(y)) for x, y in sorted_pts]
|
| 82 |
+
self._xs = np.array([p[0] for p in self._points], dtype=np.float64)
|
| 83 |
+
self._ys = np.array([p[1] for p in self._points], dtype=np.float64)
|
| 84 |
+
self._slopes = self._compute_slopes()
|
| 85 |
+
|
| 86 |
+
@property
|
| 87 |
+
def points(self) -> list[CurvePoint]:
|
| 88 |
+
return list(self._points)
|
| 89 |
+
|
| 90 |
+
def _compute_slopes(self) -> np.ndarray:
|
| 91 |
+
xs, ys = self._xs, self._ys
|
| 92 |
+
n = len(xs)
|
| 93 |
+
if n < 2:
|
| 94 |
+
return np.zeros(n, dtype=np.float64)
|
| 95 |
+
|
| 96 |
+
dx = np.diff(xs)
|
| 97 |
+
dy = np.diff(ys)
|
| 98 |
+
dx_safe = np.where(dx == 0, 1.0, dx)
|
| 99 |
+
deltas = np.where(dx == 0, 0.0, dy / dx_safe)
|
| 100 |
+
|
| 101 |
+
slopes = np.empty(n, dtype=np.float64)
|
| 102 |
+
slopes[0] = deltas[0]
|
| 103 |
+
slopes[-1] = deltas[-1]
|
| 104 |
+
for i in range(1, n - 1):
|
| 105 |
+
if deltas[i - 1] * deltas[i] <= 0:
|
| 106 |
+
slopes[i] = 0.0
|
| 107 |
+
else:
|
| 108 |
+
slopes[i] = (deltas[i - 1] + deltas[i]) / 2
|
| 109 |
+
|
| 110 |
+
for i in range(n - 1):
|
| 111 |
+
if deltas[i] == 0:
|
| 112 |
+
slopes[i] = 0.0
|
| 113 |
+
slopes[i + 1] = 0.0
|
| 114 |
+
else:
|
| 115 |
+
alpha = slopes[i] / deltas[i]
|
| 116 |
+
beta = slopes[i + 1] / deltas[i]
|
| 117 |
+
s = alpha * alpha + beta * beta
|
| 118 |
+
if s > 9:
|
| 119 |
+
t = 3 / math.sqrt(s)
|
| 120 |
+
slopes[i] = t * alpha * deltas[i]
|
| 121 |
+
slopes[i + 1] = t * beta * deltas[i]
|
| 122 |
+
return slopes
|
| 123 |
+
|
| 124 |
+
def interp(self, x: float) -> float:
|
| 125 |
+
xs, ys, slopes = self._xs, self._ys, self._slopes
|
| 126 |
+
n = len(xs)
|
| 127 |
+
if n == 0:
|
| 128 |
+
return 0.0
|
| 129 |
+
if n == 1:
|
| 130 |
+
return float(ys[0])
|
| 131 |
+
if x <= xs[0]:
|
| 132 |
+
return float(ys[0])
|
| 133 |
+
if x >= xs[-1]:
|
| 134 |
+
return float(ys[-1])
|
| 135 |
+
|
| 136 |
+
hi = int(np.searchsorted(xs, x, side='right'))
|
| 137 |
+
hi = min(hi, n - 1)
|
| 138 |
+
lo = hi - 1
|
| 139 |
+
|
| 140 |
+
dx = xs[hi] - xs[lo]
|
| 141 |
+
if dx == 0:
|
| 142 |
+
return float(ys[lo])
|
| 143 |
+
|
| 144 |
+
t = (x - xs[lo]) / dx
|
| 145 |
+
t2 = t * t
|
| 146 |
+
t3 = t2 * t
|
| 147 |
+
h00 = 2 * t3 - 3 * t2 + 1
|
| 148 |
+
h10 = t3 - 2 * t2 + t
|
| 149 |
+
h01 = -2 * t3 + 3 * t2
|
| 150 |
+
h11 = t3 - t2
|
| 151 |
+
return float(h00 * ys[lo] + h10 * dx * slopes[lo] + h01 * ys[hi] + h11 * dx * slopes[hi])
|
| 152 |
+
|
| 153 |
+
def interp_array(self, xs_in: np.ndarray) -> np.ndarray:
|
| 154 |
+
"""Fully vectorised evaluation using numpy."""
|
| 155 |
+
xs, ys, slopes = self._xs, self._ys, self._slopes
|
| 156 |
+
n = len(xs)
|
| 157 |
+
if n == 0:
|
| 158 |
+
return np.zeros_like(xs_in, dtype=np.float64)
|
| 159 |
+
if n == 1:
|
| 160 |
+
return np.full_like(xs_in, ys[0], dtype=np.float64)
|
| 161 |
+
|
| 162 |
+
hi = np.searchsorted(xs, xs_in, side='right').clip(1, n - 1)
|
| 163 |
+
lo = hi - 1
|
| 164 |
+
|
| 165 |
+
dx = xs[hi] - xs[lo]
|
| 166 |
+
dx_safe = np.where(dx == 0, 1.0, dx)
|
| 167 |
+
t = np.where(dx == 0, 0.0, (xs_in - xs[lo]) / dx_safe)
|
| 168 |
+
t2 = t * t
|
| 169 |
+
t3 = t2 * t
|
| 170 |
+
|
| 171 |
+
h00 = 2 * t3 - 3 * t2 + 1
|
| 172 |
+
h10 = t3 - 2 * t2 + t
|
| 173 |
+
h01 = -2 * t3 + 3 * t2
|
| 174 |
+
h11 = t3 - t2
|
| 175 |
+
|
| 176 |
+
result = h00 * ys[lo] + h10 * dx * slopes[lo] + h01 * ys[hi] + h11 * dx * slopes[hi]
|
| 177 |
+
result = np.where(xs_in <= xs[0], ys[0], result)
|
| 178 |
+
result = np.where(xs_in >= xs[-1], ys[-1], result)
|
| 179 |
+
return result
|
| 180 |
+
|
| 181 |
+
def __repr__(self) -> str:
|
| 182 |
+
return f"MonotoneCubicCurve(points={self._points})"
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
class LinearCurve(CurveInput):
|
| 186 |
+
"""Piecewise linear interpolation over control points.
|
| 187 |
+
|
| 188 |
+
Mirrors the frontend ``createLinearInterpolator`` in
|
| 189 |
+
``ComfyUI_frontend/src/components/curve/curveUtils.ts``.
|
| 190 |
+
"""
|
| 191 |
+
|
| 192 |
+
def __init__(self, control_points: list[CurvePoint]):
|
| 193 |
+
sorted_pts = sorted(control_points, key=lambda p: p[0])
|
| 194 |
+
self._points = [(float(x), float(y)) for x, y in sorted_pts]
|
| 195 |
+
self._xs = np.array([p[0] for p in self._points], dtype=np.float64)
|
| 196 |
+
self._ys = np.array([p[1] for p in self._points], dtype=np.float64)
|
| 197 |
+
|
| 198 |
+
@property
|
| 199 |
+
def points(self) -> list[CurvePoint]:
|
| 200 |
+
return list(self._points)
|
| 201 |
+
|
| 202 |
+
def interp(self, x: float) -> float:
|
| 203 |
+
xs, ys = self._xs, self._ys
|
| 204 |
+
n = len(xs)
|
| 205 |
+
if n == 0:
|
| 206 |
+
return 0.0
|
| 207 |
+
if n == 1:
|
| 208 |
+
return float(ys[0])
|
| 209 |
+
return float(np.interp(x, xs, ys))
|
| 210 |
+
|
| 211 |
+
def interp_array(self, xs_in: np.ndarray) -> np.ndarray:
|
| 212 |
+
if len(self._xs) == 0:
|
| 213 |
+
return np.zeros_like(xs_in, dtype=np.float64)
|
| 214 |
+
if len(self._xs) == 1:
|
| 215 |
+
return np.full_like(xs_in, self._ys[0], dtype=np.float64)
|
| 216 |
+
return np.interp(xs_in, self._xs, self._ys)
|
| 217 |
+
|
| 218 |
+
def __repr__(self) -> str:
|
| 219 |
+
return f"LinearCurve(points={self._points})"
|
comfy_api/latest/_input/range_types.py
ADDED
|
@@ -0,0 +1,70 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import logging
|
| 4 |
+
import math
|
| 5 |
+
import numpy as np
|
| 6 |
+
|
| 7 |
+
logger = logging.getLogger(__name__)
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class RangeInput:
|
| 11 |
+
"""Represents a levels/range adjustment: input range [min, max] with
|
| 12 |
+
optional midpoint (gamma control).
|
| 13 |
+
|
| 14 |
+
Generates a 1D LUT identical to GIMP's levels mapping:
|
| 15 |
+
1. Normalize input to [0, 1] using [min, max]
|
| 16 |
+
2. Apply gamma correction: pow(value, 1/gamma)
|
| 17 |
+
3. Clamp to [0, 1]
|
| 18 |
+
|
| 19 |
+
The midpoint field is a position in [0, 1] representing where the
|
| 20 |
+
midtone falls within [min, max]. It maps to gamma via:
|
| 21 |
+
gamma = -log2(midpoint)
|
| 22 |
+
So midpoint=0.5 → gamma=1.0 (linear).
|
| 23 |
+
"""
|
| 24 |
+
|
| 25 |
+
def __init__(self, min_val: float, max_val: float, midpoint: float | None = None):
|
| 26 |
+
self.min_val = min_val
|
| 27 |
+
self.max_val = max_val
|
| 28 |
+
self.midpoint = midpoint
|
| 29 |
+
|
| 30 |
+
@staticmethod
|
| 31 |
+
def from_raw(data) -> RangeInput:
|
| 32 |
+
if isinstance(data, RangeInput):
|
| 33 |
+
return data
|
| 34 |
+
if isinstance(data, dict):
|
| 35 |
+
return RangeInput(
|
| 36 |
+
min_val=float(data.get("min", 0.0)),
|
| 37 |
+
max_val=float(data.get("max", 1.0)),
|
| 38 |
+
midpoint=float(data["midpoint"]) if data.get("midpoint") is not None else None,
|
| 39 |
+
)
|
| 40 |
+
raise TypeError(f"Cannot convert {type(data)} to RangeInput")
|
| 41 |
+
|
| 42 |
+
def to_lut(self, size: int = 256) -> np.ndarray:
|
| 43 |
+
"""Generate a float64 lookup table mapping [0, 1] input through this
|
| 44 |
+
levels adjustment.
|
| 45 |
+
|
| 46 |
+
The LUT maps normalized input values (0..1) to output values (0..1),
|
| 47 |
+
matching the GIMP levels formula.
|
| 48 |
+
"""
|
| 49 |
+
xs = np.linspace(0.0, 1.0, size, dtype=np.float64)
|
| 50 |
+
|
| 51 |
+
in_range = self.max_val - self.min_val
|
| 52 |
+
if abs(in_range) < 1e-10:
|
| 53 |
+
return np.where(xs >= self.min_val, 1.0, 0.0).astype(np.float64)
|
| 54 |
+
|
| 55 |
+
# Normalize: map [min, max] → [0, 1]
|
| 56 |
+
result = (xs - self.min_val) / in_range
|
| 57 |
+
result = np.clip(result, 0.0, 1.0)
|
| 58 |
+
|
| 59 |
+
# Gamma correction from midpoint
|
| 60 |
+
if self.midpoint is not None and self.midpoint > 0 and self.midpoint != 0.5:
|
| 61 |
+
gamma = max(-math.log2(self.midpoint), 0.001)
|
| 62 |
+
inv_gamma = 1.0 / gamma
|
| 63 |
+
mask = result > 0
|
| 64 |
+
result[mask] = np.power(result[mask], inv_gamma)
|
| 65 |
+
|
| 66 |
+
return result
|
| 67 |
+
|
| 68 |
+
def __repr__(self) -> str:
|
| 69 |
+
mid = f", midpoint={self.midpoint}" if self.midpoint is not None else ""
|
| 70 |
+
return f"RangeInput(min={self.min_val}, max={self.max_val}{mid})"
|
comfy_api/latest/_input/video_types.py
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
from abc import ABC, abstractmethod
|
| 3 |
+
from fractions import Fraction
|
| 4 |
+
from typing import Optional, Union, IO
|
| 5 |
+
import io
|
| 6 |
+
import av
|
| 7 |
+
from .._util import VideoContainer, VideoCodec, VideoComponents, normalize_crop_rect
|
| 8 |
+
|
| 9 |
+
class VideoInput(ABC):
|
| 10 |
+
"""
|
| 11 |
+
Abstract base class for video input types.
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
@abstractmethod
|
| 15 |
+
def get_components(self) -> VideoComponents:
|
| 16 |
+
"""
|
| 17 |
+
Abstract method to get the video components (images, audio, and frame rate).
|
| 18 |
+
|
| 19 |
+
Returns:
|
| 20 |
+
VideoComponents containing images, audio, and frame rate
|
| 21 |
+
"""
|
| 22 |
+
pass
|
| 23 |
+
|
| 24 |
+
@abstractmethod
|
| 25 |
+
def save_to(
|
| 26 |
+
self,
|
| 27 |
+
path: Union[str, IO[bytes]],
|
| 28 |
+
format: VideoContainer = VideoContainer.AUTO,
|
| 29 |
+
codec: VideoCodec = VideoCodec.AUTO,
|
| 30 |
+
metadata: Optional[dict] = None,
|
| 31 |
+
bit_depth: int | None = None,
|
| 32 |
+
crf: float | None = None,
|
| 33 |
+
color_space: str | None = None,
|
| 34 |
+
preset: str | None = None,
|
| 35 |
+
):
|
| 36 |
+
"""
|
| 37 |
+
Abstract method to save the video input to a file.
|
| 38 |
+
|
| 39 |
+
bit_depth selects the encoded bit depth; None keeps the video's native depth.
|
| 40 |
+
crf selects the H.264 or AV1 constant rate factor; None uses the encoder default.
|
| 41 |
+
preset selects the H.264 encoder speed/compression trade-off (e.g. "ultrafast");
|
| 42 |
+
None uses the encoder default. Ignored for other codecs.
|
| 43 |
+
color_space="sRGB" selects SDR BT.709/sRGB, "HDR" selects BT.2020/HLG, and "HDR PQ"
|
| 44 |
+
selects BT.2020/PQ. Bit depth is selected independently.
|
| 45 |
+
Tensor-created videos default to sRGB when color_space is None. Loaded videos keep matching recognized native color
|
| 46 |
+
properties; other input pixels must already use the selected color space.
|
| 47 |
+
"""
|
| 48 |
+
pass
|
| 49 |
+
|
| 50 |
+
def get_color_space(self) -> str:
|
| 51 |
+
"""Return the video's color space as sRGB, HDR, HDR PQ, or auto when unspecified."""
|
| 52 |
+
return "auto"
|
| 53 |
+
|
| 54 |
+
@abstractmethod
|
| 55 |
+
def as_trimmed(
|
| 56 |
+
self,
|
| 57 |
+
start_time: float | None = None,
|
| 58 |
+
duration: float | None = None,
|
| 59 |
+
strict_duration: bool = False,
|
| 60 |
+
) -> VideoInput | None:
|
| 61 |
+
"""
|
| 62 |
+
Create a new VideoInput which is trimmed to have the corresponding start_time and duration
|
| 63 |
+
|
| 64 |
+
Returns:
|
| 65 |
+
A new VideoInput, or None if the result would have negative duration
|
| 66 |
+
"""
|
| 67 |
+
pass
|
| 68 |
+
|
| 69 |
+
def as_cropped(
|
| 70 |
+
self,
|
| 71 |
+
x: int = 0,
|
| 72 |
+
y: int = 0,
|
| 73 |
+
width: int = 0,
|
| 74 |
+
height: int = 0,
|
| 75 |
+
) -> VideoInput:
|
| 76 |
+
"""
|
| 77 |
+
Create a new VideoInput spatially cropped to the given pixel rectangle.
|
| 78 |
+
|
| 79 |
+
The rectangle is clamped to the frame and even-aligned for encoder
|
| 80 |
+
compatibility. An empty or full-frame rectangle returns the input
|
| 81 |
+
unchanged.
|
| 82 |
+
|
| 83 |
+
Default implementation materializes the video via get_components();
|
| 84 |
+
subclasses should override with lazier strategies when possible.
|
| 85 |
+
"""
|
| 86 |
+
components = self.get_components()
|
| 87 |
+
rect = normalize_crop_rect(
|
| 88 |
+
x, y, width, height, components.images.shape[2], components.images.shape[1]
|
| 89 |
+
)
|
| 90 |
+
if rect is None:
|
| 91 |
+
return self
|
| 92 |
+
from .._input_impl.video_types import VideoFromComponents
|
| 93 |
+
|
| 94 |
+
cx, cy, cw, ch = rect
|
| 95 |
+
return VideoFromComponents(
|
| 96 |
+
VideoComponents(
|
| 97 |
+
images=components.images[:, cy:cy + ch, cx:cx + cw, :].clone(),
|
| 98 |
+
audio=components.audio,
|
| 99 |
+
frame_rate=components.frame_rate,
|
| 100 |
+
metadata=components.metadata,
|
| 101 |
+
alpha=components.alpha[:, cy:cy + ch, cx:cx + cw].clone()
|
| 102 |
+
if components.alpha is not None
|
| 103 |
+
else None,
|
| 104 |
+
),
|
| 105 |
+
bit_depth=self.get_bit_depth(),
|
| 106 |
+
)
|
| 107 |
+
|
| 108 |
+
def get_stream_source(self) -> Union[str, io.BytesIO]:
|
| 109 |
+
"""
|
| 110 |
+
Get a streamable source for the video. This allows processing without
|
| 111 |
+
loading the entire video into memory.
|
| 112 |
+
|
| 113 |
+
Returns:
|
| 114 |
+
Either a file path (str) or a BytesIO object that can be opened with av.
|
| 115 |
+
|
| 116 |
+
Default implementation creates a BytesIO buffer, but subclasses should
|
| 117 |
+
override this for better performance when possible.
|
| 118 |
+
"""
|
| 119 |
+
buffer = io.BytesIO()
|
| 120 |
+
self.save_to(buffer)
|
| 121 |
+
buffer.seek(0)
|
| 122 |
+
return buffer
|
| 123 |
+
|
| 124 |
+
def get_active_trim_window(self) -> tuple[float, float]:
|
| 125 |
+
"""Return the active trim as ``(start_time, duration)`` in seconds (start_time normalized
|
| 126 |
+
to ``>= 0``; ``duration == 0`` means "until the end"). Default: no trim; trimmable subclasses override.
|
| 127 |
+
"""
|
| 128 |
+
return 0.0, 0.0
|
| 129 |
+
|
| 130 |
+
# Provide a default implementation, but subclasses can provide optimized versions
|
| 131 |
+
# if possible.
|
| 132 |
+
def get_dimensions(self) -> tuple[int, int]:
|
| 133 |
+
"""
|
| 134 |
+
Returns the dimensions of the video input.
|
| 135 |
+
|
| 136 |
+
Returns:
|
| 137 |
+
Tuple of (width, height)
|
| 138 |
+
"""
|
| 139 |
+
components = self.get_components()
|
| 140 |
+
return components.images.shape[2], components.images.shape[1]
|
| 141 |
+
|
| 142 |
+
def get_bit_depth(self) -> int:
|
| 143 |
+
"""
|
| 144 |
+
Returns the bit depth of the video (e.g. 8 or 10).
|
| 145 |
+
|
| 146 |
+
Default implementation returns 8; subclasses report their real depth.
|
| 147 |
+
"""
|
| 148 |
+
return 8
|
| 149 |
+
|
| 150 |
+
def get_duration(self) -> float:
|
| 151 |
+
"""
|
| 152 |
+
Returns the duration of the video in seconds.
|
| 153 |
+
|
| 154 |
+
Returns:
|
| 155 |
+
Duration in seconds
|
| 156 |
+
"""
|
| 157 |
+
components = self.get_components()
|
| 158 |
+
frame_count = components.images.shape[0]
|
| 159 |
+
return float(frame_count / components.frame_rate)
|
| 160 |
+
|
| 161 |
+
def get_frame_count(self) -> int:
|
| 162 |
+
"""
|
| 163 |
+
Returns the number of frames in the video.
|
| 164 |
+
|
| 165 |
+
Default implementation uses :meth:`get_components`, which may require
|
| 166 |
+
loading all frames into memory. File-based implementations should
|
| 167 |
+
override this method and use container/stream metadata instead.
|
| 168 |
+
|
| 169 |
+
Returns:
|
| 170 |
+
Total number of frames as an integer.
|
| 171 |
+
"""
|
| 172 |
+
return int(self.get_components().images.shape[0])
|
| 173 |
+
|
| 174 |
+
def get_frame_rate(self) -> Fraction:
|
| 175 |
+
"""
|
| 176 |
+
Returns the frame rate of the video.
|
| 177 |
+
|
| 178 |
+
Default implementation materializes the video into memory via
|
| 179 |
+
`get_components()`. Subclasses that can inspect the underlying
|
| 180 |
+
container (e.g. `VideoFromFile`) should override this with a more
|
| 181 |
+
efficient implementation.
|
| 182 |
+
|
| 183 |
+
Returns:
|
| 184 |
+
Frame rate as a Fraction.
|
| 185 |
+
"""
|
| 186 |
+
return self.get_components().frame_rate
|
| 187 |
+
|
| 188 |
+
def get_container_format(self) -> str:
|
| 189 |
+
"""
|
| 190 |
+
Returns the container format of the video (e.g., 'mp4', 'mov', 'avi').
|
| 191 |
+
|
| 192 |
+
Returns:
|
| 193 |
+
Container format as string
|
| 194 |
+
"""
|
| 195 |
+
# Default implementation - subclasses should override for better performance
|
| 196 |
+
source = self.get_stream_source()
|
| 197 |
+
with av.open(source, mode="r") as container:
|
| 198 |
+
return container.format.name
|
comfy_api/latest/_input_impl/__init__.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .video_types import VideoFromFile, VideoFromComponents
|
| 2 |
+
|
| 3 |
+
__all__ = [
|
| 4 |
+
# Implementations
|
| 5 |
+
"VideoFromFile",
|
| 6 |
+
"VideoFromComponents",
|
| 7 |
+
]
|