From 250db2a0b1c35e967b72dcbd7c9ef527a90f7071 Mon Sep 17 00:00:00 2001 From: primorLee Date: Fri, 21 Aug 2026 15:28:53 +0800 Subject: [PATCH] fix: remove unsupported vanilla cfg parameter --- sgm/inference/api.py | 10 ++-------- tests/inference/test_inference.py | 10 ++++++++++ 2 files changed, 12 insertions(+), 8 deletions(-) diff --git a/sgm/inference/api.py b/sgm/inference/api.py index 2bb0e17de..bec2df45a 100644 --- a/sgm/inference/api.py +++ b/sgm/inference/api.py @@ -263,18 +263,12 @@ def get_guider_config(params: SamplingParams): elif params.guider == Guider.VANILLA: scale = params.scale - thresholder = params.thresholder - - if thresholder == Thresholder.NONE: - dyn_thresh_config = { - "target": "sgm.modules.diffusionmodules.sampling_utils.NoDynamicThresholding" - } - else: + if params.thresholder != Thresholder.NONE: raise NotImplementedError guider_config = { "target": "sgm.modules.diffusionmodules.guiders.VanillaCFG", - "params": {"scale": scale, "dyn_thresh_config": dyn_thresh_config}, + "params": {"scale": scale}, } else: raise NotImplementedError diff --git a/tests/inference/test_inference.py b/tests/inference/test_inference.py index 2b2af11e4..7d33e3667 100644 --- a/tests/inference/test_inference.py +++ b/tests/inference/test_inference.py @@ -15,6 +15,16 @@ import sgm.inference.helpers as helpers +def test_default_sampling_params_build_a_vanilla_cfg_sampler(): + from sgm.inference.api import get_sampler_config + from sgm.modules.diffusionmodules.guiders import VanillaCFG + + sampler = get_sampler_config(SamplingParams()) + + assert isinstance(sampler.guider, VanillaCFG) + assert sampler.guider.scale == 6.0 + + @pytest.mark.inference class TestInference: @fixture(scope="class", params=model_specs.keys())