Iceyy4400 commited on
Commit
8ee8fb6
·
verified ·
1 Parent(s): cb0905a

Vendor ComfyUI + custom nodes, add Gradio app with ZeroGPU + model auto-download (part 3)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +1 -0
  2. comfy/text_encoders/sa_t5.py +22 -0
  3. comfy/text_encoders/sam3_clip.py +97 -0
  4. comfy/text_encoders/sd2_clip.py +23 -0
  5. comfy/text_encoders/sd2_clip_config.json +23 -0
  6. comfy/text_encoders/sd3_clip.py +167 -0
  7. comfy/text_encoders/sensenova.py +149 -0
  8. comfy/text_encoders/spiece_tokenizer.py +59 -0
  9. comfy/text_encoders/t5.py +249 -0
  10. comfy/text_encoders/t5_config_base.json +22 -0
  11. comfy/text_encoders/t5_config_xxl.json +22 -0
  12. comfy/text_encoders/t5_old_config_xxl.json +22 -0
  13. comfy/text_encoders/t5_pile_config_xl.json +22 -0
  14. comfy/text_encoders/t5_pile_tokenizer/tokenizer.model +3 -0
  15. comfy/text_encoders/t5_tokenizer/special_tokens_map.json +125 -0
  16. comfy/text_encoders/t5_tokenizer/tokenizer.json +0 -0
  17. comfy/text_encoders/t5_tokenizer/tokenizer_config.json +939 -0
  18. comfy/text_encoders/umt5_config_base.json +22 -0
  19. comfy/text_encoders/umt5_config_xxl.json +22 -0
  20. comfy/text_encoders/wan.py +37 -0
  21. comfy/text_encoders/z_image.py +46 -0
  22. comfy/utils.py +1535 -0
  23. comfy/weight_adapter/__init__.py +42 -0
  24. comfy/weight_adapter/base.py +396 -0
  25. comfy/weight_adapter/boft.py +218 -0
  26. comfy/weight_adapter/bypass.py +441 -0
  27. comfy/weight_adapter/glora.py +290 -0
  28. comfy/weight_adapter/loha.py +378 -0
  29. comfy/weight_adapter/lokr.py +481 -0
  30. comfy/weight_adapter/lora.py +368 -0
  31. comfy/weight_adapter/oft.py +327 -0
  32. comfy_api/feature_flags.py +166 -0
  33. comfy_api/generate_api_stubs.py +86 -0
  34. comfy_api/input/__init__.py +26 -0
  35. comfy_api/input/basic_types.py +14 -0
  36. comfy_api/input/video_types.py +6 -0
  37. comfy_api/input_impl/__init__.py +7 -0
  38. comfy_api/input_impl/video_types.py +2 -0
  39. comfy_api/internal/__init__.py +150 -0
  40. comfy_api/internal/api_registry.py +39 -0
  41. comfy_api/internal/async_to_sync.py +1002 -0
  42. comfy_api/internal/singleton.py +33 -0
  43. comfy_api/latest/__init__.py +177 -0
  44. comfy_api/latest/_caching.py +42 -0
  45. comfy_api/latest/_input/__init__.py +17 -0
  46. comfy_api/latest/_input/basic_types.py +42 -0
  47. comfy_api/latest/_input/curve_types.py +219 -0
  48. comfy_api/latest/_input/range_types.py +70 -0
  49. comfy_api/latest/_input/video_types.py +198 -0
  50. 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
+ ]