update
This commit is contained in:
+47
@@ -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"]
|
||||
@@ -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)",
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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 🎚️",
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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;
|
||||
};
|
||||
}
|
||||
}
|
||||
});
|
||||
Reference in New Issue
Block a user