MiniMax H3 fixes faces with a latent upscale partway through the run, which means splitting the generation: some steps, the upscaler, then the rest. The sampler could only ever run a schedule end to end, so getting the upscaler in meant dropping shader noise out of the workflow entirely. Four inputs, the same ones KSampler (Advanced) has: add_noise, start_at_step, end_at_step and return_with_leftover_noise. They land at the tail of the optional block, because ComfyUI maps saved widget values by position. The machinery was already there. The pipeline builds one sigma schedule and samples it in segments, slicing sigmas[start:end + 1] for each; a window is that same slicing applied once more at the outer edges. Indices stay absolute, so the per-segment seed still comes out as seed + start -- which means a split run draws the same ancestral and SDE noise as the unsplit one it came from. Decisions worth naming: - Stages spread across the steps the node actually samples, not the whole schedule. Two stages over a three-step window are two stages in those three steps; spread over the schedule, one would land outside the window and never fire. merge_boundaries takes the window's own edges for this, and its first parameter is renamed end_step to say so -- measuring the tail against the schedule length instead would leave a 1-step tail inside a long window. - stage_progression keeps measuring position in the whole trajectory, so a node running the last three steps of seven gets the fine end of coarse_to_fine rather than a fresh coarse-to-fine sweep of its own. - With add_noise off the opening boundary is not painted. The latent already carries its noise from whatever ran before, and the shader on a zero tensor would add back exactly what was turned off. Later boundaries still paint, since their noise is recovered from the latent rather than made here. The refusal for an unpaintable latent weighs only the stages that will actually be painted, so it no longer rejects a latent on behalf of noise nobody makes. Two things fixed on the way: - force_full_denoise was dead. KSampler.sample only reads it alongside last_step, and this code passes sigmas= with last_step unset, so it was discarded on every call. The window's zeroed sigma tail is what does that job now. The golden fixtures recorded the flag, so they are re-captured; every sigma, noise tensor, seed and output across all thirteen is unchanged, which was checked before blessing them. - The progress bar counted the schedule rather than the steps being run. ProgressBar takes its total from the callback, so a windowed run would have started part-filled and stopped short of the end. At the defaults the numbers are identical to before. A schedule with fewer than two sigmas now returns the latent instead of being handed to the sampler, which is what denoise 0.0 builds. Verified against comfy.samplers.KSampler itself across seven start/end/leftover combinations, requiring the sigmas to match bit for bit -- the parity that lets this node sit on one side of a split and a stock KSampler (Advanced) on the other. Two golden fixtures cover the two halves. Also run for real on H3: 4 steps, latent upscale, 3 steps, 896x896x124 with audio, in the new example_workflows/MiniMaxH3_Split_Upscale_SNK_Direct.json. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
191 lines
7.7 KiB
Python
191 lines
7.7 KiB
Python
"""
|
|
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)
|
|
|
|
|
|
def test_every_advertised_shader_type_can_actually_be_resolved():
|
|
"""
|
|
The dropdown must not offer a pattern the generator registry cannot build.
|
|
|
|
It did: `temporal_coherent` shipped, was registered, and passed its own
|
|
tests, but was missing from the node's combo, so no workflow could select
|
|
it. Nothing tied the advertised list to the registry until this.
|
|
"""
|
|
from snk.core.shader_noise import resolve_generator
|
|
|
|
advertised = DirectShaderNoiseKSampler.INPUT_TYPES()["required"]["shader_type"][0]
|
|
assert advertised, "the node must advertise at least one shader type"
|
|
for shader_type in advertised:
|
|
assert callable(resolve_generator(shader_type)), shader_type
|
|
|
|
|
|
def test_the_node_offers_every_registered_generator():
|
|
"""The other direction: a shipped generator that no workflow can reach is dead weight."""
|
|
from snk.shaders.registry import get_shader, list_shaders
|
|
|
|
advertised = set(DirectShaderNoiseKSampler.INPUT_TYPES()["required"]["shader_type"][0])
|
|
# aliases point at a generator already reachable under its canonical name
|
|
canonical = {name for name in list_shaders()
|
|
if not any(get_shader(name) is get_shader(other) and other in advertised
|
|
for other in advertised)}
|
|
assert not canonical - advertised, f"registered but unreachable from the node: {sorted(canonical - advertised)}"
|
|
|
|
|
|
# --- the step window reaches the pipeline ----------------------------------
|
|
|
|
WINDOW_INPUTS = ["add_noise", "start_at_step", "end_at_step", "return_with_leftover_noise"]
|
|
|
|
|
|
def test_the_window_inputs_come_last():
|
|
"""
|
|
ComfyUI maps saved widget values by position, so a new input anywhere but the
|
|
tail re-reads every stored value in every saved workflow. Walk appends its own
|
|
widgets to `required`, which the frontend orders ahead of all of `optional`, so
|
|
the tail of Direct's optional block is the last slot on both nodes.
|
|
"""
|
|
from snk.shader_noise_walk import ShaderNoiseWalk
|
|
|
|
optional = list(DirectShaderNoiseKSampler.INPUT_TYPES()["optional"])
|
|
assert optional[-4:] == WINDOW_INPUTS
|
|
assert list(ShaderNoiseWalk.INPUT_TYPES()["optional"]) == optional
|
|
|
|
|
|
def test_the_window_defaults_sample_the_whole_schedule(sampler_calls):
|
|
"""Every default has to be a no-op, or it changes what saved workflows produce."""
|
|
spec = DirectShaderNoiseKSampler.INPUT_TYPES()["optional"]
|
|
defaults = {name: spec[name][1]["default"] for name in WINDOW_INPUTS}
|
|
assert defaults == {"add_noise": True, "start_at_step": 0, "end_at_step": 10000,
|
|
"return_with_leftover_noise": False}
|
|
|
|
run_node()
|
|
whole = sampler_calls[0]["sigmas"].clone()
|
|
sampler_calls.clear()
|
|
|
|
run_node(**defaults)
|
|
assert torch.equal(sampler_calls[0]["sigmas"], whole)
|
|
|
|
|
|
def test_the_node_forwards_the_window_to_the_pipeline():
|
|
from snk import direct_shader_ksampler
|
|
|
|
with mock.patch.object(direct_shader_ksampler.standard_pipeline, "run") as run:
|
|
run.return_value = {"samples": torch.zeros(1, 4, 16, 16)}
|
|
run_node(add_noise=False, start_at_step=4, end_at_step=7,
|
|
return_with_leftover_noise=True)
|
|
|
|
assert {key: run.call_args.kwargs[key] for key in WINDOW_INPUTS} == {
|
|
"add_noise": False, "start_at_step": 4, "end_at_step": 7,
|
|
"return_with_leftover_noise": True,
|
|
}
|
|
|
|
|
|
def test_legacy_mode_ignores_the_window(sampler_calls):
|
|
"""The frozen path predates it; asking for a window there must not fail the run."""
|
|
run_node(sampling_mode="legacy", start_at_step=5, end_at_step=10,
|
|
return_with_leftover_noise=True)
|
|
|
|
assert sampler_calls and all(call["sigmas"] is None for call in sampler_calls)
|