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:
Æmotion Studio
2026-09-11 15:35:14 -07:00
co-authored by Claude Opus 5
parent cc0e1a52b9
commit 3af34338ab
4 changed files with 331 additions and 112 deletions
+92 -112
View File
@@ -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,)}
+106
View File
@@ -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)
+49
View File
@@ -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
+84
View File
@@ -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);