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)", }