2025-05-12 18:10:24 -07:00
import torch
2025-08-22 05:05:36 +03:00
from typing_extensions import override
from comfy_api . latest import ComfyExtension , io
2025-05-12 18:10:24 -07:00
def project ( v0 , v1 ) :
v1 = torch . nn . functional . normalize ( v1 , dim = [ - 1 , - 2 , - 3 ] )
v0_parallel = ( v0 * v1 ) . sum ( dim = [ - 1 , - 2 , - 3 ] , keepdim = True ) * v1
v0_orthogonal = v0 - v0_parallel
return v0_parallel , v0_orthogonal
2025-08-22 05:05:36 +03:00
class APG ( io . ComfyNode ) :
2025-05-12 18:10:24 -07:00
@classmethod
2025-08-22 05:05:36 +03:00
def define_schema ( cls ) - > io . Schema :
return io . Schema (
node_id = " APG " ,
display_name = " Adaptive Projected Guidance " ,
category = " sampling/custom_sampling " ,
2026-02-16 14:02:17 -08:00
description = " Applies Adaptive Projected Guidance to a model, decomposing CFG guidance into parallel and orthogonal components with optional momentum and norm thresholding for improved sampling quality. " ,
short_description = " Decomposes CFG guidance with projection and normalization. " ,
2025-08-22 05:05:36 +03:00
inputs = [
io . Model . Input ( " model " ) ,
io . Float . Input (
" eta " ,
default = 1.0 ,
min = - 10.0 ,
max = 10.0 ,
step = 0.01 ,
tooltip = " Controls the scale of the parallel guidance vector. Default CFG behavior at a setting of 1. " ,
) ,
io . Float . Input (
" norm_threshold " ,
default = 5.0 ,
min = 0.0 ,
max = 50.0 ,
step = 0.1 ,
tooltip = " Normalize guidance vector to this value, normalization disable at a setting of 0. " ,
) ,
io . Float . Input (
" momentum " ,
default = 0.0 ,
min = - 5.0 ,
max = 1.0 ,
step = 0.01 ,
tooltip = " Controls a running average of guidance during diffusion, disabled at a setting of 0. " ,
) ,
] ,
outputs = [ io . Model . Output ( ) ] ,
)
@classmethod
def execute ( cls , model , eta , norm_threshold , momentum ) - > io . NodeOutput :
2025-05-12 18:10:24 -07:00
running_avg = 0
prev_sigma = None
def pre_cfg_function ( args ) :
nonlocal running_avg , prev_sigma
2026-01-01 19:06:14 -08:00
if len ( args [ " conds_out " ] ) == 1 :
return args [ " conds_out " ]
2025-05-12 18:10:24 -07:00
cond = args [ " conds_out " ] [ 0 ]
uncond = args [ " conds_out " ] [ 1 ]
sigma = args [ " sigma " ] [ 0 ]
cond_scale = args [ " cond_scale " ]
if prev_sigma is not None and sigma > prev_sigma :
running_avg = 0
prev_sigma = sigma
guidance = cond - uncond
2025-05-13 10:50:32 -07:00
if momentum != 0 :
2025-05-12 18:10:24 -07:00
if not torch . is_tensor ( running_avg ) :
running_avg = guidance
else :
running_avg = momentum * running_avg + guidance
guidance = running_avg
if norm_threshold > 0 :
guidance_norm = guidance . norm ( p = 2 , dim = [ - 1 , - 2 , - 3 ] , keepdim = True )
scale = torch . minimum (
torch . ones_like ( guidance_norm ) ,
norm_threshold / guidance_norm
)
guidance = guidance * scale
guidance_parallel , guidance_orthogonal = project ( guidance , cond )
modified_guidance = guidance_orthogonal + eta * guidance_parallel
modified_cond = ( uncond + modified_guidance ) + ( cond - uncond ) / cond_scale
return [ modified_cond , uncond ] + args [ " conds_out " ] [ 2 : ]
m = model . clone ( )
m . set_model_sampler_pre_cfg_function ( pre_cfg_function )
2025-08-22 05:05:36 +03:00
return io . NodeOutput ( m )
2025-05-12 18:10:24 -07:00
2025-08-22 05:05:36 +03:00
class ApgExtension ( ComfyExtension ) :
@override
async def get_node_list ( self ) - > list [ type [ io . ComfyNode ] ] :
return [
APG ,
]
2025-05-12 18:10:24 -07:00
2025-08-22 05:05:36 +03:00
async def comfy_entrypoint ( ) - > ApgExtension :
return ApgExtension ( )