Initial version
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
from .cfgstar import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
+44
@@ -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",
|
||||
}
|
||||
@@ -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 = ""
|
||||
Reference in New Issue
Block a user