This commit is contained in:
LAOGOU-666
2025-12-24 22:16:42 +08:00
parent c6f9558202
commit ce13df1006
6 changed files with 1301 additions and 0 deletions
+47
View File
@@ -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"]
+369
View File
@@ -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)",
}
+87
View File
@@ -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",
}
+104
View File
@@ -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 🎚️",
}
+211
View File
@@ -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",
}
+483
View File
@@ -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;
};
}
}
});