Initial version

This commit is contained in:
bvhari
2025-03-31 00:47:33 +05:30
parent 78f2053746
commit a1c943b672
3 changed files with 62 additions and 0 deletions
+4
View File
@@ -0,0 +1,4 @@
from .cfgstar import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+44
View File
@@ -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",
}
+14
View File
@@ -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 = ""