mirror of
https://github.com/markuryy/comfyui-conditioning-converter.git
synced 2026-06-25 18:11:02 +00:00
229 lines
8.1 KiB
Python
229 lines
8.1 KiB
Python
import torch
|
||
import torch.nn.functional as F
|
||
|
||
import comfy.sd1_clip
|
||
import comfy.model_management
|
||
|
||
|
||
PRESETS = {
|
||
"custom": None,
|
||
"klein_9b_to_krea2": (3, 4096, 12, 2560),
|
||
"klein_4b_to_krea2": (3, 2560, 12, 2560),
|
||
"krea2_to_klein_9b": (12, 2560, 3, 4096),
|
||
"flux2_mistral_to_krea2": (3, 5120, 12, 2560),
|
||
}
|
||
|
||
METHODS = ["interpolate", "repeat_nearest", "truncate_pad"]
|
||
|
||
# (layer_indices, target_hidden_dim or 0 for "keep native")
|
||
ENCODE_PRESETS = {
|
||
"klein_for_krea2": ([2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35], 2560),
|
||
"krea2_12_layers": ([2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35], 0),
|
||
"klein_3_layers": ([9, 18, 27], 0),
|
||
"flux2_3_layers": ([10, 20, 30], 0),
|
||
"uniform_6_layers": ([5, 11, 17, 23, 29, 35], 0),
|
||
"deep_3_layers": ([24, 30, 35], 0),
|
||
"custom": ([], 0),
|
||
}
|
||
|
||
|
||
def _resize_dim(tensor, dim, old_size, new_size, method):
|
||
if method == "truncate_pad":
|
||
if new_size <= old_size:
|
||
slices = [slice(None)] * tensor.ndim
|
||
slices[dim] = slice(0, new_size)
|
||
return tensor[tuple(slices)]
|
||
pad_size = new_size - old_size
|
||
pad_shape = list(tensor.shape)
|
||
pad_shape[dim] = pad_size
|
||
padding = torch.zeros(pad_shape, dtype=tensor.dtype, device=tensor.device)
|
||
return torch.cat([tensor, padding], dim=dim)
|
||
|
||
if method == "repeat_nearest":
|
||
repeats = -(-new_size // old_size)
|
||
tensor = tensor.repeat_interleave(repeats, dim=dim)
|
||
slices = [slice(None)] * tensor.ndim
|
||
slices[dim] = slice(0, new_size)
|
||
return tensor[tuple(slices)]
|
||
|
||
perm = list(range(tensor.ndim))
|
||
perm.remove(dim)
|
||
perm.append(dim)
|
||
t = tensor.permute(*perm)
|
||
original_shape = t.shape
|
||
t = t.reshape(-1, 1, old_size)
|
||
t = F.interpolate(t, size=new_size, mode="linear", align_corners=False)
|
||
new_shape = list(original_shape[:-1]) + [new_size]
|
||
t = t.reshape(*new_shape)
|
||
|
||
inv_perm = [0] * tensor.ndim
|
||
for i, p in enumerate(perm):
|
||
inv_perm[p] = i
|
||
return t.permute(*inv_perm)
|
||
|
||
|
||
class ConditioningConverter:
|
||
@classmethod
|
||
def INPUT_TYPES(s):
|
||
return {
|
||
"required": {
|
||
"conditioning": ("CONDITIONING",),
|
||
"preset": (list(PRESETS.keys()),),
|
||
"source_layers": ("INT", {"default": 3, "min": 1, "max": 64}),
|
||
"source_hidden_dim": ("INT", {"default": 4096, "min": 1, "max": 65536}),
|
||
"target_layers": ("INT", {"default": 12, "min": 1, "max": 64}),
|
||
"target_hidden_dim": ("INT", {"default": 2560, "min": 1, "max": 65536}),
|
||
"method": (METHODS,),
|
||
}
|
||
}
|
||
|
||
RETURN_TYPES = ("CONDITIONING",)
|
||
FUNCTION = "convert"
|
||
CATEGORY = "model/conditioning/transform"
|
||
|
||
def convert(self, conditioning, preset, source_layers, source_hidden_dim,
|
||
target_layers, target_hidden_dim, method):
|
||
if preset != "custom":
|
||
source_layers, source_hidden_dim, target_layers, target_hidden_dim = PRESETS[preset]
|
||
|
||
source_features = source_layers * source_hidden_dim
|
||
target_features = target_layers * target_hidden_dim
|
||
|
||
if source_features == target_features:
|
||
return (conditioning,)
|
||
|
||
out = []
|
||
for t in conditioning:
|
||
cond_tensor = t[0]
|
||
meta = t[1].copy()
|
||
|
||
B, seq, features = cond_tensor.shape
|
||
if features != source_features:
|
||
raise ValueError(
|
||
f"ConditioningConverter: expected feature dim {source_features} "
|
||
f"({source_layers}×{source_hidden_dim}) but got {features}. "
|
||
f"Check source_layers and source_hidden_dim."
|
||
)
|
||
|
||
tensor = cond_tensor.reshape(B, seq, source_layers, source_hidden_dim)
|
||
|
||
if source_hidden_dim != target_hidden_dim:
|
||
tensor = _resize_dim(tensor, 3, source_hidden_dim, target_hidden_dim, method)
|
||
|
||
if source_layers != target_layers:
|
||
tensor = _resize_dim(tensor, 2, source_layers, target_layers, method)
|
||
|
||
tensor = tensor.reshape(B, seq, target_features)
|
||
out.append([tensor, meta])
|
||
|
||
return (out,)
|
||
|
||
|
||
class FlexibleCLIPTextEncode:
|
||
"""Encode text with a custom layer tap list, bypassing the model's
|
||
hardcoded layer reshape. Produces conditioning with real hidden
|
||
states from arbitrary transformer layers."""
|
||
|
||
@classmethod
|
||
def INPUT_TYPES(s):
|
||
return {
|
||
"required": {
|
||
"clip": ("CLIP",),
|
||
"text": ("STRING", {"multiline": True, "dynamicPrompts": True}),
|
||
"preset": (list(ENCODE_PRESETS.keys()),),
|
||
"custom_layers": ("STRING", {
|
||
"default": "2,5,8,11,14,17,20,23,26,29,32,35",
|
||
"tooltip": "Comma-separated layer indices (used when preset is 'custom')",
|
||
}),
|
||
"target_hidden_dim": ("INT", {
|
||
"default": 0, "min": 0, "max": 65536,
|
||
"tooltip": "Project per-layer hidden dim to this size. 0 = use preset value or keep native.",
|
||
}),
|
||
"projection_method": (["interpolate", "truncate_pad"],),
|
||
}
|
||
}
|
||
|
||
RETURN_TYPES = ("CONDITIONING",)
|
||
FUNCTION = "encode"
|
||
CATEGORY = "model/conditioning/transform"
|
||
|
||
def encode(self, clip, text, preset, custom_layers,
|
||
target_hidden_dim, projection_method):
|
||
preset_layers, preset_hidden = ENCODE_PRESETS[preset]
|
||
|
||
if preset == "custom":
|
||
layer_list = [int(x.strip()) for x in custom_layers.split(",") if x.strip()]
|
||
else:
|
||
layer_list = list(preset_layers)
|
||
|
||
if target_hidden_dim == 0:
|
||
target_hidden_dim = preset_hidden
|
||
|
||
if not layer_list:
|
||
raise ValueError("FlexibleCLIPTextEncode: layer list is empty.")
|
||
|
||
clip = clip.clone()
|
||
cond_stage = clip.cond_stage_model
|
||
|
||
if not isinstance(cond_stage, comfy.sd1_clip.SD1ClipModel):
|
||
raise TypeError(
|
||
f"FlexibleCLIPTextEncode: expected an SD1ClipModel-based encoder, "
|
||
f"got {type(cond_stage).__name__}. This node works with single-encoder "
|
||
f"models (Klein, Flux2, Krea2)."
|
||
)
|
||
|
||
inner_clip = getattr(cond_stage, cond_stage.clip)
|
||
original_layer = inner_clip.layer
|
||
|
||
tokens = clip.tokenize(text)
|
||
clip.load_model(tokens)
|
||
device = clip.patcher.load_device
|
||
inner_clip.execution_device = device
|
||
|
||
inner_clip.layer = layer_list
|
||
try:
|
||
with comfy.model_management.cuda_device_context(device):
|
||
out = comfy.sd1_clip.SD1ClipModel.encode_token_weights(cond_stage, tokens)
|
||
finally:
|
||
inner_clip.layer = original_layer
|
||
inner_clip.execution_device = None
|
||
|
||
z = out[0]
|
||
pooled = out[1] if len(out) > 1 else None
|
||
extra = out[2] if len(out) > 2 else {}
|
||
|
||
if z.ndim == 4:
|
||
B, n_layers, seq, hidden = z.shape
|
||
z = z.permute(0, 2, 1, 3) # (B, seq, n_layers, hidden)
|
||
|
||
if target_hidden_dim > 0 and target_hidden_dim != hidden:
|
||
z = _resize_dim(z, 3, hidden, target_hidden_dim, projection_method)
|
||
hidden = target_hidden_dim
|
||
|
||
z = z.reshape(B, seq, n_layers * hidden)
|
||
elif z.ndim == 3 and target_hidden_dim > 0:
|
||
B, seq, hidden = z.shape
|
||
if target_hidden_dim != hidden:
|
||
z = z.unsqueeze(2)
|
||
z = _resize_dim(z, 3, hidden, target_hidden_dim, projection_method)
|
||
z = z.squeeze(2)
|
||
|
||
pooled_dict = {}
|
||
if pooled is not None:
|
||
pooled_dict["pooled_output"] = pooled
|
||
if isinstance(extra, dict) and "attention_mask" in extra:
|
||
pooled_dict["attention_mask"] = extra["attention_mask"]
|
||
|
||
return ([[z, pooled_dict]],)
|
||
|
||
|
||
NODE_CLASS_MAPPINGS = {
|
||
"ConditioningConverter": ConditioningConverter,
|
||
"FlexibleCLIPTextEncode": FlexibleCLIPTextEncode,
|
||
}
|
||
|
||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||
"ConditioningConverter": "Conditioning Converter (Experimental)",
|
||
"FlexibleCLIPTextEncode": "Flexible CLIP Text Encode (Experimental)",
|
||
}
|