diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..0009878 --- /dev/null +++ b/__init__.py @@ -0,0 +1,47 @@ + +import importlib.util +import os +import sys +import json + +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} +WEB_DIRECTORY = "web" +python = sys.executable + +def get_ext_dir(subpath=None, mkdir=False): + dir = os.path.dirname(__file__) + if subpath is not None: + dir = os.path.join(dir, subpath) + + dir = os.path.abspath(dir) + + if mkdir and not os.path.exists(dir): + os.makedirs(dir) + return dir + +def serialize(obj): + if isinstance(obj, (str, int, float, bool, list, dict, type(None))): + return obj + return str(obj) + + +py = get_ext_dir("py") +files = os.listdir(py) +all_nodes = {} +for file in files: + if not file.endswith(".py"): + continue + name = os.path.splitext(file)[0] + imported_module = importlib.import_module(".py.{}".format(name), __name__) + try: + NODE_CLASS_MAPPINGS = {**NODE_CLASS_MAPPINGS, **imported_module.NODE_CLASS_MAPPINGS} + NODE_DISPLAY_NAME_MAPPINGS = {**NODE_DISPLAY_NAME_MAPPINGS, **imported_module.NODE_DISPLAY_NAME_MAPPINGS} + serialized_CLASS_MAPPINGS = {k: serialize(v) for k, v in imported_module.NODE_CLASS_MAPPINGS.items()} + serialized_DISPLAY_NAME_MAPPINGS = {k: serialize(v) for k, v in imported_module.NODE_DISPLAY_NAME_MAPPINGS.items()} + all_nodes[file]={"NODE_CLASS_MAPPINGS": serialized_CLASS_MAPPINGS, "NODE_DISPLAY_NAME_MAPPINGS": serialized_DISPLAY_NAME_MAPPINGS} + except: + pass + + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] diff --git a/py/noise_injection.py b/py/noise_injection.py new file mode 100644 index 0000000..71db31d --- /dev/null +++ b/py/noise_injection.py @@ -0,0 +1,369 @@ +""" +ZImage Noise Injection 节点 + +通过 CFG 机制将参考图像的特征注入到生成中。 +可以让生成结果"学习"参考图像的某些特质(如水珠、纹理、质感等)。 +""" + +import torch +import torch.nn.functional as F +import logging +import comfy.model_management +import comfy.latent_formats + + +class LGNoiseInjection: + """ + LG Noise Injection - 特征注入 + + 工作原理: + 参考图像的 latent 特征会被注入到 CFG 过程中, + 模型会"学习"参考图像的某些特质并应用到生成结果上。 + + 适用场景: + - 添加水珠、汗珠等表面细节 + - 添加纹理、材质感 + - 添加光泽、反射效果 + """ + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL", { + "tooltip": "模型" + }), + "vae": ("VAE", { + "tooltip": "VAE 编码器" + }), + "reference_image": ("IMAGE", { + "tooltip": "参考图像(含有你想要注入的特征)" + }), + "strength": ("FLOAT", { + "default": 0.15, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "注入强度。0.1-0.2 轻微,0.2-0.4 明显" + }), + "start_percent": ("FLOAT", { + "default": 0.0, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "开始注入的采样进度" + }), + "end_percent": ("FLOAT", { + "default": 0.6, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "结束注入的采样进度" + }), + }, + "optional": { + "mask": ("MASK", { + "tooltip": "遮罩,白色区域会被注入特征" + }), + }, + } + + RETURN_TYPES = ("MODEL",) + FUNCTION = "apply" + CATEGORY = "advanced/model" + DESCRIPTION = "将参考图像的特征(如水珠、纹理等)注入到生成结果中。" + + def apply(self, model, vae, reference_image, strength, start_percent, end_percent, mask=None): + if strength <= 0: + return (model,) + + m = model.clone() + + # 编码参考图像 + ref_latent = self._encode_reference(vae, reference_image) + + # 预处理 mask + mask_latent = None + if mask is not None: + mask_latent = self._prepare_mask(mask) + + step_counter = [0] + + def cfg_function(args): + """在 CFG 合并后注入参考特征""" + cond = args["cond"] + uncond = args["uncond"] + cond_scale = args["cond_scale"] + timestep = args["timestep"] + + step_counter[0] += 1 + + # 标准 CFG + cfg_result = uncond + cond_scale * (cond - uncond) + + # 计算进度 + sigma = float(timestep[0]) if timestep.dim() > 0 else float(timestep) + progress = 1.0 - min(sigma, 1.0) + + # 检查是否在作用范围内 + if progress < start_percent or progress > end_percent: + return cfg_result + + # 准备 reference + ref = ref_latent.to(device=cfg_result.device, dtype=cfg_result.dtype) + + if ref.shape[2:] != cfg_result.shape[2:]: + ref = F.interpolate(ref, size=cfg_result.shape[2:], mode='bilinear', align_corners=False) + + if ref.shape[0] != cfg_result.shape[0]: + if ref.shape[0] == 1: + ref = ref.expand(cfg_result.shape[0], -1, -1, -1) + else: + ref = ref[:cfg_result.shape[0]] + + # 准备 mask + current_mask = None + if mask_latent is not None: + current_mask = mask_latent.to(device=cfg_result.device, dtype=cfg_result.dtype) + # 调整 mask 尺寸到 latent 空间 + if current_mask.shape[2:] != cfg_result.shape[2:]: + current_mask = F.interpolate(current_mask, size=cfg_result.shape[2:], mode='bilinear', align_corners=False) + # 调整 batch size + if current_mask.shape[0] != cfg_result.shape[0]: + if current_mask.shape[0] == 1: + current_mask = current_mask.expand(cfg_result.shape[0], -1, -1, -1) + else: + current_mask = current_mask[:cfg_result.shape[0]] + + # 计算有效强度(线性衰减) + if end_percent > start_percent: + range_progress = (progress - start_percent) / (end_percent - start_percent) + decay = 1.0 - range_progress + else: + decay = 1.0 + effective_strength = strength * decay + + # 计算特征注入方向 + feature_direction = ref - cfg_result + + # 控制注入幅度,避免过度偏移 + cfg_std = cfg_result.std() + feature_std = feature_direction.std() + if feature_std > cfg_std * 3: + feature_direction = feature_direction * (cfg_std * 3 / feature_std) + + # 应用遮罩 + if current_mask is not None: + # mask: 1 = 注入区域, 0 = 保持原样 + feature_direction = feature_direction * current_mask + + # 应用特征注入 + injected = cfg_result + feature_direction * effective_strength + + if step_counter[0] <= 3: + mask_info = "with mask" if current_mask is not None else "no mask" + logging.warning(f"[FeatureInj] step={step_counter[0]} | progress={progress:.2f} | eff_str={effective_strength:.3f} | {mask_info}") + + return injected + + m.set_model_sampler_cfg_function(cfg_function) + + return (m,) + + def _encode_reference(self, vae, reference_image): + """编码参考图像""" + loaded_models = comfy.model_management.loaded_models(only_currently_used=True) + latent = vae.encode(reference_image) + latent = comfy.latent_formats.Flux().process_in(latent) + comfy.model_management.load_models_gpu(loaded_models) + logging.warning(f"[FeatureInj] Reference encoded: shape={latent.shape}") + return latent + + def _prepare_mask(self, mask): + """准备遮罩,转换为 latent 空间格式""" + # mask 输入格式: [B, H, W] 或 [H, W] + if mask.dim() == 2: + mask = mask.unsqueeze(0) # [H, W] -> [1, H, W] + + # 添加 channel 维度: [B, H, W] -> [B, 1, H, W] + mask = mask.unsqueeze(1) + + logging.warning(f"[FeatureInj] Mask prepared: shape={mask.shape}") + return mask + + +class LGNoiseInjectionLatent: + """ + LG Noise Injection (Latent) - 直接使用 Latent 的特征注入 + + 工作原理: + 直接输入参考 latent,其特征会被注入到 CFG 过程中。 + 如果 latent 包含 noise_mask,则自动使用该遮罩。 + + 适用场景: + - 添加水珠、汗珠等表面细节 + - 添加纹理、材质感 + - 添加光泽、反射效果 + """ + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL", { + "tooltip": "模型" + }), + "reference_latent": ("LATENT", { + "tooltip": "参考 latent(含有你想要注入的特征)" + }), + "strength": ("FLOAT", { + "default": 0.15, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "注入强度。0.1-0.2 轻微,0.2-0.4 明显" + }), + "start_percent": ("FLOAT", { + "default": 0.0, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "开始注入的采样进度" + }), + "end_percent": ("FLOAT", { + "default": 0.6, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "结束注入的采样进度" + }), + }, + } + + RETURN_TYPES = ("MODEL",) + FUNCTION = "apply" + CATEGORY = "advanced/model" + DESCRIPTION = "直接输入 latent 进行特征注入,自动使用 latent 的 noise_mask 作为遮罩。" + + def apply(self, model, reference_latent, strength, start_percent, end_percent): + if strength <= 0: + return (model,) + + m = model.clone() + + # 获取 latent samples + ref_latent = reference_latent["samples"] + + # 获取 noise_mask(如果存在) + mask_latent = None + if "noise_mask" in reference_latent: + mask_latent = self._prepare_mask(reference_latent["noise_mask"]) + logging.warning(f"[FeatureInjLatent] Using noise_mask from latent") + + logging.warning(f"[FeatureInjLatent] Reference latent: shape={ref_latent.shape}") + + step_counter = [0] + + def cfg_function(args): + """在 CFG 合并后注入参考特征""" + cond = args["cond"] + uncond = args["uncond"] + cond_scale = args["cond_scale"] + timestep = args["timestep"] + + step_counter[0] += 1 + + # 标准 CFG + cfg_result = uncond + cond_scale * (cond - uncond) + + # 计算进度 + sigma = float(timestep[0]) if timestep.dim() > 0 else float(timestep) + progress = 1.0 - min(sigma, 1.0) + + # 检查是否在作用范围内 + if progress < start_percent or progress > end_percent: + return cfg_result + + # 准备 reference + ref = ref_latent.to(device=cfg_result.device, dtype=cfg_result.dtype) + + if ref.shape[2:] != cfg_result.shape[2:]: + ref = F.interpolate(ref, size=cfg_result.shape[2:], mode='bilinear', align_corners=False) + + if ref.shape[0] != cfg_result.shape[0]: + if ref.shape[0] == 1: + ref = ref.expand(cfg_result.shape[0], -1, -1, -1) + else: + ref = ref[:cfg_result.shape[0]] + + # 准备 mask + current_mask = None + if mask_latent is not None: + current_mask = mask_latent.to(device=cfg_result.device, dtype=cfg_result.dtype) + # 调整 mask 尺寸到 latent 空间 + if current_mask.shape[2:] != cfg_result.shape[2:]: + current_mask = F.interpolate(current_mask, size=cfg_result.shape[2:], mode='bilinear', align_corners=False) + # 调整 batch size + if current_mask.shape[0] != cfg_result.shape[0]: + if current_mask.shape[0] == 1: + current_mask = current_mask.expand(cfg_result.shape[0], -1, -1, -1) + else: + current_mask = current_mask[:cfg_result.shape[0]] + + # 计算有效强度(线性衰减) + if end_percent > start_percent: + range_progress = (progress - start_percent) / (end_percent - start_percent) + decay = 1.0 - range_progress + else: + decay = 1.0 + effective_strength = strength * decay + + # 计算特征注入方向 + feature_direction = ref - cfg_result + + # 控制注入幅度,避免过度偏移 + cfg_std = cfg_result.std() + feature_std = feature_direction.std() + if feature_std > cfg_std * 3: + feature_direction = feature_direction * (cfg_std * 3 / feature_std) + + # 应用遮罩 + if current_mask is not None: + # mask: 1 = 注入区域, 0 = 保持原样 + feature_direction = feature_direction * current_mask + + # 应用特征注入 + injected = cfg_result + feature_direction * effective_strength + + if step_counter[0] <= 3: + mask_info = "with noise_mask" if current_mask is not None else "no mask" + logging.warning(f"[FeatureInjLatent] step={step_counter[0]} | progress={progress:.2f} | eff_str={effective_strength:.3f} | {mask_info}") + + return injected + + m.set_model_sampler_cfg_function(cfg_function) + + return (m,) + + def _prepare_mask(self, mask): + """准备遮罩,转换为 latent 空间格式""" + # mask 输入格式: [B, H, W] 或 [H, W] + if mask.dim() == 2: + mask = mask.unsqueeze(0) # [H, W] -> [1, H, W] + + # 添加 channel 维度: [B, H, W] -> [B, 1, H, W] + if mask.dim() == 3: + mask = mask.unsqueeze(1) + + return mask + + +NODE_CLASS_MAPPINGS = { + "LGNoiseInjection": LGNoiseInjection, + "LGNoiseInjectionLatent": LGNoiseInjectionLatent, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LGNoiseInjection": "🎈LG Noise Injection", + "LGNoiseInjectionLatent": "🎈LG Noise Injection (Latent)", +} diff --git a/py/sampling_offset.py b/py/sampling_offset.py new file mode 100644 index 0000000..42626b6 --- /dev/null +++ b/py/sampling_offset.py @@ -0,0 +1,87 @@ +""" +ZImage Model Sampling 节点 + +用于调整 ZImage/Lumina2 模型的采样参数。 +关键区别:ZImage 使用 multiplier=1.0,而 SD3 使用 multiplier=1000。 +""" + +import torch +import logging +import comfy.model_sampling + + +class ModelSamplingZImage: + """ + ZImage/Lumina2 采样参数调整节点 + + 与 ModelSamplingSD3 的关键区别: + - ZImage 使用 multiplier=1.0(timestep 范围 0-1) + - SD3 使用 multiplier=1000(timestep 范围 0-1000) + + shift 参数控制噪声调度的偏移(通过 time_snr_shift 函数): + - shift=1.0: 无偏移,线性调度 + - shift>1.0: 向高噪声偏移,早期步骤更激进 + - ZImage 默认 shift=3.0 + """ + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "shift": ("FLOAT", { + "default": 3.0, + "min": 0.0, + "max": 100.0, + "step": 0.01, + "tooltip": "噪声调度偏移量。ZImage 默认 3.0。shift=1.0 为线性,>1.0 向高噪声偏移" + }), + "multiplier": ("FLOAT", { + "default": 1.0, + "min": 0.001, + "max": 10000.0, + "step": 0.001, + "tooltip": "Timestep 乘数。ZImage/AuraFlow=1.0,SD3/Flux=1000" + }), + } + } + + RETURN_TYPES = ("MODEL",) + FUNCTION = "patch" + CATEGORY = "advanced/model" + DESCRIPTION = "调整 ZImage/Lumina2 模型的采样参数。可设置 shift 和 multiplier。" + + def patch(self, model, shift, multiplier): + m = model.clone() + + sampling_base = comfy.model_sampling.ModelSamplingDiscreteFlow + sampling_type = comfy.model_sampling.CONST + + class ModelSamplingAdvanced(sampling_base, sampling_type): + pass + + model_sampling = ModelSamplingAdvanced(model.model.model_config) + model_sampling.set_parameters(shift=shift, multiplier=multiplier) + m.add_object_patch("model_sampling", model_sampling) + + # 调试输出 + logging.info(f"[ModelSamplingZImage] Applied: shift={shift}, multiplier={multiplier}") + logging.info(f"[ModelSamplingZImage] sigma_min={model_sampling.sigma_min:.6f}, sigma_max={model_sampling.sigma_max:.6f}") + logging.info(f"[ModelSamplingZImage] sigmas[0:5]={model_sampling.sigmas[:5].tolist()}") + + # 验证 patch 是否成功 + patched_ms = m.get_model_object("model_sampling") + logging.info(f"[ModelSamplingZImage] Patch verification: patched shift={patched_ms.shift}, multiplier={patched_ms.multiplier}") + logging.info(f"[ModelSamplingZImage] Patched sigmas[0:5]={patched_ms.sigmas[:5].tolist()}") + + return (m,) + + +# 注册节点 +NODE_CLASS_MAPPINGS = { + "ModelSamplingZImage": ModelSamplingZImage, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "ModelSamplingZImage": "Model Sampling ZImage", +} diff --git a/py/sigmas_editor.py b/py/sigmas_editor.py new file mode 100644 index 0000000..7c6c0ca --- /dev/null +++ b/py/sigmas_editor.py @@ -0,0 +1,104 @@ +""" +Interactive Sigmas Editor Node +Allows real-time adjustment of sigmas curve by dragging points +""" + +import torch +import numpy as np +import json +import os +import folder_paths +from server import PromptServer +from aiohttp import web + +# Set web directory for custom UI +WEB_DIRECTORY = "./web" + + +class SigmasEditor: + """Interactive editor for adjusting sigmas curve""" + + # 类级别缓存,存储每个节点上次接收的输入sigmas(用于检测输入是否变化) + _last_sent_data = {} + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "sigmas": ("SIGMAS", {"tooltip": "Input sigmas schedule to edit"}), + "sigmas_adjustments": ("STRING", { + "default": "[]", + "multiline": False, + "dynamicPrompts": False, + "tooltip": "JSON array of adjusted sigma values for each step" + }), + }, + "hidden": { + "unique_id": "UNIQUE_ID", + }, + } + + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("adjusted_sigmas",) + FUNCTION = "adjust_sigmas" + CATEGORY = "sampling/custom_sampling/sigmas" + DESCRIPTION = "Interactively adjust sigmas curve by dragging control points" + + def adjust_sigmas(self, sigmas, sigmas_adjustments="[]", unique_id=None): + # Convert sigmas to numpy + if isinstance(sigmas, torch.Tensor): + sigmas_np = sigmas.cpu().numpy() + else: + sigmas_np = np.array(sigmas) + + # Parse adjusted sigma values from JSON + try: + adjusted_values = json.loads(sigmas_adjustments) + except: + adjusted_values = [] + + # If no adjustments or length mismatch, use original sigmas + if len(adjusted_values) != len(sigmas_np): + adjusted_sigmas = sigmas_np.copy() + else: + # Use the adjusted sigma values directly + adjusted_sigmas = np.array(adjusted_values, dtype=np.float64) + + # Ensure last sigma is still 0 if original was 0 + if sigmas_np[-1] == 0: + adjusted_sigmas[-1] = 0 + + result_tensor = torch.FloatTensor(adjusted_sigmas) + + # Send sigmas data to frontend via PromptServer (只有输入sigmas改变时才发送) + if unique_id is not None: + # 只根据输入的sigmas创建缓存键(不包括adjustments) + current_sigmas_key = tuple(sigmas_np.tolist()) + + # 检查是否与上次输入的sigmas相同 + last_sigmas_key = self._last_sent_data.get(unique_id) + + # 只有输入sigmas改变时才发送数据到前端 + if last_sigmas_key != current_sigmas_key: + PromptServer.instance.send_sync("sigmas_editor_update", { + "node_id": unique_id, + "sigmas_data": { + "original": sigmas_np.tolist(), + "adjusted": adjusted_sigmas.tolist(), + } + }) + + # 更新缓存(只缓存输入的sigmas) + self._last_sent_data[unique_id] = current_sigmas_key + + return (result_tensor,) + + +NODE_CLASS_MAPPINGS = { + "SigmasEditor": SigmasEditor, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "SigmasEditor": "Sigmas Editor 🎚️", +} + diff --git a/py/timestep_noise.py b/py/timestep_noise.py new file mode 100644 index 0000000..431d490 --- /dev/null +++ b/py/timestep_noise.py @@ -0,0 +1,211 @@ +""" +ZImage/Lumina2 采样扰动节点 + +这个节点可以在采样过程中注入噪声扰动, +打破模型的同质化输出,使不同种子能产生更明显的差异。 +""" + +import torch +import torch.nn.functional as F +import logging + + +class ZImageTimestepNoise: + """ + 对 timestep/sigma 添加噪声扰动 + 这会改变模型对当前去噪步骤的感知,产生不同的输出 + + 支持两种模式: + - sigma: 适用于传统扩散模型,使用乘性噪声 + - flow: 适用于 Flow Matching 模型(如 ZImage/Lumina2),使用加性噪声 + """ + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "sigmas": ("SIGMAS",), + "mode": (["sigma", "flow"], { + "default": "flow", + "tooltip": "sigma: 传统扩散模型(乘性噪声); flow: Flow Matching 模型(加性噪声)" + }), + "noise_strength": ("FLOAT", { + "default": 0.05, + "min": 0.0, + "max": 2.0, + "step": 0.01, + "tooltip": "噪声强度" + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 0xffffffffffffffff, + "tooltip": "噪声种子" + }), + "start_percent": ("FLOAT", { + "default": 0.0, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "开始应用噪声的采样进度 (0.0 = 开始)" + }), + "end_percent": ("FLOAT", { + "default": 0.5, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "停止应用噪声的采样进度 (1.0 = 结束)" + }), + }, + "optional": { + "mask": ("MASK",), + } + } + + RETURN_TYPES = ("MODEL",) + FUNCTION = "patch" + CATEGORY = "advanced/model" + DESCRIPTION = "对 timestep 添加噪声扰动,改变模型对去噪步骤的感知。sigma 模式适用于传统扩散模型,flow 模式适用于 Flow Matching 模型(如 ZImage/Lumina2)。可选遮罩限制影响区域。" + + def patch(self, model, sigmas, mode, noise_strength, seed, start_percent, end_percent, mask=None): + m = model.clone() + + if noise_strength <= 0: + return (m,) + + # 将 sigmas 转换为 list 以便查找 + sigma_list = sigmas.tolist() + total_steps = len(sigma_list) - 1 # 最后一个是 0 + + stored_mode = mode + stored_strength = noise_strength + stored_seed = seed + stored_start = start_percent + stored_end = end_percent + stored_sigmas = sigma_list + stored_total_steps = total_steps + stored_mask = mask + step_counter = [0] + + def unet_wrapper(apply_model_func, args): + input_x = args["input"] + timestep = args["timestep"] + c = args["c"] + + step_counter[0] += 1 + current_step = step_counter[0] + + # 获取当前 sigma 值 + sigma_val = float(timestep[0]) if timestep.dim() > 0 else float(timestep) + + # 在 sigma_list 中查找当前 sigma 的位置来确定进度 + min_diff = float('inf') + matched_idx = 0 + for i, s in enumerate(stored_sigmas): + diff = abs(s - sigma_val) + if diff < min_diff: + min_diff = diff + matched_idx = i + + # 计算进度 + progress = matched_idx / stored_total_steps if stored_total_steps > 0 else 0.0 + progress = max(0.0, min(1.0, progress)) + + # 检查是否在指定范围内 + in_range = progress >= stored_start and progress <= stored_end + + # 准备遮罩(如果有) + latent_mask = None + if stored_mask is not None: + mask_tensor = stored_mask + if mask_tensor.dim() == 2: + mask_tensor = mask_tensor.unsqueeze(0) + latent_h, latent_w = input_x.shape[2], input_x.shape[3] + latent_mask = F.interpolate( + mask_tensor.unsqueeze(1), + size=(latent_h, latent_w), + mode='bilinear', + align_corners=False + ) + if latent_mask.shape[0] == 1 and input_x.shape[0] > 1: + latent_mask = latent_mask.expand(input_x.shape[0], -1, -1, -1) + latent_mask = latent_mask.to(device=input_x.device, dtype=input_x.dtype) + + # 计算噪声(无论是否在范围内都计算,用于调试显示) + generator = torch.Generator(device=timestep.device) + generator.manual_seed(stored_seed + current_step) + + t_orig = float(timestep[0]) if timestep.dim() > 0 else float(timestep) + + if stored_mode == "sigma": + noise_factor = 1.0 + (torch.rand(1, generator=generator, device=timestep.device).item() - 0.5) * stored_strength + noisy_timestep = timestep * noise_factor + else: # flow 模式 + # 计算可用空间 + headroom_up = 1.0 - t_orig # 向上的空间 + headroom_down = t_orig - 0.0 # 向下的空间 + + # 确定实际可用的噪声范围 + actual_up = min(stored_strength, headroom_up) + actual_down = min(stored_strength, headroom_down) + + # 生成 [0, 1] 随机数,然后映射到 [-actual_down, +actual_up] + raw = torch.rand(1, generator=generator, device=timestep.device).item() + total_range = actual_down + actual_up + + if total_range > 1e-6: + # 映射到可用范围:raw=0 -> -actual_down, raw=1 -> +actual_up + actual_delta = raw * total_range - actual_down + else: + actual_delta = 0.0 + + noisy_timestep = timestep + actual_delta + # 安全 clamp + noisy_timestep = torch.clamp(noisy_timestep, 0.0, 1.0) + + # 获取数值用于日志 + t_noisy = float(noisy_timestep[0]) if noisy_timestep.dim() > 0 else float(noisy_timestep) + delta = t_noisy - t_orig + + # 调试输出 + if current_step <= 5: + status = "✓ APPLIED" if in_range else "✗ SKIPPED" + has_mask = "mask=YES" if latent_mask is not None else "mask=NO" + if stored_mode == "flow": + logging.info(f"[ZImageTimestepNoise] step={current_step}/{stored_total_steps} | progress={progress:.2f} | range=[{stored_start:.2f}, {stored_end:.2f}] | {status}") + logging.info(f"[ZImageTimestepNoise] timestep: {t_orig:.6f} -> {t_noisy:.6f} (delta={delta:+.6f}) | headroom=[↓{headroom_down:.3f}, ↑{headroom_up:.3f}] | {has_mask}") + else: + logging.info(f"[ZImageTimestepNoise] step={current_step}/{stored_total_steps} | progress={progress:.2f} | range=[{stored_start:.2f}, {stored_end:.2f}] | {status}") + logging.info(f"[ZImageTimestepNoise] timestep: {t_orig:.6f} -> {t_noisy:.6f} (delta={delta:+.6f}) | factor={noise_factor:.4f} | {has_mask}") + + if not in_range: + # 不在范围内,使用原始 timestep + return apply_model_func(input_x, timestep, **c) + + # 在范围内,应用噪声 + if latent_mask is not None: + # 计算两种 timestep 下的结果并混合 + result_original = apply_model_func(input_x, timestep, **c) + result_noisy = apply_model_func(input_x, noisy_timestep, **c) + result = result_original * (1 - latent_mask) + result_noisy * latent_mask + return result + else: + # 无遮罩,直接使用扰动后的 timestep + return apply_model_func(input_x, noisy_timestep, **c) + + return apply_model_func(input_x, timestep, **c) + + m.set_model_unet_function_wrapper(unet_wrapper) + + return (m,) + + +# 注册节点 +NODE_CLASS_MAPPINGS = { + "ZImageTimestepNoise": ZImageTimestepNoise, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "ZImageTimestepNoise": "ZImage Timestep Noise", +} diff --git a/web/sigmas_editor.js b/web/sigmas_editor.js new file mode 100644 index 0000000..0c86e69 --- /dev/null +++ b/web/sigmas_editor.js @@ -0,0 +1,483 @@ +import { app } from "../../scripts/app.js"; +import { api } from "../../scripts/api.js"; + +app.registerExtension({ + name: "SigmasEditor.Interactive", + async beforeRegisterNodeDef(nodeType, nodeData) { + if (nodeData.name === "SigmasEditor") { + + // 扩展节点的构造函数 + const onNodeCreated = nodeType.prototype.onNodeCreated; + nodeType.prototype.onNodeCreated = function() { + const result = onNodeCreated?.apply(this, arguments); + + // 初始化节点数据 + this.sigmas_data = null; + this.adjustments = []; // 存储实际调整后的sigma值 + this.dragging_point = -1; + this.isAdjusting = false; + + // 允许调整节点大小 + this.resizable = true; + + // 设置初始大小(宽度自动,高度400) + this.size = [this.size[0] || 600, 400]; + + // 设置WebSocket监听 + this.setupWebSocket(); + + return result; + }; + + // 添加WebSocket设置方法 + nodeType.prototype.setupWebSocket = function() { + const messageHandler = (event) => { + const data = event.detail; + + if (!data || !data.node_id || !data.sigmas_data) { + return; + } + + // 通过node_id查找对应的节点 + const targetNode = app.graph.getNodeById(parseInt(data.node_id)); + + // 检查是否是当前节点 + if (targetNode && targetNode === this) { + this.sigmas_data = data.sigmas_data.original; + // 从后端接收调整后的值,如果为空则使用原始值 + this.adjustments = data.sigmas_data.adjusted || data.sigmas_data.original.slice(); + + // 第一次执行时,触发 sigmas_adjustments 组件更新 + this.updateAdjustmentsWidget(); + + // 立即更新画布(强制重新计算尺寸) + if (this.canvas) { + this.updateCanvas(true); + } else { + // 如果画布还没创建,等待一下再尝试 + setTimeout(() => { + if (this.canvas) { + this.updateCanvas(true); + } + }, 100); + } + this.setDirtyCanvas(true, true); + } + }; + + api.addEventListener("sigmas_editor_update", messageHandler); + + // 存储handler引用以便后续清理 + this._sigmasEditorMessageHandler = messageHandler; + }; + + // 添加节点时的处理 + const onAdded = nodeType.prototype.onAdded; + nodeType.prototype.onAdded = function() { + const result = onAdded?.apply(this, arguments); + + if (!this.canvasContainer && this.id !== undefined && this.id !== -1) { + // 创建画布容器(高度由节点自适应) + const container = document.createElement("div"); + container.style.position = "relative"; + container.style.width = "100%"; + container.style.height = "100%"; + container.style.minHeight = "300px"; + container.style.backgroundColor = "#1e1e1e"; + container.style.borderRadius = "8px"; + container.style.overflow = "hidden"; + + // 创建画布 + const canvas = document.createElement("canvas"); + canvas.style.width = "100%"; + canvas.style.height = "100%"; + canvas.style.cursor = "crosshair"; + + container.appendChild(canvas); + this.canvas = canvas; + this.canvasContainer = container; + + // 添加鼠标事件监听 + this.addCanvasEventListeners(); + + // 添加DOM组件 + this.widgets ||= []; + this.widgets_up = true; + + requestAnimationFrame(() => { + if (this.widgets) { + this.canvasWidget = this.addDOMWidget("sigmas_canvas", "canvas", container); + + // 初始化时强制计算画布尺寸 + this.updateCanvas(true); + this.setDirtyCanvas(true, true); + } + }); + } + + return result; + }; + + // 添加画布事件监听器 + nodeType.prototype.addCanvasEventListeners = function() { + const canvas = this.canvas; + + // 获取缩放后的鼠标坐标 + const getScaledMousePos = (e) => { + const rect = canvas.getBoundingClientRect(); + // 计算缩放比例 + const scaleX = canvas.width / rect.width; + const scaleY = canvas.height / rect.height; + + // 鼠标相对于canvas显示区域的位置 + const displayX = e.clientX - rect.left; + const displayY = e.clientY - rect.top; + + // 转换为实际canvas坐标 + const canvasX = displayX * scaleX; + const canvasY = displayY * scaleY; + + return { x: canvasX, y: canvasY }; + }; + + // 鼠标按下 + canvas.addEventListener("mousedown", (e) => { + const pos = getScaledMousePos(e); + const pointIdx = this.findNearestPoint(pos.x, pos.y); + + if (pointIdx !== -1) { + this.dragging_point = pointIdx; + this.isAdjusting = true; + this.updateCanvas(); // 只重绘,不resize + } + }); + + // 鼠标移动 + canvas.addEventListener("mousemove", (e) => { + if (this.dragging_point !== -1 && this.isAdjusting) { + const pos = getScaledMousePos(e); + this.updatePointAdjustment(this.dragging_point, pos.y); + this.updateAdjustmentsWidget(); // 实时同步更新组件值 + this.updateCanvas(); // 只重绘,不resize + } + }); + + // 鼠标释放 + canvas.addEventListener("mouseup", (e) => { + if (this.dragging_point !== -1) { + this.dragging_point = -1; + this.isAdjusting = false; + + // 更新adjustments widget + this.updateAdjustmentsWidget(); + this.updateCanvas(); // 只重绘,不resize + } + }); + + // 鼠标离开 + canvas.addEventListener("mouseleave", (e) => { + if (this.dragging_point !== -1) { + this.dragging_point = -1; + this.isAdjusting = false; + this.updateAdjustmentsWidget(); + this.updateCanvas(); // 只重绘,不resize + } + }); + }; + + // 查找最近的点 + nodeType.prototype.findNearestPoint = function(mouseX, mouseY) { + if (!this.sigmas_data || this.sigmas_data.length === 0) return -1; + + const paddingLeft = 60; + const paddingTop = 50; + const paddingRight = 20; + const paddingBottom = 50; + + const canvas = this.canvas; + const chartWidth = canvas.width - paddingLeft - paddingRight; + const chartHeight = canvas.height - paddingTop - paddingBottom; + const chartX = paddingLeft; + const chartY = paddingTop; + + const steps = this.sigmas_data.length; + + let closestDist = Infinity; + let closestIdx = -1; + + for (let i = 0; i < steps; i++) { + const x = chartX + (chartWidth / (steps - 1)) * i; + // adjustments存储的是实际调整后的sigma值 + const adjustedValue = this.adjustments[i] !== undefined ? this.adjustments[i] : this.sigmas_data[i]; + const y = chartY + chartHeight - (adjustedValue * chartHeight); + + const dist = Math.sqrt(Math.pow(mouseX - x, 2) + Math.pow(mouseY - y, 2)); + if (dist < 15 && dist < closestDist) { + closestDist = dist; + closestIdx = i; + } + } + + return closestIdx; + }; + + // 更新点的调整值 + nodeType.prototype.updatePointAdjustment = function(pointIdx, mouseY) { + if (!this.sigmas_data || pointIdx < 0 || pointIdx >= this.sigmas_data.length) return; + + const paddingTop = 50; + const paddingBottom = 50; + + const canvas = this.canvas; + const chartHeight = canvas.height - paddingTop - paddingBottom; + const chartY = paddingTop; + + // 计算新的sigma值(范围0-1) + const clampedY = Math.max(chartY, Math.min(chartY + chartHeight, mouseY)); + const newSigmaValue = (chartY + chartHeight - clampedY) / chartHeight; + + // 直接存储调整后的sigma值,限制在0-1范围内 + this.adjustments[pointIdx] = Math.max(0.0, Math.min(1.0, newSigmaValue)); + }; + + // 向上取整到指定小数位 + nodeType.prototype.ceilToFixed = function(value, decimals) { + const multiplier = Math.pow(10, decimals); + return Math.ceil(value * multiplier) / multiplier; + }; + + // 更新adjustments widget(向上取整到4位小数) + nodeType.prototype.updateAdjustmentsWidget = function() { + const widget = this.widgets?.find(w => w.name === "sigmas_adjustments"); + if (widget) { + // 对所有调整后的值进行向上取整到4位小数,并格式化为字符串保留末尾的0 + const formattedValues = this.adjustments.map(v => { + const rounded = this.ceilToFixed(v, 4); + return rounded.toFixed(4); + }); + // 手动构建 JSON 数组字符串,保留4位小数格式 + widget.value = '[' + formattedValues.join(', ') + ']'; + } + }; + + // 更新画布 + nodeType.prototype.updateCanvas = function(forceResize = false) { + if (!this.canvas) return; + + requestAnimationFrame(() => { + const canvas = this.canvas; + const ctx = canvas.getContext("2d"); + + // 只在必要时重新设置画布尺寸 + if (forceResize || !this._canvasInitialized) { + const rect = canvas.getBoundingClientRect(); + + // 确保画布有合理的尺寸 + const width = rect.width > 0 ? rect.width : 600; + const height = rect.height > 0 ? rect.height : 300; + + // 只有尺寸真正改变时才重新设置 + if (canvas.width !== width || canvas.height !== height) { + canvas.width = width; + canvas.height = height; + } + + this._canvasInitialized = true; + } + + // 增加padding以容纳坐标轴标签 + const paddingLeft = 60; + const paddingRight = 20; + const paddingTop = 50; + const paddingBottom = 50; + + const chartWidth = canvas.width - paddingLeft - paddingRight; + const chartHeight = canvas.height - paddingTop - paddingBottom; + const chartX = paddingLeft; + const chartY = paddingTop; + + // 清空画布 + ctx.fillStyle = "#1e1e1e"; + ctx.fillRect(0, 0, canvas.width, canvas.height); + + // 如果没有数据,显示提示 + if (!this.sigmas_data || this.sigmas_data.length === 0) { + ctx.fillStyle = "#999"; + ctx.font = "14px Arial"; + ctx.textAlign = "center"; + ctx.fillText("Connect Sigmas Input & Execute Workflow", canvas.width / 2, canvas.height / 2); + return; + } + + // 绘制图表区域 + ctx.fillStyle = "#2a2a2a"; + ctx.fillRect(chartX, chartY, chartWidth, chartHeight); + + // 绘制网格 + ctx.strokeStyle = "#444"; + ctx.lineWidth = 1; + for (let i = 0; i <= 10; i++) { + const y = chartY + (chartHeight / 10) * i; + ctx.beginPath(); + ctx.moveTo(chartX, y); + ctx.lineTo(chartX + chartWidth, y); + ctx.stroke(); + } + + const steps = this.sigmas_data.length; + + // 确保adjustments数组长度正确,初始化为原始sigma值 + if (this.adjustments.length !== steps) { + this.adjustments = this.sigmas_data.slice(); + } + + // 绘制Y轴刻度和标签(Sigma值 0-1,刻度0.1) + ctx.fillStyle = "#999"; + ctx.font = "10px Arial"; + ctx.textAlign = "right"; + for (let i = 0; i <= 10; i++) { + const value = 1.0 - (i / 10); + const y = chartY + (chartHeight / 10) * i; + ctx.fillText(value.toFixed(1), chartX - 5, y + 3); + } + + // Y轴标签 + ctx.save(); + ctx.translate(15, chartY + chartHeight / 2); + ctx.rotate(-Math.PI / 2); + ctx.textAlign = "center"; + ctx.font = "12px Arial"; + ctx.fillStyle = "#ccc"; + ctx.fillText("Sigma Value", 0, 0); + ctx.restore(); + + // 绘制X轴刻度和标签 + ctx.fillStyle = "#999"; + ctx.font = "10px Arial"; + ctx.textAlign = "center"; + + // 根据步数决定显示哪些刻度 + let stepInterval = 1; + if (steps > 30) stepInterval = 2; + if (steps > 50) stepInterval = 5; + if (steps > 100) stepInterval = 10; + + for (let i = 0; i < steps; i += stepInterval) { + const x = chartX + (chartWidth / (steps - 1)) * i; + ctx.fillText(i.toString(), x, chartY + chartHeight + 15); + } + // 确保显示最后一步 + if ((steps - 1) % stepInterval !== 0) { + const x = chartX + chartWidth; + ctx.fillText((steps - 1).toString(), x, chartY + chartHeight + 15); + } + + // X轴标签 + ctx.font = "12px Arial"; + ctx.fillStyle = "#ccc"; + ctx.fillText("Steps", chartX + chartWidth / 2, chartY + chartHeight + 30); + + // 绘制曲线 + ctx.strokeStyle = "#4a9eff"; + ctx.lineWidth = 2; + ctx.beginPath(); + + for (let i = 0; i < steps; i++) { + const x = chartX + (chartWidth / (steps - 1)) * i; + // adjustments存储的是实际调整后的sigma值 + const adjustedValue = this.adjustments[i] !== undefined ? this.adjustments[i] : this.sigmas_data[i]; + // 限制在0-1范围内 + const clampedValue = Math.max(0, Math.min(1, adjustedValue)); + const y = chartY + chartHeight - (clampedValue * chartHeight); + + if (i === 0) { + ctx.moveTo(x, y); + } else { + ctx.lineTo(x, y); + } + } + ctx.stroke(); + + // 绘制控制点 + for (let i = 0; i < steps; i++) { + const x = chartX + (chartWidth / (steps - 1)) * i; + const adjustedValue = this.adjustments[i] !== undefined ? this.adjustments[i] : this.sigmas_data[i]; + const clampedValue = Math.max(0, Math.min(1, adjustedValue)); + const y = chartY + chartHeight - (clampedValue * chartHeight); + + ctx.fillStyle = this.dragging_point === i ? "#ff6b6b" : "#4a9eff"; + ctx.beginPath(); + ctx.arc(x, y, 5, 0, Math.PI * 2); + ctx.fill(); + } + + // 绘制标题 + ctx.fillStyle = "#ccc"; + ctx.font = "14px Arial"; + ctx.textAlign = "center"; + ctx.fillText("Sigmas Schedule Editor", canvas.width / 2, 20); + + // 显示当前拖拽点的信息 + if (this.dragging_point !== -1) { + const orig = this.sigmas_data[this.dragging_point]; + const adjusted = this.adjustments[this.dragging_point]; + const multiplier = orig > 0 ? (adjusted / orig) : 1.0; + + // 向上取整到4位小数显示 + const origCeil = this.ceilToFixed(orig, 4); + const adjustedCeil = this.ceilToFixed(adjusted, 4); + + ctx.fillStyle = "#ff6b6b"; + ctx.textAlign = "center"; + ctx.font = "11px Arial"; + ctx.fillText( + `Step ${this.dragging_point}: Original=${origCeil.toFixed(4)}, Adjusted=${adjustedCeil.toFixed(4)}, Multiplier=${multiplier.toFixed(2)}x`, + canvas.width / 2, + 35 + ); + } + }); + }; + + // 监听节点尺寸变化 + const onResize = nodeType.prototype.onResize; + nodeType.prototype.onResize = function(size) { + const result = onResize?.apply(this, arguments); + + // 节点尺寸改变时,强制重新计算画布尺寸 + if (this.canvas) { + this._canvasInitialized = false; // 重置标志 + this.updateCanvas(true); + } + + return result; + }; + + // 节点移除时的处理 + const onRemoved = nodeType.prototype.onRemoved; + nodeType.prototype.onRemoved = function() { + const result = onRemoved?.apply(this, arguments); + + // 清理画布 + if (this && this.canvas) { + const ctx = this.canvas.getContext("2d"); + if (ctx) { + ctx.clearRect(0, 0, this.canvas.width, this.canvas.height); + } + this.canvas = null; + } + if (this) { + this.canvasContainer = null; + } + + // 清理事件监听器 + if (this._sigmasEditorMessageHandler) { + api.removeEventListener("sigmas_editor_update", this._sigmasEditorMessageHandler); + this._sigmasEditorMessageHandler = null; + } + + return result; + }; + } + } +});