diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..f797810 --- /dev/null +++ b/__init__.py @@ -0,0 +1,4 @@ +from .cfgstar import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/cfgstar.py b/cfgstar.py new file mode 100644 index 0000000..39ead92 --- /dev/null +++ b/cfgstar.py @@ -0,0 +1,44 @@ +import torch + + +class CFGStar: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + } + } + + RETURN_TYPES = ("MODEL",) + FUNCTION = "patch" + + CATEGORY = "advanced/model" + + def patch(self, model): + def custom_cfg_function(args): + cond = args["cond_denoised"] + uncond = args["uncond_denoised"] + cond_scale = args["cond_scale"] + x = args["input"] + + dot_product = torch.sum(cond * uncond, dim=(-1, -2)) + squared_norm = torch.sum(uncond * uncond, dim=(-1, -2)) + 1e-8 + s = dot_product / squared_norm + uncond_scaled = uncond * s[:, :, None, None] + + return x - (uncond_scaled + cond_scale * (cond - uncond_scaled)) + + m = model.clone() + m.set_model_sampler_cfg_function(custom_cfg_function) + + return (m,) + + +NODE_CLASS_MAPPINGS = { + "CFGStar": CFGStar, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "CFGStar": "CFGStar", +} diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..629ef65 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,14 @@ +[project] +name = "comfyui_cfgstar" +description = "A per channel implementation of the scaled CFG from this paper: https://arxiv.org/abs/2503.18886" +version = "1.0.0" +license = { text = "GNU General Public License v3.0" } + +[project.urls] +Repository = "https://github.com/bvhari/ComfyUI_CFGStar" +# Used by Comfy Registry https://comfyregistry.org + +[tool.comfy] +PublisherId = "bvhari" +DisplayName = "ComfyUI_CFGStar" +Icon = ""