feat: add sampling_mode and route the node to the corrected pipeline
The Direct node now dispatches: "standard" (default) runs the new pipeline, "legacy" runs the frozen pre-2.0 one. Measured through the node with a recording sampler, 20 steps: standard, 1 stage 1 call x 20 steps, from sigma 14.61 standard, 2 stages 2 calls x 10 steps, from 14.61 then 1.48 standard, denoise 0.6 1 call x 20 steps, from 2.23 standard, 3 injections 4 calls x 5 steps, 14.61 / 3.87 / 1.48 / 0.60 legacy, 2 stages 2 calls, each rebuilding a full schedule The 1.48 is the fix: stages continue one trajectory instead of restarting at maximum noise, where flow models discard the previous stage entirely. The 2.23 is denoise finally reaching the schedule. Also on the node: - sequential_distribution, injection_distribution and fast_high_channel_noise become real optional inputs. As V1 `hidden` tuple inputs ComfyUI never delivered them, so they were stuck at their defaults. - the debug/visualisation hidden inputs are gone; they drove stub no-ops. - IS_CHANGED is removed: it only restated widget values that are already part of the cache key, and would have rejected the new input. - new widgets are appended last, so saved workflows keep their widget order. web/src/sampling_mode_migration.ts switches nodes loaded from pre-2.0 workflows to "legacy", recognising them by the absence of the snk_version property, so existing seeds keep reproducing. 186 Python tests and 85 web tests pass; the 11 legacy goldens now exercise the legacy branch through this dispatch. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
cc0e1a52b9
commit
3af34338ab
+92
-112
@@ -1,8 +1,10 @@
|
||||
import comfy.sample
|
||||
from .shader_params_reader import get_shader_params, ShaderParamsReader
|
||||
from .shader_noise_ksampler import ShaderNoiseKSampler, get_visualizer, set_debug_level
|
||||
from .pipelines import standard as standard_pipeline
|
||||
|
||||
class DirectShaderNoiseKSampler(ShaderNoiseKSampler):
|
||||
|
||||
class DirectShaderNoiseKSampler(ShaderNoiseKSampler):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
@@ -23,154 +25,126 @@ class DirectShaderNoiseKSampler(ShaderNoiseKSampler):
|
||||
"blend_mode": (["normal", "add", "multiply", "screen", "overlay", "soft_light", "hard_light", "difference"], {"default": "multiply", "tooltip": "Method used to blend the shader noise with the base noise"}),
|
||||
"noise_transform": (["none", "reverse", "inverse", "absolute", "square", "sqrt", "log", "sin", "cos"], {"default": "none", "tooltip": "Apply mathematical transformations to the noise for creative effects"}),
|
||||
"use_temporal_coherence": ("BOOLEAN", {"default": False, "tooltip": "Ensures consistent noise patterns. For sequences like video frames, it helps maintain frame-to-frame consistency (e.g., using the same base seed and 4D noise). For single image generations, it ensures the base noise is derived consistently from the main seed."}),
|
||||
|
||||
|
||||
# New direct shader parameters
|
||||
"shader_type": (["domain_warp", "tensor_field", "curl_noise"], {"default": "domain_warp", "tooltip": "Select the type of shader noise pattern to use in the visualization [different types have different characteristic outputs]"}),
|
||||
"shape_type": (["none", "radial", "linear", "spiral", "checkerboard", "spots", "hexgrid", "stripes", "gradient", "vignette", "cross", "stars", "triangles", "concentric", "rays", "zigzag"], {"default": "none", "tooltip": "Apply a shape mask to the shader noise pattern to create more complex structures [not post processing - is applied to shader noise pattern before rendering]"}),
|
||||
"color_scheme": (["none", "blue_red", "viridis", "plasma", "inferno", "magma", "turbo", "jet", "rainbow", "cool", "hot", "parula", "hsv", "autumn", "winter", "spring", "summer", "copper", "pink", "bone", "ocean", "terrain", "neon", "fire"], {"default": "none", "tooltip": "Choose a color palette to apply to the shader noise visualization [not post processing - is applied to the shader noise pattern before rendering]"}),
|
||||
"noise_scale": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.001, "tooltip": "Adjust the scale of the shader noise pattern - lower values create larger, zoomed-in features; higher values create smaller, zoomed-out features [small value shifts can lead to larger variations]"}),
|
||||
"octaves": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 8.0, "step": 0.1, "tooltip": "Number of shader noise layers to combine - higher values add more detail and complexity [small value shifts can lead to larger variations]"}),
|
||||
"octaves": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 8.0, "step": 0.1, "tooltip": "Number of shader noise layers to combine - higher values add more detail and complexity. Fractional values blend between two layer counts (standard sampling only)"}),
|
||||
"warp_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 5.0, "step": 0.001, "tooltip": "Control how much the shader noise pattern warps and distorts - higher values create more swirling or complex transformations [small adjustments are good for subtle variations]"}),
|
||||
"shape_mask_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.0001, "tooltip": "Adjust the intensity of the shape mask\'s effect on the shader noise pattern - higher values make the shape more prominent [small adjustments are good for subtle variations - not effective without shape mask]"}),
|
||||
"phase_shift": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 2.0, "step": 0.0001, "tooltip": "Shift the phase of the shader noise pattern to create different variations or animate patterns over time [small adjustments are good for subtle variations]"}),
|
||||
"color_intensity": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 1.0, "step": 0.0001, "tooltip": "Adjust the intensity of the color scheme application - lower values are more desaturated, higher values are more vibrant [small adjustments are good for subtle variations - not effective without color scheme]"}),
|
||||
},
|
||||
# Appended after the required widgets on purpose: saved workflows map
|
||||
# widget values by position, so new widgets must come last.
|
||||
"optional": {
|
||||
"custom_sigmas": ("SIGMAS", {"tooltip": "Optional custom sigma schedule to override the model's default schedule"}),
|
||||
},
|
||||
"hidden": {
|
||||
"debug_level": (["0-Off", "1-Basic", "2-Detailed", "3-Verbose"], {"default": "0-Off", "tooltip": "Enable debugging at specified level to understand what's happening during shader generation and sampling."}),
|
||||
"fast_high_channel_noise": ("BOOLEAN", {"default": False, "tooltip": "Use a faster, simplified noise generation method for models with many channels (>16), like LTXV."}),
|
||||
"sampling_mode": (["standard", "legacy"], {"default": "standard", "tooltip": "standard: stages are segments of one sampling run, denoise and custom sigmas are honoured, and blended noise keeps the distribution the model expects. legacy: the pre-2.0 behaviour, kept so older workflows reproduce their seeds."}),
|
||||
"sequential_distribution": (["uniform", "linear_decrease", "linear_increase", "gaussian", "first_stronger", "last_stronger"], {"default": "linear_decrease", "tooltip": "How shader strength is distributed across sequential stages"}),
|
||||
"injection_distribution": (["uniform", "linear_decrease", "linear_increase", "gaussian", "first_stronger", "last_stronger"], {"default": "linear_decrease", "tooltip": "How shader strength is distributed across injection stages"}),
|
||||
"denoise_visualization_frequency": (["Every step", "25% intervals", "10% intervals", "4 steps", "2 steps"], {"default": "Every step", "tooltip": "How often to save images during the denoising process. Higher frequency means more images but slower generation."}),
|
||||
"target_attribute_changes": ("STRING", {"forceInput": True, "tooltip": "Connect output from ParameterResponseMapperNode here"}),
|
||||
}
|
||||
"fast_high_channel_noise": ("BOOLEAN", {"default": False, "tooltip": "Use a faster, simplified noise generation method for models with many channels (>16), like LTXV"}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
@classmethod
|
||||
def REGISTER_MATRIX_BUTTON(s):
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image,
|
||||
denoise=1.0, sequential_stages=2, injection_stages=3, shader_strength=0.3, blend_mode="multiply",
|
||||
noise_transform="none", sequential_distribution="linear_decrease", injection_distribution="linear_decrease",
|
||||
use_temporal_coherence=False, debug_level="0-Off", fast_high_channel_noise=False,
|
||||
denoise_visualization_frequency="25% intervals", custom_sigmas=None, target_attribute_changes="",
|
||||
shader_type="domain_warp", shape_type="none", color_scheme="none", noise_scale=1.0, octaves=1,
|
||||
warp_strength=0.5, shape_mask_strength=1.0, phase_shift=0.5, color_intensity=0.8):
|
||||
# Ensure all direct parameters also trigger re-execution when changed
|
||||
return (seed, steps, cfg, sampler_name, scheduler, denoise, sequential_stages,
|
||||
injection_stages, shader_strength, blend_mode, noise_transform,
|
||||
sequential_distribution, injection_distribution, use_temporal_coherence,
|
||||
debug_level, denoise_visualization_frequency, custom_sigmas,
|
||||
target_attribute_changes, fast_high_channel_noise,
|
||||
# Include direct shader parameters
|
||||
shader_type, shape_type, color_scheme, noise_scale, octaves,
|
||||
warp_strength, shape_mask_strength, phase_shift, color_intensity)
|
||||
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image,
|
||||
denoise=1.0, sequential_stages=1, injection_stages=0, shader_strength=0.3, blend_mode="multiply",
|
||||
noise_transform="none", use_temporal_coherence=False,
|
||||
shader_type="domain_warp", shape_type="none", color_scheme="none", noise_scale=1.0, octaves=1.0,
|
||||
warp_strength=0.5, shape_mask_strength=1.0, phase_shift=0.5, color_intensity=0.8,
|
||||
sampling_mode="standard", sequential_distribution="linear_decrease",
|
||||
injection_distribution="linear_decrease", fast_high_channel_noise=False, custom_sigmas=None,
|
||||
# Accepted for the legacy path and for older callers; not exposed as inputs.
|
||||
debug_level="0-Off", denoise_visualization_frequency="25% intervals", target_attribute_changes=""):
|
||||
"""Run the shader noise sampler with direct parameter inputs."""
|
||||
debugger = set_debug_level(int(debug_level.split("-")[0]))
|
||||
get_visualizer()
|
||||
|
||||
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image,
|
||||
denoise=1.0, sequential_stages=2, injection_stages=3, shader_strength=0.3, blend_mode="multiply",
|
||||
noise_transform="none", sequential_distribution="linear_decrease", injection_distribution="linear_decrease",
|
||||
use_temporal_coherence=False, debug_level="0-Off", fast_high_channel_noise=False,
|
||||
denoise_visualization_frequency="25% intervals", custom_sigmas=None, target_attribute_changes="",
|
||||
shader_type="tensor_field", shape_type="none", color_scheme="none", noise_scale=1.0, octaves=3.0,
|
||||
warp_strength=0.5, shape_mask_strength=1.0, phase_shift=0.0, color_intensity=0.8):
|
||||
"""
|
||||
Run the multi-stage shader noise k-sampler with direct parameter inputs
|
||||
"""
|
||||
# Parse debug level from the selected option
|
||||
debug_level_value = int(debug_level.split("-")[0])
|
||||
|
||||
# Set the debug level in the shader debugger
|
||||
debugger = set_debug_level(debug_level_value)
|
||||
|
||||
# Get the visualizer
|
||||
visualizer = get_visualizer()
|
||||
|
||||
# Get device early from latent_image
|
||||
device = latent_image["samples"].device
|
||||
if debugger.enabled:
|
||||
print(f"ℹ️ Using device: {device}")
|
||||
|
||||
# Get the shader parameters, but we'll modify them with our direct inputs
|
||||
# Start from the saved params file, then override with this node's inputs.
|
||||
shader_params = get_shader_params()
|
||||
|
||||
# Override shader parameters with direct inputs - ensure all naming variants are set
|
||||
# Shader Type (set all variants)
|
||||
|
||||
# Every generator reads a different spelling of these, so set all variants.
|
||||
shader_params["shader_type"] = shader_type
|
||||
shader_params["shaderType"] = shader_type
|
||||
|
||||
# Shape Type (set all variants)
|
||||
|
||||
shader_params["shape_type"] = shape_type
|
||||
shader_params["shaderShapeType"] = shape_type
|
||||
|
||||
# Color Scheme
|
||||
|
||||
shader_params["colorScheme"] = color_scheme
|
||||
shader_params["color_scheme"] = color_scheme
|
||||
|
||||
# Noise Scale (set all variants)
|
||||
|
||||
shader_params["scale"] = noise_scale
|
||||
shader_params["shaderScale"] = noise_scale
|
||||
|
||||
# Octaves (set all variants)
|
||||
shader_params["octaves"] = float(octaves) # Ensure float for compatibility
|
||||
|
||||
shader_params["octaves"] = float(octaves)
|
||||
shader_params["shaderOctaves"] = float(octaves)
|
||||
|
||||
# Warp Strength (set all variants)
|
||||
|
||||
shader_params["warp_strength"] = warp_strength
|
||||
shader_params["shaderWarpStrength"] = warp_strength
|
||||
|
||||
# Shape Mask Strength (set all variants)
|
||||
|
||||
shader_params["shapemaskstrength"] = shape_mask_strength
|
||||
shader_params["shaderShapeStrength"] = shape_mask_strength
|
||||
shader_params["shapeMaskStrength"] = shape_mask_strength # Added capital M version
|
||||
shader_params["shape_mask_strength"] = shape_mask_strength # Added underscore version
|
||||
shader_params["shape_strength"] = shape_mask_strength # Added alternative name checked in shader_to_tensor.py
|
||||
|
||||
# Phase Shift (set all variants)
|
||||
shader_params["shapeMaskStrength"] = shape_mask_strength
|
||||
shader_params["shape_mask_strength"] = shape_mask_strength
|
||||
shader_params["shape_strength"] = shape_mask_strength
|
||||
|
||||
shader_params["phase_shift"] = phase_shift
|
||||
shader_params["shaderPhaseShift"] = phase_shift
|
||||
|
||||
# Color Intensity (set all variants)
|
||||
|
||||
shader_params["intensity"] = color_intensity
|
||||
shader_params["shaderColorIntensity"] = color_intensity
|
||||
|
||||
# Other parameters
|
||||
shader_params["time"] = shader_params.get("time", 0.0) # Keep existing time or set to 0
|
||||
shader_params["base_seed"] = seed # Set base seed for temporal coherence
|
||||
|
||||
shader_params["time"] = shader_params.get("time", 0.0)
|
||||
shader_params["base_seed"] = seed
|
||||
shader_params["useTemporalCoherence"] = use_temporal_coherence
|
||||
shader_params["temporal_coherence"] = use_temporal_coherence # Add alternative name
|
||||
shader_params["temporal_coherence"] = use_temporal_coherence
|
||||
shader_params["fast_high_channel_noise"] = fast_high_channel_noise
|
||||
|
||||
# Visualization type (default to 3/ellipses as in the default parameters)
|
||||
shader_params["visualization_type"] = shader_params.get("visualization_type", 3)
|
||||
|
||||
# Apply security validation and sanitization to the overridden parameters
|
||||
# This prevents DoS attacks (e.g. excessive octaves) and ensures parameter safety
|
||||
# Clamp octaves, seeds and enum values before they reach noise generation.
|
||||
shader_params = ShaderParamsReader.validate_and_sanitize_params(shader_params)
|
||||
# Sanitising truncates octaves to an integer; the standard pipeline
|
||||
# interpolates between integer renders, so keep the requested value.
|
||||
shader_params["octaves"] = float(octaves)
|
||||
|
||||
if debugger.enabled:
|
||||
# Debug output for direct parameters
|
||||
print(f"🔧 Direct Shader Parameters:")
|
||||
print(f" Shader Type: {shader_type}")
|
||||
print(f" Shape Type: {shape_type}")
|
||||
print(f" Color Scheme: {color_scheme}")
|
||||
print(f" Noise Scale: {noise_scale}")
|
||||
print(f" Octaves: {octaves}")
|
||||
print(f" Warp Strength: {warp_strength}")
|
||||
print(f" Shape Mask Strength: {shape_mask_strength}")
|
||||
print(f" Phase Shift: {phase_shift}")
|
||||
print(f" Color Intensity: {color_intensity}")
|
||||
|
||||
# Call parent class sample method with the modified shader params
|
||||
# Use the parent class implementation from ShaderNoiseKSampler
|
||||
result = super().sample(
|
||||
print(f"🔧 Direct shader parameters: type={shader_type} shape={shape_type} colour={color_scheme} "
|
||||
f"scale={noise_scale} octaves={octaves} warp={warp_strength} phase={phase_shift}")
|
||||
|
||||
if sampling_mode == "legacy":
|
||||
return super().sample(
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
sequential_stages=sequential_stages,
|
||||
injection_stages=injection_stages,
|
||||
shader_strength=shader_strength,
|
||||
blend_mode=blend_mode,
|
||||
noise_transform=noise_transform,
|
||||
sequential_distribution=sequential_distribution,
|
||||
injection_distribution=injection_distribution,
|
||||
use_temporal_coherence=use_temporal_coherence,
|
||||
debug_level=debug_level,
|
||||
fast_high_channel_noise=fast_high_channel_noise,
|
||||
denoise_visualization_frequency=denoise_visualization_frequency,
|
||||
custom_sigmas=custom_sigmas,
|
||||
target_attribute_changes=target_attribute_changes,
|
||||
shader_params_override=shader_params,
|
||||
)
|
||||
|
||||
result = standard_pipeline.run(
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
@@ -179,22 +153,28 @@ class DirectShaderNoiseKSampler(ShaderNoiseKSampler):
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
latent_image=latent_image,
|
||||
latent=latent_image,
|
||||
denoise=denoise,
|
||||
sequential_stages=sequential_stages,
|
||||
injection_stages=injection_stages,
|
||||
shader_strength=shader_strength,
|
||||
blend_mode=blend_mode,
|
||||
noise_transform=noise_transform,
|
||||
shader_params=shader_params,
|
||||
shader_type=shader_type,
|
||||
sequential_distribution=sequential_distribution,
|
||||
injection_distribution=injection_distribution,
|
||||
use_temporal_coherence=use_temporal_coherence,
|
||||
debug_level=debug_level,
|
||||
fast_high_channel_noise=fast_high_channel_noise,
|
||||
denoise_visualization_frequency=denoise_visualization_frequency,
|
||||
custom_sigmas=custom_sigmas,
|
||||
target_attribute_changes=target_attribute_changes,
|
||||
shader_params_override=shader_params, # Pass our modified shader params to parent
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
shader_info = {
|
||||
"shader_type": shader_type,
|
||||
"shader_strength": shader_strength,
|
||||
"sequential_stages": sequential_stages,
|
||||
"injection_stages": injection_stages,
|
||||
"blend_mode": blend_mode,
|
||||
"noise_transform": noise_transform,
|
||||
"sampling_mode": sampling_mode,
|
||||
}
|
||||
return {"ui": {"images": [], "shader_info": shader_info}, "result": (result,)}
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
"""
|
||||
The node routes to the right pipeline, and standard mode samples one trajectory.
|
||||
|
||||
Measured through the node with a recording sampler:
|
||||
|
||||
standard, 1 stage 1 call x 20 steps, from sigma 14.61
|
||||
standard, 2 stages 2 calls x 10 steps, from 14.61 then 1.48
|
||||
standard, denoise 0.6 1 call x 20 steps, from 2.23
|
||||
standard, 3 injections 4 calls x 5 steps, 14.61 / 3.87 / 1.48 / 0.60
|
||||
legacy, 2 stages 2 calls, each building its own full schedule
|
||||
|
||||
The second sigma in the two-stage standard run is the point: legacy restarted
|
||||
every stage at 14.61, which on flow models discards the previous stage entirely.
|
||||
"""
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import comfy.sample
|
||||
from helpers import FakeModel
|
||||
from snk.direct_shader_ksampler import DirectShaderNoiseKSampler
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sampler_calls():
|
||||
calls = []
|
||||
|
||||
def fake_sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative,
|
||||
latent_image, denoise=1.0, disable_noise=False, start_step=None, last_step=None,
|
||||
force_full_denoise=False, noise_mask=None, sigmas=None, callback=None,
|
||||
disable_pbar=False, seed=None):
|
||||
calls.append({
|
||||
"steps": steps,
|
||||
"sigmas": None if sigmas is None else sigmas.detach().clone(),
|
||||
"denoise": denoise,
|
||||
})
|
||||
result = latent_image + 0.1 * noise
|
||||
if callback is not None:
|
||||
callback(max(steps - 1, 0), result * 0.5, result, steps)
|
||||
return result
|
||||
|
||||
with mock.patch.object(comfy.sample, "sample", fake_sample):
|
||||
yield calls
|
||||
|
||||
|
||||
def run_node(**overrides):
|
||||
kwargs = dict(
|
||||
model=FakeModel("eps"), seed=8888, steps=20, cfg=7.0, sampler_name="euler",
|
||||
scheduler="normal", positive=[], negative=[],
|
||||
latent_image={"samples": torch.zeros(1, 4, 16, 16)}, denoise=1.0,
|
||||
sequential_stages=1, injection_stages=0, shader_strength=0.3, blend_mode="multiply",
|
||||
noise_transform="none", use_temporal_coherence=False, shader_type="domain_warp",
|
||||
shape_type="none", color_scheme="none", noise_scale=1.0, octaves=1.0,
|
||||
warp_strength=0.5, shape_mask_strength=1.0, phase_shift=0.5, color_intensity=0.8,
|
||||
)
|
||||
kwargs.update(overrides)
|
||||
return DirectShaderNoiseKSampler().sample(**kwargs)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ["standard", "legacy"])
|
||||
def test_both_modes_return_a_latent(sampler_calls, mode):
|
||||
result = run_node(sampling_mode=mode)
|
||||
assert "result" in result and isinstance(result["result"], tuple)
|
||||
assert result["result"][0]["samples"].shape == (1, 4, 16, 16)
|
||||
|
||||
|
||||
def test_standard_mode_samples_one_schedule(sampler_calls):
|
||||
run_node(sampling_mode="standard", sequential_stages=2)
|
||||
|
||||
assert len(sampler_calls) == 2
|
||||
first, second = sampler_calls
|
||||
assert first["steps"] == second["steps"] == 10
|
||||
assert first["sigmas"] is not None
|
||||
# The second segment continues where the first stopped instead of restarting.
|
||||
assert float(second["sigmas"][0]) < float(first["sigmas"][0])
|
||||
assert torch.equal(first["sigmas"][-1], second["sigmas"][0])
|
||||
|
||||
|
||||
def test_legacy_mode_still_builds_its_own_schedule(sampler_calls):
|
||||
"""The frozen path passes no sigmas, so KSampler rebuilds a full one per stage."""
|
||||
run_node(sampling_mode="legacy", sequential_stages=2)
|
||||
|
||||
assert len(sampler_calls) == 2
|
||||
assert all(call["sigmas"] is None for call in sampler_calls)
|
||||
|
||||
|
||||
def test_denoise_only_reaches_the_schedule_in_standard_mode(sampler_calls):
|
||||
run_node(sampling_mode="standard", denoise=1.0)
|
||||
full_start = float(sampler_calls[0]["sigmas"][0])
|
||||
|
||||
sampler_calls.clear()
|
||||
run_node(sampling_mode="standard", denoise=0.6)
|
||||
assert float(sampler_calls[0]["sigmas"][0]) < full_start
|
||||
|
||||
|
||||
def test_injection_stages_never_leave_a_one_step_segment(sampler_calls):
|
||||
run_node(sampling_mode="standard", injection_stages=3)
|
||||
|
||||
assert sum(call["steps"] for call in sampler_calls) == 20
|
||||
assert all(call["steps"] >= 2 for call in sampler_calls)
|
||||
|
||||
|
||||
def test_standard_mode_is_the_default(sampler_calls):
|
||||
run_node(sequential_stages=2)
|
||||
assert all(call["sigmas"] is not None for call in sampler_calls)
|
||||
@@ -0,0 +1,49 @@
|
||||
// Import app from ComfyUI at runtime (this import is resolved by the browser)
|
||||
// @ts-ignore - ComfyUI provides this at runtime
|
||||
import { app } from '../../scripts/app.js';
|
||||
const NODE_NAME = 'ShaderNoiseKSamplerDirect';
|
||||
const SAMPLING_MODE_WIDGET = 'sampling_mode';
|
||||
const VERSION_PROPERTY = 'snk_version';
|
||||
const CURRENT_VERSION = 2;
|
||||
/** True when this serialized node predates the sampling modes. */
|
||||
export function needsLegacySampling(info) {
|
||||
const properties = info?.properties;
|
||||
return properties?.[VERSION_PROPERTY] === undefined;
|
||||
}
|
||||
/** Set the sampling_mode widget, if the node has one. */
|
||||
export function setSamplingMode(node, mode) {
|
||||
const widget = node.widgets?.find((w) => w.name === SAMPLING_MODE_WIDGET);
|
||||
if (widget)
|
||||
widget.value = mode;
|
||||
}
|
||||
/** Record that this node was written by a version that understands sampling modes. */
|
||||
function stampVersion(node) {
|
||||
node.properties = node.properties || {};
|
||||
node.properties[VERSION_PROPERTY] = CURRENT_VERSION;
|
||||
}
|
||||
const extension = {
|
||||
name: 'ShaderNoiseKSampler.SamplingModeMigration',
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, _app) {
|
||||
if (nodeData.name !== NODE_NAME)
|
||||
return;
|
||||
const origOnNodeCreated = nodeType.prototype.onNodeCreated;
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
if (origOnNodeCreated)
|
||||
origOnNodeCreated.call(this);
|
||||
stampVersion(this);
|
||||
};
|
||||
const origOnConfigure = nodeType.prototype.onConfigure;
|
||||
nodeType.prototype.onConfigure = function (info) {
|
||||
if (origOnConfigure)
|
||||
origOnConfigure.call(this, info);
|
||||
// Read the marker from the saved data, not from this.properties:
|
||||
// onNodeCreated has already stamped the live node by this point.
|
||||
if (needsLegacySampling(info)) {
|
||||
setSamplingMode(this, 'legacy');
|
||||
}
|
||||
stampVersion(this);
|
||||
};
|
||||
},
|
||||
};
|
||||
app.registerExtension(extension);
|
||||
//# sourceMappingURL=sampling_mode_migration.js.map
|
||||
@@ -0,0 +1,84 @@
|
||||
/**
|
||||
* sampling_mode_migration.ts - keeps workflows saved before 2.0 on legacy sampling.
|
||||
*
|
||||
* 2.0 changed how stages sample: they are now segments of one run rather than
|
||||
* independent restarts, denoise and custom sigmas take effect, and blended noise
|
||||
* keeps the distribution the model expects. That changes what an existing seed
|
||||
* produces, so a node loaded from a workflow saved before 2.0 is switched to
|
||||
* "legacy" and reproduces its original output. Newly added nodes keep the
|
||||
* default, "standard".
|
||||
*
|
||||
* Saved nodes are recognised by the absence of the snk_version property, which
|
||||
* only 2.0 and later write.
|
||||
*/
|
||||
// Type imports from our local type definitions
|
||||
import type {
|
||||
ComfyApp,
|
||||
ComfyNodeData,
|
||||
ComfyExtension,
|
||||
NodeTypeConstructor,
|
||||
LGraphNode,
|
||||
} from '../types/comfyui';
|
||||
|
||||
// Import app from ComfyUI at runtime (this import is resolved by the browser)
|
||||
// @ts-ignore - ComfyUI provides this at runtime
|
||||
import { app } from '../../scripts/app.js';
|
||||
|
||||
const NODE_NAME = 'ShaderNoiseKSamplerDirect';
|
||||
const SAMPLING_MODE_WIDGET = 'sampling_mode';
|
||||
const VERSION_PROPERTY = 'snk_version';
|
||||
const CURRENT_VERSION = 2;
|
||||
|
||||
/** Node carrying the properties this migration stamps */
|
||||
interface MigratableNode extends LGraphNode {
|
||||
properties: Record<string, unknown>;
|
||||
}
|
||||
|
||||
/** True when this serialized node predates the sampling modes. */
|
||||
export function needsLegacySampling(info: unknown): boolean {
|
||||
const properties = (info as { properties?: Record<string, unknown> } | null | undefined)?.properties;
|
||||
return properties?.[VERSION_PROPERTY] === undefined;
|
||||
}
|
||||
|
||||
/** Set the sampling_mode widget, if the node has one. */
|
||||
export function setSamplingMode(node: MigratableNode, mode: string): void {
|
||||
const widget = node.widgets?.find((w) => w.name === SAMPLING_MODE_WIDGET);
|
||||
if (widget) widget.value = mode;
|
||||
}
|
||||
|
||||
/** Record that this node was written by a version that understands sampling modes. */
|
||||
function stampVersion(node: MigratableNode): void {
|
||||
node.properties = node.properties || {};
|
||||
node.properties[VERSION_PROPERTY] = CURRENT_VERSION;
|
||||
}
|
||||
|
||||
const extension: ComfyExtension = {
|
||||
name: 'ShaderNoiseKSampler.SamplingModeMigration',
|
||||
async beforeRegisterNodeDef(
|
||||
nodeType: NodeTypeConstructor,
|
||||
nodeData: ComfyNodeData,
|
||||
_app: ComfyApp
|
||||
): Promise<void> {
|
||||
if (nodeData.name !== NODE_NAME) return;
|
||||
|
||||
const origOnNodeCreated = nodeType.prototype.onNodeCreated;
|
||||
nodeType.prototype.onNodeCreated = function (this: MigratableNode): void {
|
||||
if (origOnNodeCreated) origOnNodeCreated.call(this);
|
||||
stampVersion(this);
|
||||
};
|
||||
|
||||
const origOnConfigure = nodeType.prototype.onConfigure;
|
||||
nodeType.prototype.onConfigure = function (this: MigratableNode, info: unknown): void {
|
||||
if (origOnConfigure) origOnConfigure.call(this, info);
|
||||
|
||||
// Read the marker from the saved data, not from this.properties:
|
||||
// onNodeCreated has already stamped the live node by this point.
|
||||
if (needsLegacySampling(info)) {
|
||||
setSamplingMode(this, 'legacy');
|
||||
}
|
||||
stampVersion(this);
|
||||
};
|
||||
},
|
||||
};
|
||||
|
||||
app.registerExtension(extension);
|
||||
Reference in New Issue
Block a user