Merge remote-tracking branch 'upstream/master' into sam3d_body
This commit is contained in:
commit
23c22fb443
|
|
@ -23,9 +23,9 @@ jobs:
|
|||
# SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up
|
||||
# upstream changes; keep `workflows_ref` matching so prompts/scripts load
|
||||
# from the same commit as the workflow definition.
|
||||
uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@047ca48febe3a6647608ed2e0c4331b491cb9d6a # github-workflows#9
|
||||
uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@964d5aad37cbfb57c5b23961d42c2fd85868bf1d # github-workflows main (964d5aa)
|
||||
with:
|
||||
workflows_ref: 047ca48febe3a6647608ed2e0c4331b491cb9d6a
|
||||
workflows_ref: 964d5aad37cbfb57c5b23961d42c2fd85868bf1d
|
||||
diff_excludes: >-
|
||||
:!**/.claude/**
|
||||
:!**/dist/**
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
* @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai
|
||||
|
||||
/CODEOWNERS @comfyanonymous
|
||||
/AGENTS.md @comfyanonymous
|
||||
/.ci/ @comfyanonymous
|
||||
/.github/ @comfyanonymous
|
||||
|
|
|
|||
|
|
@ -0,0 +1,278 @@
|
|||
import re
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
import comfy.ops
|
||||
import comfy.utils
|
||||
|
||||
|
||||
MODULE_PATTERN = re.compile(r"lllite_dit_blocks_(\d+)_(self_attn_[qkv]_proj|cross_attn_q_proj|mlp_layer1)$")
|
||||
|
||||
|
||||
def _group_norm(channels, device=None, dtype=None, operations=None):
|
||||
groups = 8
|
||||
while groups > 1 and channels % groups != 0:
|
||||
groups //= 2
|
||||
return operations.GroupNorm(groups, channels, device=device, dtype=dtype)
|
||||
|
||||
|
||||
class AnimaLLLiteResBlock(nn.Module):
|
||||
def __init__(self, channels, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.norm1 = _group_norm(channels, device=device, dtype=dtype, operations=operations)
|
||||
self.conv1 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype)
|
||||
self.norm2 = _group_norm(channels, device=device, dtype=dtype, operations=operations)
|
||||
self.conv2 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x):
|
||||
h = self.conv1(F.silu(self.norm1(x)))
|
||||
h = self.conv2(F.silu(self.norm2(h)))
|
||||
return x + h
|
||||
|
||||
|
||||
class AnimaLLLiteASPP(nn.Module):
|
||||
def __init__(self, channels, dilations, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
branches = []
|
||||
for dilation in dilations:
|
||||
if dilation == 1:
|
||||
conv = operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype)
|
||||
else:
|
||||
conv = operations.Conv2d(channels, channels, kernel_size=3, padding=dilation, dilation=dilation, device=device, dtype=dtype)
|
||||
branches.append(nn.Sequential(conv, _group_norm(channels, device=device, dtype=dtype, operations=operations), nn.SiLU()))
|
||||
self.branches = nn.ModuleList(branches)
|
||||
self.global_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.global_conv = nn.Sequential(
|
||||
operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype),
|
||||
_group_norm(channels, device=device, dtype=dtype, operations=operations),
|
||||
nn.SiLU(),
|
||||
)
|
||||
self.proj = nn.Sequential(
|
||||
operations.Conv2d(channels * (len(dilations) + 1), channels, kernel_size=1, device=device, dtype=dtype),
|
||||
_group_norm(channels, device=device, dtype=dtype, operations=operations),
|
||||
nn.SiLU(),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
height, width = x.shape[-2:]
|
||||
outputs = [branch(x) for branch in self.branches]
|
||||
pooled = self.global_conv(self.global_pool(x))
|
||||
outputs.append(F.interpolate(pooled, size=(height, width), mode="bilinear", align_corners=False))
|
||||
return self.proj(torch.cat(outputs, dim=1))
|
||||
|
||||
|
||||
class AnimaLLLiteConditioning(nn.Module):
|
||||
def __init__(self, cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
half_dim = cond_dim // 2
|
||||
self.conv1 = operations.Conv2d(cond_in_channels, half_dim, kernel_size=4, stride=4, device=device, dtype=dtype)
|
||||
self.norm1 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations)
|
||||
self.conv2 = operations.Conv2d(half_dim, half_dim, kernel_size=3, padding=1, device=device, dtype=dtype)
|
||||
self.norm2 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations)
|
||||
self.conv3 = operations.Conv2d(half_dim, cond_dim, kernel_size=4, stride=4, device=device, dtype=dtype)
|
||||
self.norm3 = _group_norm(cond_dim, device=device, dtype=dtype, operations=operations)
|
||||
self.resblocks = nn.ModuleList([
|
||||
AnimaLLLiteResBlock(cond_dim, device=device, dtype=dtype, operations=operations)
|
||||
for _ in range(cond_resblocks)
|
||||
])
|
||||
self.aspp = AnimaLLLiteASPP(cond_dim, aspp_dilations, device=device, dtype=dtype, operations=operations) if aspp_dilations else None
|
||||
self.proj = operations.Conv2d(cond_dim, cond_emb_dim, kernel_size=1, device=device, dtype=dtype)
|
||||
self.out_norm = operations.LayerNorm(cond_emb_dim, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.silu(self.norm1(self.conv1(x)))
|
||||
x = F.silu(self.norm2(self.conv2(x)))
|
||||
x = F.silu(self.norm3(self.conv3(x)))
|
||||
for block in self.resblocks:
|
||||
x = block(x)
|
||||
if self.aspp is not None:
|
||||
x = self.aspp(x)
|
||||
x = self.proj(x).flatten(2).transpose(1, 2).contiguous()
|
||||
return self.out_norm(x)
|
||||
|
||||
|
||||
class AnimaLLLiteModule(nn.Module):
|
||||
def __init__(self, in_dim, cond_emb_dim, mlp_dim, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.down = operations.Linear(in_dim, mlp_dim, device=device, dtype=dtype)
|
||||
self.mid = operations.Linear(mlp_dim + cond_emb_dim, mlp_dim, device=device, dtype=dtype)
|
||||
self.cond_to_film = operations.Linear(cond_emb_dim, 2 * mlp_dim, device=device, dtype=dtype)
|
||||
self.up = operations.Linear(mlp_dim, in_dim, device=device, dtype=dtype)
|
||||
self.depth_embed = nn.Parameter(torch.empty(cond_emb_dim, device=device, dtype=dtype), requires_grad=False)
|
||||
|
||||
def forward(self, x, cond_emb, strength):
|
||||
original_shape = x.shape
|
||||
if x.ndim == 5:
|
||||
x = x.flatten(1, 3)
|
||||
|
||||
if x.shape[0] != cond_emb.shape[0]:
|
||||
if x.shape[0] % cond_emb.shape[0] != 0:
|
||||
raise ValueError(f"Anima LLLite batch mismatch: model input batch {x.shape[0]}, control batch {cond_emb.shape[0]}")
|
||||
cond_emb = cond_emb.repeat(x.shape[0] // cond_emb.shape[0], 1, 1)
|
||||
if x.shape[1] != cond_emb.shape[1]:
|
||||
raise ValueError(f"Anima LLLite sequence mismatch: model input has {x.shape[1]} tokens, control has {cond_emb.shape[1]}")
|
||||
|
||||
cond_local = cond_emb + comfy.ops.cast_to_input(self.depth_embed, cond_emb)
|
||||
hidden = F.silu(self.down(x))
|
||||
gamma, beta = self.cond_to_film(cond_local).chunk(2, dim=-1)
|
||||
hidden = self.mid(torch.cat((cond_local, hidden), dim=-1))
|
||||
hidden = F.silu(hidden * (1 + gamma) + beta)
|
||||
x = x + self.up(hidden) * strength
|
||||
|
||||
if len(original_shape) == 5:
|
||||
x = x.reshape(original_shape)
|
||||
return x
|
||||
|
||||
|
||||
class AnimaLLLite(nn.Module):
|
||||
def __init__(self, state_dict, metadata, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
metadata = metadata or {}
|
||||
version = metadata.get("lllite.version", "2")
|
||||
if version != "2":
|
||||
raise ValueError(f"Unsupported Anima LLLite version {version!r}; only named-key v2 checkpoints are supported")
|
||||
|
||||
module_names = sorted({key.split(".", 1)[0] for key in state_dict if key.startswith("lllite_dit_blocks_")})
|
||||
if not module_names:
|
||||
raise ValueError("Anima LLLite checkpoint has no lllite_dit_blocks_* modules")
|
||||
|
||||
cond_in_channels = state_dict["lllite_conditioning1.conv1.weight"].shape[1]
|
||||
cond_dim = state_dict["lllite_conditioning1.conv3.weight"].shape[0]
|
||||
cond_emb_dim = state_dict["lllite_conditioning1.proj.weight"].shape[0]
|
||||
resblock_ids = {int(key.split(".")[2]) for key in state_dict if key.startswith("lllite_conditioning1.resblocks.")}
|
||||
cond_resblocks = max(resblock_ids) + 1 if resblock_ids else 0
|
||||
use_aspp = any(key.startswith("lllite_conditioning1.aspp.") for key in state_dict)
|
||||
dilation_string = metadata.get("lllite.aspp_dilations", "1,2,4,8")
|
||||
aspp_dilations = tuple(int(value) for value in dilation_string.split(",") if value.strip()) if use_aspp else ()
|
||||
|
||||
self.cond_in_channels = cond_in_channels
|
||||
self.inpaint_masked_input = metadata.get("lllite.inpaint_masked_input", "false").lower() == "true"
|
||||
self.lllite_conditioning1 = AnimaLLLiteConditioning(
|
||||
cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations,
|
||||
device=device, dtype=dtype, operations=operations,
|
||||
)
|
||||
|
||||
self.module_names = set()
|
||||
self.block_count = 0
|
||||
self.model_dim = None
|
||||
for name in module_names:
|
||||
match = MODULE_PATTERN.fullmatch(name)
|
||||
if match is None:
|
||||
raise ValueError(f"Unsupported Anima LLLite module name: {name}")
|
||||
down_shape = state_dict[f"{name}.down.weight"].shape
|
||||
mlp_dim, in_dim = down_shape
|
||||
module_cond_dim = state_dict[f"{name}.cond_to_film.weight"].shape[1]
|
||||
if module_cond_dim != cond_emb_dim:
|
||||
raise ValueError(f"Anima LLLite conditioning dimension mismatch in {name}: {module_cond_dim} != {cond_emb_dim}")
|
||||
if self.model_dim is None:
|
||||
self.model_dim = in_dim
|
||||
elif self.model_dim != in_dim:
|
||||
raise ValueError(f"Anima LLLite model dimension mismatch in {name}: {in_dim} != {self.model_dim}")
|
||||
self.add_module(name, AnimaLLLiteModule(in_dim, cond_emb_dim, mlp_dim, device=device, dtype=dtype, operations=operations))
|
||||
self.module_names.add(name)
|
||||
self.block_count = max(self.block_count, int(match.group(1)) + 1)
|
||||
|
||||
def encode_conditioning(self, image):
|
||||
return self.lllite_conditioning1(image)
|
||||
|
||||
def apply(self, x, cond_emb, block_index, target, strength):
|
||||
name = f"lllite_dit_blocks_{block_index}_{target}"
|
||||
if name not in self.module_names:
|
||||
return x
|
||||
return self.get_submodule(name)(x, cond_emb, strength)
|
||||
|
||||
|
||||
class AnimaLLLitePatch:
|
||||
def __init__(self, model_patch, image, mask, strength, sigma_start, sigma_end):
|
||||
self.model_patch = model_patch
|
||||
self.image = image
|
||||
self.mask = mask
|
||||
self.strength = strength
|
||||
self.sigma_start = sigma_start
|
||||
self.sigma_end = sigma_end
|
||||
|
||||
def __call__(self, args):
|
||||
x = args["x"]
|
||||
transformer_options = args["transformer_options"]
|
||||
if self.strength == 0.0:
|
||||
return args
|
||||
sigmas = transformer_options.get("sigmas")
|
||||
if sigmas is not None:
|
||||
sigma = float(sigmas.max().item())
|
||||
if not self.sigma_end <= sigma <= self.sigma_start:
|
||||
return args
|
||||
if x.shape[2] != 1:
|
||||
raise ValueError(f"Anima LLLite only supports T=1, got T={x.shape[2]}")
|
||||
|
||||
target_height = x.shape[-2] * 8
|
||||
target_width = x.shape[-1] * 8
|
||||
image = comfy.utils.common_upscale(
|
||||
self.image.movedim(-1, 1), target_width, target_height, "bicubic", crop="center"
|
||||
).clamp(0.0, 1.0)
|
||||
image = image.to(device=x.device, dtype=x.dtype) * 2.0 - 1.0
|
||||
|
||||
if self.model_patch.model.cond_in_channels == 4:
|
||||
mask = self.mask
|
||||
if mask.ndim == 3:
|
||||
mask = mask.unsqueeze(1)
|
||||
if mask.ndim != 4 or mask.shape[1] != 1:
|
||||
raise ValueError(f"Anima LLLite mask must have one channel, got shape {tuple(mask.shape)}")
|
||||
mask = comfy.utils.common_upscale(
|
||||
mask.float(), target_width, target_height, "nearest-exact", crop="center"
|
||||
)
|
||||
if mask.shape[0] != image.shape[0]:
|
||||
if image.shape[0] % mask.shape[0] != 0:
|
||||
raise ValueError(
|
||||
f"Anima LLLite mask batch {mask.shape[0]} cannot be broadcast to image batch {image.shape[0]}"
|
||||
)
|
||||
mask = mask.repeat(image.shape[0] // mask.shape[0], 1, 1, 1)
|
||||
mask = (mask >= 0.5).to(device=x.device, dtype=x.dtype)
|
||||
if self.model_patch.model.inpaint_masked_input:
|
||||
image = image * (mask < 0.5).to(image.dtype)
|
||||
image = torch.cat((image, mask * 2.0 - 1.0), dim=1)
|
||||
|
||||
cond_emb = self.model_patch.model.encode_conditioning(image)
|
||||
transformer_options["model_patch_data"][self] = cond_emb
|
||||
return args
|
||||
|
||||
def to(self, device_or_dtype):
|
||||
return self
|
||||
|
||||
def models(self):
|
||||
return [self.model_patch]
|
||||
|
||||
|
||||
class AnimaLLLiteAttentionPatch:
|
||||
def __init__(self, patch, targets):
|
||||
self.patch = patch
|
||||
self.targets = targets
|
||||
|
||||
def __call__(self, q, k, v, pe=None, attn_mask=None, extra_options=None):
|
||||
cond_emb = extra_options["model_patch_data"].get(self.patch)
|
||||
if cond_emb is None:
|
||||
return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask}
|
||||
|
||||
block_index = extra_options["block_index"]
|
||||
values = {"q": q, "k": k, "v": v}
|
||||
for value_name, target in self.targets.items():
|
||||
values[value_name] = self.patch.model_patch.model.apply(
|
||||
values[value_name], cond_emb, block_index, target, self.patch.strength
|
||||
)
|
||||
|
||||
return {"q": values["q"], "k": values["k"], "v": values["v"], "pe": pe, "attn_mask": attn_mask}
|
||||
|
||||
|
||||
class AnimaLLLiteMLPPatch:
|
||||
def __init__(self, patch):
|
||||
self.patch = patch
|
||||
|
||||
def __call__(self, args):
|
||||
cond_emb = args["transformer_options"]["model_patch_data"].get(self.patch)
|
||||
if cond_emb is None:
|
||||
return args
|
||||
args["x"] = self.patch.model_patch.model.apply(
|
||||
args["x"], cond_emb, args["transformer_options"]["block_index"], "mlp_layer1", self.patch.strength
|
||||
)
|
||||
return args
|
||||
|
|
@ -149,11 +149,29 @@ class Attention(nn.Module):
|
|||
x: torch.Tensor,
|
||||
context: Optional[torch.Tensor] = None,
|
||||
rope_emb: Optional[torch.Tensor] = None,
|
||||
transformer_options: Optional[dict] = {},
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
q = self.q_proj(x)
|
||||
context = x if context is None else context
|
||||
k = self.k_proj(context)
|
||||
v = self.v_proj(context)
|
||||
q_input = x
|
||||
k_input = context
|
||||
v_input = context
|
||||
|
||||
transformer_patches = transformer_options.get("patches", {})
|
||||
patch_name = "attn1_patch" if self.is_selfattn else "attn2_patch"
|
||||
if patch_name in transformer_patches:
|
||||
extra_options = transformer_options.copy()
|
||||
extra_options["n_heads"] = self.n_heads
|
||||
extra_options["dim_head"] = self.head_dim
|
||||
for patch in transformer_patches[patch_name]:
|
||||
out = patch(q_input, k_input, v_input, pe=rope_emb, attn_mask=None, extra_options=extra_options)
|
||||
q_input = out.get("q", q_input)
|
||||
k_input = out.get("k", k_input)
|
||||
v_input = out.get("v", v_input)
|
||||
rope_emb = out.get("pe", rope_emb)
|
||||
|
||||
q = self.q_proj(q_input)
|
||||
k = self.k_proj(k_input)
|
||||
v = self.v_proj(v_input)
|
||||
q, k, v = map(
|
||||
lambda t: rearrange(t, "b ... (h d) -> b ... h d", h=self.n_heads, d=self.head_dim),
|
||||
(q, k, v),
|
||||
|
|
@ -194,7 +212,7 @@ class Attention(nn.Module):
|
|||
x (Tensor): The query tensor of shape [B, Mq, K]
|
||||
context (Optional[Tensor]): The key tensor of shape [B, Mk, K] or use x as context [self attention] if None
|
||||
"""
|
||||
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb)
|
||||
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb, transformer_options=transformer_options)
|
||||
return self.compute_attention(q, k, v, transformer_options=transformer_options)
|
||||
|
||||
|
||||
|
|
@ -561,8 +579,14 @@ class Block(nn.Module):
|
|||
self.layer_norm_mlp,
|
||||
scale_mlp_B_T_1_1_D,
|
||||
shift_mlp_B_T_1_1_D,
|
||||
)
|
||||
result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D.to(compute_dtype))
|
||||
).to(compute_dtype)
|
||||
patches = transformer_options.get("patches", {})
|
||||
if "mlp_patch" in patches:
|
||||
args = {"x": normalized_x_B_T_H_W_D, "transformer_options": transformer_options}
|
||||
for patch in patches["mlp_patch"]:
|
||||
args = patch(args)
|
||||
normalized_x_B_T_H_W_D = args["x"]
|
||||
result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D)
|
||||
x_B_T_H_W_D = torch.addcmul(x_B_T_H_W_D, gate_mlp_B_T_1_1_D.to(residual_dtype), result_B_T_H_W_D.to(residual_dtype))
|
||||
return x_B_T_H_W_D
|
||||
|
||||
|
|
@ -869,11 +893,22 @@ class MiniTrainDIT(nn.Module):
|
|||
x_B_T_H_W_D.shape == extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape
|
||||
), f"{x_B_T_H_W_D.shape} != {extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape}"
|
||||
|
||||
transformer_options = kwargs.get("transformer_options", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
if "post_input" in patches:
|
||||
transformer_options = transformer_options.copy()
|
||||
transformer_options["model_patch_data"] = {}
|
||||
|
||||
if "post_input" in patches:
|
||||
for patch in patches["post_input"]:
|
||||
out = patch({"img": x_B_T_H_W_D, "x": x_B_C_T_H_W, "transformer_options": transformer_options})
|
||||
x_B_T_H_W_D = out["img"]
|
||||
|
||||
block_kwargs = {
|
||||
"rope_emb_L_1_1_D": rope_emb_L_1_1_D.unsqueeze(1).unsqueeze(0),
|
||||
"adaln_lora_B_T_3D": adaln_lora_B_T_3D,
|
||||
"extra_per_block_pos_emb": extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D,
|
||||
"transformer_options": kwargs.get("transformer_options", {}),
|
||||
"transformer_options": transformer_options,
|
||||
}
|
||||
|
||||
# The residual stream for this model has large values. To make fp16 compute_dtype work, we keep the residual stream
|
||||
|
|
@ -883,7 +918,8 @@ class MiniTrainDIT(nn.Module):
|
|||
if x_B_T_H_W_D.dtype == torch.float16:
|
||||
x_B_T_H_W_D = x_B_T_H_W_D.float()
|
||||
|
||||
for block in self.blocks:
|
||||
for block_index, block in enumerate(self.blocks):
|
||||
transformer_options["block_index"] = block_index
|
||||
x_B_T_H_W_D = block(
|
||||
x_B_T_H_W_D,
|
||||
t_embedding_B_T_D,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import torch.nn.functional as F
|
|||
|
||||
from comfy.ldm.modules.attention import optimized_attention
|
||||
import comfy.model_management
|
||||
import comfy.ops
|
||||
import comfy.quant_ops
|
||||
|
||||
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
|
||||
|
|
@ -111,11 +112,17 @@ class ErnieImageAttention(nn.Module):
|
|||
query = q_flat.view(B, S, self.heads, self.head_dim)
|
||||
key = k_flat.view(B, S, self.heads, self.head_dim)
|
||||
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb)
|
||||
if image_rotary_emb is not None and not comfy.model_management.in_training:
|
||||
q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, query, offloadable=True)
|
||||
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, key, offloadable=True)
|
||||
query, key = comfy.quant_ops.ck.rms_rope_split_half(query, key, image_rotary_emb, q_scale, k_scale, self.norm_q.eps)
|
||||
comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream)
|
||||
comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream)
|
||||
else:
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
if image_rotary_emb is not None:
|
||||
query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb)
|
||||
|
||||
q_flat = query.reshape(B, S, -1)
|
||||
k_flat = key.reshape(B, S, -1)
|
||||
|
|
|
|||
|
|
@ -94,12 +94,21 @@ class JoyImageAttention(nn.Module):
|
|||
txt_k = txt_k.unflatten(-1, (heads, -1))
|
||||
txt_v = txt_v.unflatten(-1, (heads, -1))
|
||||
|
||||
img_q = self.img_attn_q_norm(img_q)
|
||||
img_k = self.img_attn_k_norm(img_k)
|
||||
txt_q = self.txt_attn_q_norm(txt_q)
|
||||
txt_k = self.txt_attn_k_norm(txt_k)
|
||||
|
||||
img_q, img_k = comfy_kitchen.apply_rope(img_q, img_k, image_rotary_emb)
|
||||
img_q_scale, _, img_q_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_q_norm, img_q, offloadable=True)
|
||||
img_k_scale, _, img_k_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_k_norm, img_k, offloadable=True)
|
||||
img_q, img_k = comfy_kitchen.rms_rope(
|
||||
img_q,
|
||||
img_k,
|
||||
image_rotary_emb,
|
||||
img_q_scale,
|
||||
img_k_scale,
|
||||
self.img_attn_q_norm.eps,
|
||||
)
|
||||
comfy.ops.uncast_bias_weight(self.img_attn_q_norm, img_q_scale, None, img_q_offload_stream)
|
||||
comfy.ops.uncast_bias_weight(self.img_attn_k_norm, img_k_scale, None, img_k_offload_stream)
|
||||
|
||||
joint_q = torch.cat([img_q, txt_q], dim=1)
|
||||
joint_k = torch.cat([img_k, txt_k], dim=1)
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from einops import rearrange
|
|||
import comfy.model_management
|
||||
import comfy.patcher_extension
|
||||
import comfy.ldm.common_dit
|
||||
import comfy.utils
|
||||
from comfy.ldm.flux.layers import EmbedND, timestep_embedding
|
||||
from comfy.ldm.flux.math import apply_rope
|
||||
from comfy.ldm.modules.attention import optimized_attention_masked
|
||||
|
|
@ -73,11 +74,20 @@ class Attention(nn.Module):
|
|||
self.wo = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x, freqs=None, mask=None, transformer_options={}):
|
||||
transformer_patches = transformer_options.get("patches", {})
|
||||
extra_options = transformer_options.copy()
|
||||
q, k, v, gate = self.wq(x), self.wk(x), self.wv(x), self.gate(x)
|
||||
q = rearrange(q, "B L (H D) -> B H L D", H=self.heads)
|
||||
k = rearrange(k, "B L (H D) -> B H L D", H=self.kvheads)
|
||||
v = rearrange(v, "B L (H D) -> B H L D", H=self.kvheads)
|
||||
q, k = self.qknorm(q, k)
|
||||
|
||||
if "block_index" in transformer_options and "attn1_patch" in transformer_patches:
|
||||
for p in transformer_patches["attn1_patch"]:
|
||||
out = p(q, k, v, pe=freqs, attn_mask=mask, extra_options=extra_options)
|
||||
q, k, v = out.get("q", q), out.get("k", k), out.get("v", v)
|
||||
freqs, mask = out.get("pe", freqs), out.get("attn_mask", mask)
|
||||
|
||||
if freqs is not None:
|
||||
q, k = apply_rope(q, k, freqs)
|
||||
if self.kvheads != self.heads:
|
||||
|
|
@ -86,6 +96,11 @@ class Attention(nn.Module):
|
|||
v = v.repeat_interleave(rep, dim=1)
|
||||
out = optimized_attention_masked(q, k, v, self.heads, mask=mask, skip_reshape=True,
|
||||
transformer_options=transformer_options)
|
||||
|
||||
if "block_index" in transformer_options and "attn1_output_patch" in transformer_patches:
|
||||
for p in transformer_patches["attn1_output_patch"]:
|
||||
out = p(out, extra_options)
|
||||
|
||||
return self.wo(out * F.sigmoid(gate))
|
||||
|
||||
|
||||
|
|
@ -158,8 +173,44 @@ class SingleStreamBlock(nn.Module):
|
|||
self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations)
|
||||
self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations)
|
||||
|
||||
def forward(self, x, vec, freqs, mask=None, transformer_options={}):
|
||||
def forward(self, x, vec, freqs, mask=None, timestep_zero_index=None, transformer_options={}):
|
||||
prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec)
|
||||
if timestep_zero_index is not None:
|
||||
bs = x.shape[0]
|
||||
ref_prescale = prescale[bs:]
|
||||
ref_preshift = preshift[bs:]
|
||||
ref_pregate = pregate[bs:]
|
||||
ref_postscale = postscale[bs:]
|
||||
ref_postshift = postshift[bs:]
|
||||
ref_postgate = postgate[bs:]
|
||||
prescale = prescale[:bs]
|
||||
preshift = preshift[:bs]
|
||||
pregate = pregate[:bs]
|
||||
postscale = postscale[:bs]
|
||||
postshift = postshift[:bs]
|
||||
postgate = postgate[:bs]
|
||||
|
||||
pre = self.prenorm(x)
|
||||
pre[:, :timestep_zero_index].mul_(1 + prescale).add_(preshift)
|
||||
pre[:, timestep_zero_index:].mul_(1 + ref_prescale).add_(ref_preshift)
|
||||
attn = self.attn(pre, freqs, mask, transformer_options=transformer_options)
|
||||
del pre
|
||||
attn[:, :timestep_zero_index].mul_(pregate)
|
||||
attn[:, timestep_zero_index:].mul_(ref_pregate)
|
||||
x = x + attn
|
||||
del attn
|
||||
|
||||
post = self.postnorm(x)
|
||||
post[:, :timestep_zero_index].mul_(1 + postscale).add_(postshift)
|
||||
post[:, timestep_zero_index:].mul_(1 + ref_postscale).add_(ref_postshift)
|
||||
mlp = self.mlp(post)
|
||||
del post
|
||||
mlp[:, :timestep_zero_index].mul_(postgate)
|
||||
mlp[:, timestep_zero_index:].mul_(ref_postgate)
|
||||
x = x + mlp
|
||||
del mlp
|
||||
return x
|
||||
|
||||
x = x + pregate * self.attn((1 + prescale) * self.prenorm(x) + preshift, freqs, mask, transformer_options=transformer_options)
|
||||
x = x + postgate * self.mlp((1 + postscale) * self.postnorm(x) + postshift)
|
||||
return x
|
||||
|
|
@ -181,7 +232,7 @@ class LastLayer(nn.Module):
|
|||
class SingleStreamDiT(nn.Module):
|
||||
def __init__(self, features=6144, tdim=256, txtdim=2560, heads=48, kvheads=12, multiplier=4,
|
||||
layers=28, patch=2, channels=16, bias=False, theta=1e3, txtlayers=12,
|
||||
txtheads=20, txtkvheads=20, image_model=None,
|
||||
txtheads=20, txtkvheads=20, default_ref_method=None, image_model=None,
|
||||
device=None, dtype=None, operations=None, **kwargs):
|
||||
super().__init__()
|
||||
self.dtype = dtype
|
||||
|
|
@ -191,6 +242,7 @@ class SingleStreamDiT(nn.Module):
|
|||
self.heads = heads
|
||||
self.txtdim = txtdim
|
||||
self.txtlayers = txtlayers
|
||||
self.default_ref_method = default_ref_method
|
||||
|
||||
headdim = features // heads
|
||||
axes = [headdim - 12 * (headdim // 16), 6 * (headdim // 16), 6 * (headdim // 16)]
|
||||
|
|
@ -221,61 +273,110 @@ class SingleStreamDiT(nn.Module):
|
|||
operations.Linear(features, features * 6, device=device, dtype=dtype),
|
||||
)
|
||||
|
||||
def forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs):
|
||||
def forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
|
||||
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
||||
self._forward,
|
||||
self,
|
||||
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options),
|
||||
).execute(x, timesteps, context, attention_mask, transformer_options, **kwargs)
|
||||
).execute(x, timesteps, context, attention_mask, ref_latents, transformer_options, **kwargs)
|
||||
|
||||
def _forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs):
|
||||
def process_img(self, x, index=0):
|
||||
patch = self.patch
|
||||
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch))
|
||||
h, w = x.shape[-2] // patch, x.shape[-1] // patch
|
||||
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch)
|
||||
|
||||
img_ids = torch.zeros(h, w, 3, device=x.device, dtype=torch.float32)
|
||||
img_ids[..., 0] = index
|
||||
img_ids[..., 1] = torch.arange(h, device=x.device, dtype=torch.float32)[:, None]
|
||||
img_ids[..., 2] = torch.arange(w, device=x.device, dtype=torch.float32)[None, :]
|
||||
return img, img_ids.reshape(1, h * w, 3).repeat(x.shape[0], 1, 1), h, w
|
||||
|
||||
def _forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
|
||||
transformer_options = transformer_options.copy()
|
||||
temporal = x.ndim == 5
|
||||
if temporal:
|
||||
b5, c5, t5, h5, w5 = x.shape
|
||||
x = x.reshape(b5 * t5, c5, h5, w5)
|
||||
bs, c, H_orig, W_orig = x.shape
|
||||
bs, _, h_orig, w_orig = x.shape
|
||||
patch = self.patch
|
||||
# Pad the latent up to a multiple of patch (as Flux/Lumina/QwenImage do); crop back at the end.
|
||||
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch))
|
||||
H, W = x.shape[-2], x.shape[-1]
|
||||
h_, w_ = H // patch, W // patch
|
||||
|
||||
# context arrives as (B, seq, txtlayers*txtdim); reshape to (B, txtlayers, seq, txtdim).
|
||||
context = self._unpack_context(context)
|
||||
|
||||
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch)
|
||||
img, imgpos, h_, w_ = self.process_img(x)
|
||||
img_tokens = img.shape[1]
|
||||
timestep_zero_index = None
|
||||
ref_method = kwargs.get("ref_latents_method", self.default_ref_method)
|
||||
if ref_method is not None and ref_latents is not None and len(ref_latents) > 0:
|
||||
ref_tokens = []
|
||||
ref_pos = []
|
||||
ref_num_tokens = []
|
||||
for index, ref in enumerate(ref_latents, 1):
|
||||
if ref.ndim == 5:
|
||||
rb, rc, rt, rh5, rw5 = ref.shape
|
||||
ref = ref.reshape(rb * rt, rc, rh5, rw5)
|
||||
ref = comfy.utils.repeat_to_batch_size(ref, bs)
|
||||
kontext, kontext_ids, _, _ = self.process_img(ref, index=index)
|
||||
ref_tokens.append(kontext)
|
||||
ref_pos.append(kontext_ids)
|
||||
ref_num_tokens.append(kontext.shape[1])
|
||||
img = torch.cat([img] + ref_tokens, dim=1)
|
||||
imgpos = torch.cat([imgpos] + ref_pos, dim=1)
|
||||
del ref_tokens, ref_pos
|
||||
if ref_method == "index_timestep_zero":
|
||||
timestep_zero_index = img_tokens
|
||||
transformer_options["reference_image_num_tokens"] = ref_num_tokens
|
||||
|
||||
img = self.first(img)
|
||||
|
||||
t = self.tmlp(timestep_embedding(timesteps, self.tdim).unsqueeze(1).to(img.dtype))
|
||||
tvec = self.tproj(t)
|
||||
if timestep_zero_index is not None:
|
||||
t0 = self.tmlp(timestep_embedding(torch.zeros_like(timesteps), self.tdim).unsqueeze(1).to(img.dtype))
|
||||
tvec = torch.cat((tvec, self.tproj(t0)), dim=0)
|
||||
|
||||
context = self.txtfusion(context, mask=None, transformer_options=transformer_options)
|
||||
context = self.txtmlp(context)
|
||||
|
||||
txtlen, imglen = context.shape[1], img.shape[1]
|
||||
txtlen = context.shape[1]
|
||||
device = context.device
|
||||
txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32)
|
||||
|
||||
patches = transformer_options.get("patches", {})
|
||||
if "post_input" in patches:
|
||||
for p in patches["post_input"]:
|
||||
out = p({"img": img, "txt": context, "img_ids": imgpos, "txt_ids": txtpos, "transformer_options": transformer_options})
|
||||
img, context = out["img"], out["txt"]
|
||||
imgpos, txtpos = out["img_ids"], out["txt_ids"]
|
||||
|
||||
combined = torch.cat((context, img), dim=1)
|
||||
del context, img
|
||||
if timestep_zero_index is not None:
|
||||
timestep_zero_index += txtlen
|
||||
|
||||
# Position ids: text at 0, image at (0, h_idx, w_idx).
|
||||
device = combined.device
|
||||
txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32)
|
||||
imgids = torch.zeros(h_, w_, 3, device=device, dtype=torch.float32)
|
||||
imgids[..., 1] = torch.arange(h_, device=device, dtype=torch.float32)[:, None]
|
||||
imgids[..., 2] = torch.arange(w_, device=device, dtype=torch.float32)[None, :]
|
||||
imgpos = imgids.reshape(1, h_ * w_, 3).repeat(bs, 1, 1)
|
||||
pos = torch.cat((txtpos, imgpos), dim=1)
|
||||
del txtpos, imgpos
|
||||
|
||||
freqs = self.pe_embedder(pos)
|
||||
del pos
|
||||
|
||||
for block in self.blocks:
|
||||
combined = block(combined, tvec, freqs, None, transformer_options=transformer_options)
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "single"
|
||||
transformer_options["img_slice"] = [txtlen, combined.shape[1]]
|
||||
for i, block in enumerate(self.blocks):
|
||||
transformer_options["block_index"] = i
|
||||
combined = block(combined, tvec, freqs, None, timestep_zero_index=timestep_zero_index, transformer_options=transformer_options)
|
||||
|
||||
final = self.last(combined, t)
|
||||
out = final[:, txtlen:txtlen + imglen, :]
|
||||
del combined
|
||||
out = final[:, txtlen:txtlen + img_tokens, :]
|
||||
out = rearrange(out, "b (h w) (c ph pw) -> b c (h ph) (w pw)",
|
||||
h=h_, w=w_, ph=patch, pw=patch, c=self.channels)
|
||||
out = out[:, :, :H_orig, :W_orig] # crop padding back off
|
||||
out = out[:, :, :h_orig, :w_orig] # crop padding back off
|
||||
if temporal:
|
||||
out = out.reshape(b5, t5, self.channels, H_orig, W_orig).movedim(1, 2)
|
||||
out = out.reshape(b5, t5, self.channels, h_orig, w_orig).movedim(1, 2)
|
||||
return out
|
||||
|
||||
def _unpack_context(self, context):
|
||||
|
|
|
|||
|
|
@ -6,6 +6,9 @@ import torch
|
|||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import comfy.ldm.common_dit
|
||||
import comfy.model_management
|
||||
import comfy.ops
|
||||
import comfy.quant_ops
|
||||
|
||||
from comfy.ldm.modules.diffusionmodules.mmdit import TimestepEmbedder
|
||||
from comfy.ldm.modules.attention import optimized_attention_masked
|
||||
|
|
@ -97,6 +100,7 @@ class JointAttention(nn.Module):
|
|||
self.n_local_kv_heads = self.n_kv_heads
|
||||
self.n_rep = self.n_local_heads // self.n_local_kv_heads
|
||||
self.head_dim = dim // n_heads
|
||||
self.qk_norm = qk_norm
|
||||
|
||||
self.qkv = operation_settings.get("operations").Linear(
|
||||
dim,
|
||||
|
|
@ -151,10 +155,21 @@ class JointAttention(nn.Module):
|
|||
xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)
|
||||
xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)
|
||||
|
||||
xq = self.q_norm(xq)
|
||||
xk = self.k_norm(xk)
|
||||
|
||||
xq, xk = apply_rope(xq, xk, freqs_cis)
|
||||
if self.qk_norm and not comfy.model_management.in_training:
|
||||
q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, xq, offloadable=True)
|
||||
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, xk, offloadable=True)
|
||||
epsilon = self.q_norm.eps if self.q_norm.eps is not None else torch.finfo(torch.float32).eps
|
||||
if self.n_local_heads == self.n_local_kv_heads:
|
||||
xq, xk = comfy.quant_ops.ck.rms_rope(xq, xk, freqs_cis, q_scale, k_scale, epsilon)
|
||||
else:
|
||||
xq = comfy.quant_ops.ck.rms_rope1(xq, freqs_cis, q_scale, epsilon)
|
||||
xk = comfy.quant_ops.ck.rms_rope1(xk, freqs_cis, k_scale, epsilon)
|
||||
comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream)
|
||||
comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream)
|
||||
else:
|
||||
xq = self.q_norm(xq)
|
||||
xk = self.k_norm(xk)
|
||||
xq, xk = apply_rope(xq, xk, freqs_cis)
|
||||
|
||||
n_rep = self.n_local_heads // self.n_local_kv_heads
|
||||
if n_rep >= 1:
|
||||
|
|
|
|||
|
|
@ -552,6 +552,7 @@ class WanModel(torch.nn.Module):
|
|||
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
grid_sizes = x.shape[2:]
|
||||
transformer_options["grid_sizes"] = grid_sizes
|
||||
|
|
@ -564,11 +565,13 @@ class WanModel(torch.nn.Module):
|
|||
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
|
||||
|
||||
full_ref = None
|
||||
img_offset = 0
|
||||
if self.ref_conv is not None:
|
||||
full_ref = kwargs.get("reference_latent", None)
|
||||
if full_ref is not None:
|
||||
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
|
||||
x = torch.concat((full_ref, x), dim=1)
|
||||
img_offset = full_ref.shape[1]
|
||||
|
||||
# In-context reference (Bernini)
|
||||
context_latents = kwargs.get("context_latents", None)
|
||||
|
|
@ -589,6 +592,7 @@ class WanModel(torch.nn.Module):
|
|||
context_img_len = clip_fea.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
|
|
@ -604,6 +608,11 @@ class WanModel(torch.nn.Module):
|
|||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
|
|
@ -777,6 +786,7 @@ class VaceWanModel(WanModel):
|
|||
**kwargs,
|
||||
):
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
grid_sizes = x.shape[2:]
|
||||
transformer_options["grid_sizes"] = grid_sizes
|
||||
|
|
@ -807,6 +817,7 @@ class VaceWanModel(WanModel):
|
|||
x_orig = x
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
|
|
@ -822,6 +833,11 @@ class VaceWanModel(WanModel):
|
|||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
ii = self.vace_layers_mapping.get(i, None)
|
||||
if ii is not None:
|
||||
for iii in range(len(c)):
|
||||
|
|
@ -887,6 +903,7 @@ class CameraWanModel(WanModel):
|
|||
**kwargs,
|
||||
):
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
if self.control_adapter is not None and camera_conditions is not None:
|
||||
x = x + self.control_adapter(camera_conditions).to(x.dtype)
|
||||
|
|
@ -909,6 +926,7 @@ class CameraWanModel(WanModel):
|
|||
context_img_len = clip_fea.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
|
|
@ -924,6 +942,11 @@ class CameraWanModel(WanModel):
|
|||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
|
|
@ -1335,6 +1358,7 @@ class WanModel_S2V(WanModel):
|
|||
|
||||
# embeddings
|
||||
bs, _, time, height, width = x.shape
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
if control_video is not None:
|
||||
x = x + self.cond_encoder(control_video)
|
||||
|
|
@ -1379,6 +1403,7 @@ class WanModel_S2V(WanModel):
|
|||
context = self.text_embedding(context)
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
|
|
@ -1393,6 +1418,12 @@ class WanModel_S2V(WanModel):
|
|||
x = out["img"]
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
if audio_emb is not None:
|
||||
x = self.audio_injector(x, i, audio_emb, audio_emb_global, seq_len)
|
||||
# head
|
||||
|
|
@ -1599,6 +1630,7 @@ class HumoWanModel(WanModel):
|
|||
bs, _, time, height, width = x.shape
|
||||
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
grid_sizes = x.shape[2:]
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
|
|
@ -1630,6 +1662,7 @@ class HumoWanModel(WanModel):
|
|||
audio = None
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
|
|
@ -1645,6 +1678,11 @@ class HumoWanModel(WanModel):
|
|||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, audio=audio, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
|
|
@ -1660,8 +1698,14 @@ class SCAILWanModel(WanModel):
|
|||
|
||||
def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs):
|
||||
|
||||
x_input = x
|
||||
|
||||
img_offset = 0
|
||||
if reference_latent is not None:
|
||||
x = torch.cat((reference_latent, x), dim=2)
|
||||
img_offset = (reference_latent.shape[2] // self.patch_size[0]) * \
|
||||
(reference_latent.shape[3] // self.patch_size[1]) * \
|
||||
(reference_latent.shape[4] // self.patch_size[2])
|
||||
|
||||
# embeddings
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
|
|
@ -1697,6 +1741,7 @@ class SCAILWanModel(WanModel):
|
|||
context_img_len = clip_fea.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
|
|
@ -1712,6 +1757,11 @@ class SCAILWanModel(WanModel):
|
|||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
|
|
|
|||
|
|
@ -493,6 +493,7 @@ class AnimateWanModel(WanModel):
|
|||
**kwargs,
|
||||
):
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
x, motion_vec = self.after_patch_embedding(x, pose_latents, face_pixel_values)
|
||||
grid_sizes = x.shape[2:]
|
||||
|
|
@ -505,11 +506,13 @@ class AnimateWanModel(WanModel):
|
|||
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
|
||||
|
||||
full_ref = None
|
||||
img_offset = 0
|
||||
if self.ref_conv is not None:
|
||||
full_ref = kwargs.get("reference_latent", None)
|
||||
if full_ref is not None:
|
||||
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
|
||||
x = torch.concat((full_ref, x), dim=1)
|
||||
img_offset = full_ref.shape[1]
|
||||
|
||||
# context
|
||||
context = self.text_embedding(context)
|
||||
|
|
@ -522,6 +525,7 @@ class AnimateWanModel(WanModel):
|
|||
context_img_len = clip_fea.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
|
|
@ -537,6 +541,11 @@ class AnimateWanModel(WanModel):
|
|||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
if i % 5 == 0 and motion_vec is not None:
|
||||
x = x + self.face_adapter.fuser_blocks[i // 5](x, motion_vec)
|
||||
|
||||
|
|
|
|||
|
|
@ -111,6 +111,7 @@ class WanDancerModel(WanModel):
|
|||
|
||||
def forward_orig(self, x, t, context, clip_fea=None, clip_fea_ref=None, freqs=None, audio_embed=None, fps=30, audio_inject_scale=1.0, transformer_options={}, **kwargs):
|
||||
# embeddings
|
||||
x_input = x
|
||||
if int(fps + 0.5) != 30:
|
||||
x = self.patch_embedding_global(x.float()).to(x.dtype)
|
||||
else:
|
||||
|
|
@ -128,11 +129,13 @@ class WanDancerModel(WanModel):
|
|||
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
|
||||
|
||||
full_ref = None
|
||||
img_offset = 0
|
||||
if self.ref_conv is not None: # model has the weight, but this wasn't used in the original pipeline
|
||||
full_ref = kwargs.get("reference_latent", None)
|
||||
if full_ref is not None:
|
||||
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
|
||||
x = torch.concat((full_ref, x), dim=1)
|
||||
img_offset = full_ref.shape[1]
|
||||
|
||||
# context
|
||||
context = self.text_embedding(context)
|
||||
|
|
@ -163,6 +166,7 @@ class WanDancerModel(WanModel):
|
|||
context_img_len += clip_fea_ref.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
|
|
@ -177,6 +181,12 @@ class WanDancerModel(WanModel):
|
|||
x = out["img"]
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
if audio_emb is not None:
|
||||
x = self.music_injector(x, i, audio_emb, audio_emb_global=None, seq_len=seq_len, scale=audio_inject_scale)
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,149 @@
|
|||
# Uni3C controlnet for Wan 2.1: https://github.com/ewrfcas/Uni3C
|
||||
# Converted from the original diffusers based implementation.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.ldm.flux.layers import EmbedND
|
||||
from .model import WanSelfAttention
|
||||
|
||||
|
||||
class Uni3CLayerNormZero(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
conditioning_dim,
|
||||
embedding_dim,
|
||||
eps=1e-5,
|
||||
device=None, dtype=None, operations=None
|
||||
):
|
||||
super().__init__()
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = operations.Linear(conditioning_dim, 3 * embedding_dim, device=device, dtype=dtype)
|
||||
self.norm = operations.LayerNorm(embedding_dim, eps=eps, elementwise_affine=True, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x, temb):
|
||||
shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1)
|
||||
x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||
return x, gate[:, None, :]
|
||||
|
||||
|
||||
class Uni3CAttentionBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
ffn_dim,
|
||||
num_heads,
|
||||
time_embed_dim=5120,
|
||||
eps=1e-6,
|
||||
device=None, dtype=None, operations=None
|
||||
):
|
||||
super().__init__()
|
||||
operation_settings = {"operations": operations, "device": device, "dtype": dtype}
|
||||
self.norm1 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
|
||||
self.self_attn = WanSelfAttention(dim, num_heads, qk_norm=True, eps=eps, operation_settings=operation_settings)
|
||||
self.norm2 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
|
||||
self.ffn = nn.Sequential(
|
||||
operations.Linear(dim, ffn_dim, device=device, dtype=dtype), nn.GELU(approximate='tanh'),
|
||||
operations.Linear(ffn_dim, dim, device=device, dtype=dtype))
|
||||
|
||||
def forward(self, x, temb, freqs):
|
||||
norm_x, gate_msa = self.norm1(x, temb)
|
||||
x = x + gate_msa * self.self_attn(norm_x, freqs)
|
||||
norm_x, gate_ff = self.norm2(x, temb)
|
||||
x = x + gate_ff * self.ffn(norm_x)
|
||||
return x
|
||||
|
||||
|
||||
class MaskCamEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
add_channels=7,
|
||||
mid_channels=256,
|
||||
conv_out_dim=5120,
|
||||
device=None, dtype=None, operations=None
|
||||
):
|
||||
super().__init__()
|
||||
self.mask_padding = [0, 0, 0, 0, 3, 0] # first frame conditioning
|
||||
self.mask_proj = nn.Sequential(
|
||||
operations.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8), device=device, dtype=dtype),
|
||||
operations.GroupNorm(mid_channels // 8, mid_channels, device=device, dtype=dtype),
|
||||
nn.SiLU())
|
||||
self.mask_zero_proj = operations.Conv3d(mid_channels, conv_out_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), device=device, dtype=dtype)
|
||||
|
||||
def forward(self, add_inputs):
|
||||
add_padded = torch.nn.functional.pad(add_inputs, self.mask_padding, mode="constant", value=0)
|
||||
add_embeds = self.mask_proj(add_padded)
|
||||
add_embeds = self.mask_zero_proj(add_embeds)
|
||||
add_embeds = add_embeds.flatten(2).transpose(1, 2)
|
||||
return add_embeds
|
||||
|
||||
|
||||
class WanUni3CControlnet(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels=36,
|
||||
conv_out_dim=5120,
|
||||
dim=1024,
|
||||
ffn_dim=8192,
|
||||
num_heads=16,
|
||||
num_layers=20,
|
||||
time_embed_dim=5120,
|
||||
out_proj_dim=5120,
|
||||
add_channels=7,
|
||||
mid_channels=256,
|
||||
device=None, dtype=None, operations=None
|
||||
):
|
||||
super().__init__()
|
||||
patch_size = (1, 2, 2)
|
||||
self.num_layers = num_layers
|
||||
|
||||
self.controlnet_patch_embedding = operations.Conv3d(
|
||||
in_channels, conv_out_dim, kernel_size=patch_size, stride=patch_size, device=device, dtype=torch.float32)
|
||||
self.controlnet_mask_embedding = MaskCamEmbed(add_channels, mid_channels, conv_out_dim, device=device, dtype=dtype, operations=operations)
|
||||
|
||||
if conv_out_dim != dim:
|
||||
self.proj_in = operations.Linear(conv_out_dim, dim, device=device, dtype=dtype)
|
||||
else:
|
||||
self.proj_in = nn.Identity()
|
||||
|
||||
self.controlnet_blocks = nn.ModuleList([
|
||||
Uni3CAttentionBlock(dim, ffn_dim, num_heads, time_embed_dim, device=device, dtype=dtype, operations=operations)
|
||||
for _ in range(num_layers)])
|
||||
self.proj_out = nn.ModuleList([
|
||||
operations.Linear(dim, out_proj_dim, device=device, dtype=dtype)
|
||||
for _ in range(num_layers)])
|
||||
|
||||
head_dim = dim // num_heads
|
||||
self.rope_embedder = EmbedND(dim=head_dim, theta=10000.0, axes_dim=[head_dim - 4 * (head_dim // 6), 2 * (head_dim // 6), 2 * (head_dim // 6)])
|
||||
|
||||
def rope_encode(self, t_len, h_len, w_len, device=None, dtype=None):
|
||||
img_ids = torch.zeros((t_len, h_len, w_len, 3), device=device, dtype=dtype)
|
||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.arange(t_len, device=device, dtype=dtype).reshape(-1, 1, 1)
|
||||
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.arange(h_len, device=device, dtype=dtype).reshape(1, -1, 1)
|
||||
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.arange(w_len, device=device, dtype=dtype).reshape(1, 1, -1)
|
||||
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
|
||||
freqs = self.rope_embedder(img_ids).movedim(1, 2)
|
||||
return freqs
|
||||
|
||||
def process_input(self, control_input, render_mask=None, camera_embedding=None):
|
||||
# render_mask/camera_embedding are the checkpoint's extra conditioning path, not wired up yet
|
||||
hidden = self.controlnet_patch_embedding(control_input.float()).to(control_input.dtype)
|
||||
t_len, h_len, w_len = hidden.shape[2:]
|
||||
freqs = self.rope_encode(t_len, h_len, w_len, device=hidden.device, dtype=hidden.dtype)
|
||||
hidden = hidden.flatten(2).transpose(1, 2)
|
||||
|
||||
add_inputs = None
|
||||
if camera_embedding is not None and render_mask is not None:
|
||||
add_inputs = torch.cat([render_mask, camera_embedding], dim=1)
|
||||
elif render_mask is not None:
|
||||
add_inputs = render_mask
|
||||
|
||||
if add_inputs is not None:
|
||||
hidden = hidden + self.controlnet_mask_embedding(add_inputs.to(hidden.dtype))
|
||||
|
||||
hidden = self.proj_in(hidden)
|
||||
return hidden, freqs
|
||||
|
||||
def forward_block(self, block_index, hidden, temb, freqs):
|
||||
hidden = self.controlnet_blocks[block_index](hidden, temb, freqs)
|
||||
residual = self.proj_out[block_index](hidden)
|
||||
return hidden, residual
|
||||
|
|
@ -2024,11 +2024,11 @@ class WAN22_WanDancer(WAN21):
|
|||
|
||||
fps = kwargs.get("fps", None)
|
||||
if fps is not None:
|
||||
out['fps'] = comfy.conds.CONDRegular(torch.FloatTensor([fps]))
|
||||
out['fps'] = comfy.conds.CONDConstant(fps)
|
||||
|
||||
audio_inject_scale = kwargs.get("audio_inject_scale", None)
|
||||
if audio_inject_scale is not None:
|
||||
out['audio_inject_scale'] = comfy.conds.CONDRegular(torch.FloatTensor([audio_inject_scale]))
|
||||
out['audio_inject_scale'] = comfy.conds.CONDConstant(audio_inject_scale)
|
||||
return out
|
||||
|
||||
class Hunyuan3Dv2(BaseModel):
|
||||
|
|
@ -2227,10 +2227,7 @@ class Omnigen2(BaseModel):
|
|||
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
|
||||
ref_latents = kwargs.get("reference_latents", None)
|
||||
if ref_latents is not None:
|
||||
latents = []
|
||||
for lat in ref_latents:
|
||||
latents.append(self.process_latent_in(lat))
|
||||
out['ref_latents'] = comfy.conds.CONDList(latents)
|
||||
out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents])
|
||||
return out
|
||||
|
||||
def extra_conds_shapes(self, **kwargs):
|
||||
|
|
@ -2317,12 +2314,30 @@ class Ideogram4(BaseModel):
|
|||
class Krea2(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.FLUX, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.krea2.model.SingleStreamDiT)
|
||||
self.memory_usage_factor_conds = ("ref_latents",)
|
||||
|
||||
def extra_conds(self, **kwargs):
|
||||
out = super().extra_conds(**kwargs)
|
||||
cross_attn = kwargs.get("cross_attn", None)
|
||||
if cross_attn is not None:
|
||||
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
|
||||
ref_latents = kwargs.get("reference_latents", None)
|
||||
if ref_latents is not None:
|
||||
latents = []
|
||||
for lat in ref_latents:
|
||||
latents.append(self.process_latent_in(lat))
|
||||
out['ref_latents'] = comfy.conds.CONDList(latents)
|
||||
|
||||
ref_latents_method = kwargs.get("reference_latents_method", None)
|
||||
if ref_latents_method is not None:
|
||||
out['ref_latents_method'] = comfy.conds.CONDConstant(ref_latents_method)
|
||||
return out
|
||||
|
||||
def extra_conds_shapes(self, **kwargs):
|
||||
out = {}
|
||||
ref_latents = kwargs.get("reference_latents", None)
|
||||
if ref_latents is not None:
|
||||
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
|
||||
return out
|
||||
|
||||
class HunyuanImage21(BaseModel):
|
||||
|
|
|
|||
|
|
@ -473,7 +473,7 @@ except:
|
|||
|
||||
SUPPORT_FP8_OPS = args.supports_fp8_compute
|
||||
|
||||
AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]
|
||||
AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1035", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]
|
||||
AMD_ENABLE_MIOPEN_ENV = 'COMFYUI_ENABLE_MIOPEN'
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -1434,6 +1434,7 @@ class TEModel(Enum):
|
|||
GPT_OSS_20B = 33
|
||||
QWEN3VL_4B = 34
|
||||
QWEN3VL_8B = 35
|
||||
GEMMA_4_12B = 36
|
||||
|
||||
|
||||
def detect_te_model(sd):
|
||||
|
|
@ -1463,6 +1464,9 @@ def detect_te_model(sd):
|
|||
if 'model.layers.0.post_feedforward_layernorm.weight' in sd:
|
||||
if 'model.layers.59.self_attn.q_norm.weight' in sd:
|
||||
return TEModel.GEMMA_4_31B
|
||||
# Gemma4 12B Unified: 48 layers, encoder-free; global layers drop v_proj (attention_k_eq_v).
|
||||
if 'model.layers.47.self_attn.q_norm.weight' in sd and 'model.layers.5.self_attn.v_proj.weight' not in sd:
|
||||
return TEModel.GEMMA_4_12B
|
||||
if 'model.layers.41.self_attn.q_norm.weight' in sd and 'model.layers.47.self_attn.q_norm.weight' not in sd:
|
||||
return TEModel.GEMMA_4_E4B
|
||||
if 'model.layers.34.self_attn.q_norm.weight' in sd and 'model.layers.41.self_attn.q_norm.weight' not in sd:
|
||||
|
|
@ -1618,10 +1622,11 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
|
|||
clip_target.clip = comfy.text_encoders.sa3.SAT5GemmaModel
|
||||
clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer
|
||||
tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None)
|
||||
elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B):
|
||||
elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B, TEModel.GEMMA_4_12B):
|
||||
variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B,
|
||||
TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B,
|
||||
TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B}[te_model]
|
||||
TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B,
|
||||
TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model]
|
||||
clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant)
|
||||
clip_target.tokenizer = variant.tokenizer
|
||||
tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None)
|
||||
|
|
|
|||
|
|
@ -1,11 +1,15 @@
|
|||
import torch
|
||||
import torch.nn as nn
|
||||
import torchaudio.functional as AF
|
||||
import torchvision.transforms.functional as TVF
|
||||
import numpy as np
|
||||
from tokenizers import Tokenizer
|
||||
from dataclasses import dataclass
|
||||
import math
|
||||
|
||||
from comfy import sd1_clip
|
||||
import comfy.model_management
|
||||
import comfy.ops
|
||||
from comfy.ldm.modules.attention import optimized_attention_for_device
|
||||
from comfy.rmsnorm import rms_norm
|
||||
from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding
|
||||
|
|
@ -21,6 +25,10 @@ GEMMA4_VISION_CONFIG = {"hidden_size": 768, "image_size": 896, "intermediate_siz
|
|||
GEMMA4_VISION_31B_CONFIG = {"hidden_size": 1152, "image_size": 896, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 16, "head_dim": 72, "rms_norm_eps": 1e-6, "position_embedding_size": 10240, "pooling_kernel_size": 3}
|
||||
GEMMA4_AUDIO_CONFIG = {"hidden_size": 1024, "num_hidden_layers": 12, "num_attention_heads": 8, "intermediate_size": 4096, "conv_kernel_size": 5, "attention_chunk_size": 12, "attention_context_left": 13, "attention_context_right": 0, "attention_logit_cap": 50.0, "output_proj_dims": 1536, "rms_norm_eps": 1e-6, "residual_weight": 0.5}
|
||||
|
||||
# Encoder-free (gemma4_unified) multimodal embedders: raw patches/waveform projected directly into LM space.
|
||||
GEMMA4_UNIFIED_VISION_CONFIG = {"model_patch_size": 48, "patch_size": 16, "pooling_kernel_size": 3, "mm_embed_dim": 3840, "mm_posemb_size": 1120, "output_proj_dims": 3840, "rms_norm_eps": 1e-6}
|
||||
GEMMA4_UNIFIED_AUDIO_CONFIG = {"audio_samples_per_token": 640, "output_proj_dims": 640, "rms_norm_eps": 1e-6}
|
||||
|
||||
@dataclass
|
||||
class Gemma4Config:
|
||||
vocab_size: int = 262144
|
||||
|
|
@ -35,6 +43,9 @@ class Gemma4Config:
|
|||
transformer_type: str = "gemma4"
|
||||
head_dim = 256
|
||||
global_head_dim = 512
|
||||
num_global_key_value_heads = None
|
||||
attention_k_eq_v = False
|
||||
vision_bidirectional = False
|
||||
rms_norm_add = False
|
||||
mlp_activation = "gelu_pytorch_tanh"
|
||||
qkv_bias = False
|
||||
|
|
@ -51,6 +62,7 @@ class Gemma4Config:
|
|||
num_kv_shared_layers: int = 18
|
||||
use_double_wide_mlp: bool = False
|
||||
stop_tokens = [1, 50, 106]
|
||||
suppress_tokens = []
|
||||
vision_config = GEMMA4_VISION_CONFIG
|
||||
audio_config = GEMMA4_AUDIO_CONFIG
|
||||
mm_tokens_per_image = 280
|
||||
|
|
@ -72,12 +84,30 @@ class Gemma4_31B_Config(Gemma4Config):
|
|||
num_hidden_layers: int = 60
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: int = 16
|
||||
vision_bidirectional = True
|
||||
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
|
||||
hidden_size_per_layer_input: int = 0
|
||||
num_kv_shared_layers: int = 0
|
||||
audio_config = None
|
||||
vision_config = GEMMA4_VISION_31B_CONFIG
|
||||
|
||||
@dataclass
|
||||
class Gemma4_12B_Config(Gemma4Config):
|
||||
hidden_size: int = 3840
|
||||
intermediate_size: int = 15360
|
||||
num_hidden_layers: int = 48
|
||||
num_attention_heads: int = 16
|
||||
num_key_value_heads: int = 8
|
||||
num_global_key_value_heads = 1
|
||||
attention_k_eq_v = True
|
||||
vision_bidirectional = True
|
||||
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
|
||||
hidden_size_per_layer_input: int = 0
|
||||
num_kv_shared_layers: int = 0
|
||||
audio_config = GEMMA4_UNIFIED_AUDIO_CONFIG
|
||||
vision_config = GEMMA4_UNIFIED_VISION_CONFIG
|
||||
suppress_tokens = [258883, 258882]
|
||||
|
||||
|
||||
# unfused RoPE as addcmul_ RoPE diverges from reference code
|
||||
def _apply_rotary_pos_emb(x, freqs_cis):
|
||||
|
|
@ -89,17 +119,18 @@ def _apply_rotary_pos_emb(x, freqs_cis):
|
|||
return out
|
||||
|
||||
class Gemma4Attention(nn.Module):
|
||||
def __init__(self, config, head_dim, device=None, dtype=None, ops=None):
|
||||
def __init__(self, config, head_dim, num_kv_heads=None, k_eq_v=False, device=None, dtype=None, ops=None):
|
||||
super().__init__()
|
||||
self.num_heads = config.num_attention_heads
|
||||
self.num_kv_heads = config.num_key_value_heads
|
||||
self.num_kv_heads = num_kv_heads if num_kv_heads is not None else config.num_key_value_heads
|
||||
self.hidden_size = config.hidden_size
|
||||
self.head_dim = head_dim
|
||||
self.inner_size = self.num_heads * head_dim
|
||||
|
||||
self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype)
|
||||
self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
|
||||
self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
|
||||
# k_eq_v: V reuses the K projection (no separate v_proj weight)
|
||||
self.v_proj = None if k_eq_v else ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
|
||||
self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype)
|
||||
|
||||
self.q_norm = None
|
||||
|
|
@ -133,7 +164,10 @@ class Gemma4Attention(nn.Module):
|
|||
shareable_kv = None
|
||||
else:
|
||||
xk = self.k_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
|
||||
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
|
||||
if self.v_proj is not None:
|
||||
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
|
||||
else:
|
||||
xv = xk # k_eq_v: V is the raw K projection (before k_norm/RoPE)
|
||||
if self.k_norm is not None:
|
||||
xk = self.k_norm(xk)
|
||||
xv = rms_norm(xv)
|
||||
|
|
@ -186,7 +220,10 @@ class TransformerBlockGemma4(nn.Module):
|
|||
|
||||
head_dim = config.head_dim if self.sliding_attention else config.global_head_dim
|
||||
|
||||
self.self_attn = Gemma4Attention(config, head_dim=head_dim, device=device, dtype=dtype, ops=ops)
|
||||
# k_eq_v only on global layers, which then use num_global_key_value_heads
|
||||
k_eq_v = config.attention_k_eq_v and not self.sliding_attention
|
||||
num_kv_heads = config.num_global_key_value_heads if k_eq_v else config.num_key_value_heads
|
||||
self.self_attn = Gemma4Attention(config, head_dim=head_dim, num_kv_heads=num_kv_heads, k_eq_v=k_eq_v, device=device, dtype=dtype, ops=ops)
|
||||
|
||||
num_kv_shared = config.num_kv_shared_layers
|
||||
first_kv_shared = config.num_hidden_layers - num_kv_shared
|
||||
|
|
@ -203,9 +240,9 @@ class TransformerBlockGemma4(nn.Module):
|
|||
self.per_layer_input_gate = ops.Linear(config.hidden_size, self.hidden_size_per_layer_input, bias=False, device=device, dtype=dtype)
|
||||
self.per_layer_projection = ops.Linear(self.hidden_size_per_layer_input, config.hidden_size, bias=False, device=device, dtype=dtype)
|
||||
self.post_per_layer_input_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype)
|
||||
self.register_buffer("layer_scalar", torch.ones(1, device=device, dtype=dtype))
|
||||
else:
|
||||
self.layer_scalar = None
|
||||
|
||||
# layer_scalar exists on every gemma4 variant, independent of per-layer input
|
||||
self.register_buffer("layer_scalar", torch.empty(1, device=device, dtype=dtype))
|
||||
|
||||
def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, per_layer_input=None, shared_kv=None):
|
||||
sliding_window = None
|
||||
|
|
@ -244,8 +281,7 @@ class TransformerBlockGemma4(nn.Module):
|
|||
x = self.post_per_layer_input_norm(x)
|
||||
x = residual + x
|
||||
|
||||
if self.layer_scalar is not None:
|
||||
x = x * self.layer_scalar
|
||||
x = x * comfy.ops.cast_to_input(self.layer_scalar, x)
|
||||
|
||||
return x, present_key_value, shareable_kv
|
||||
|
||||
|
|
@ -334,6 +370,19 @@ class Gemma4Transformer(nn.Module):
|
|||
causal_mask.masked_fill_(torch.ones_like(causal_mask, dtype=torch.bool).triu_(1), min_val)
|
||||
mask = mask + causal_mask if mask is not None else causal_mask
|
||||
|
||||
# Bidirectional attention within each image soft-token block (prefill only; text/audio stay causal).
|
||||
if self.config.vision_bidirectional and past_len == 0 and embeds_info:
|
||||
block_ids = torch.full((seq_len,), -1, dtype=torch.long, device=x.device)
|
||||
group = 0
|
||||
for info in embeds_info:
|
||||
if info.get("type") == "image":
|
||||
start = info["index"]
|
||||
block_ids[start:start + info["size"]] = group
|
||||
group += 1
|
||||
if group > 0:
|
||||
same_block = (block_ids[:, None] == block_ids[None, :]) & (block_ids[:, None] >= 0)
|
||||
mask = mask.masked_fill(same_block, 0.0)
|
||||
|
||||
# Per-layer inputs
|
||||
per_layer_inputs = None
|
||||
if self.hidden_size_per_layer_input:
|
||||
|
|
@ -354,8 +403,24 @@ class Gemma4Transformer(nn.Module):
|
|||
shared_global_kv = None # KV from last non-shared global layer
|
||||
|
||||
intermediate = None
|
||||
all_intermediate = None
|
||||
only_layers = None
|
||||
if intermediate_output is not None:
|
||||
if isinstance(intermediate_output, list):
|
||||
all_intermediate = []
|
||||
only_layers = {len(self.layers) + layer if layer < 0 else layer for layer in intermediate_output}
|
||||
elif intermediate_output == "all":
|
||||
all_intermediate = []
|
||||
intermediate_output = None
|
||||
elif intermediate_output < 0:
|
||||
intermediate_output = len(self.layers) + intermediate_output
|
||||
|
||||
next_key_values = []
|
||||
for i, layer in enumerate(self.layers):
|
||||
if all_intermediate is not None:
|
||||
if only_layers is None or (i in only_layers):
|
||||
all_intermediate.append(x.unsqueeze(1).clone())
|
||||
|
||||
past_kv = past_key_values[i] if past_key_values is not None and len(past_key_values) > 0 else None
|
||||
|
||||
layer_kwargs = {}
|
||||
|
|
@ -385,7 +450,18 @@ class Gemma4Transformer(nn.Module):
|
|||
if self.norm is not None:
|
||||
x = self.norm(x)
|
||||
|
||||
if len(next_key_values) > 0:
|
||||
if all_intermediate is not None:
|
||||
if only_layers is None or (len(self.layers) in only_layers):
|
||||
all_intermediate.append(x.unsqueeze(1).clone())
|
||||
if len(all_intermediate) > 0:
|
||||
intermediate = torch.cat(all_intermediate, dim=1)
|
||||
|
||||
if intermediate is not None and final_layer_norm_intermediate and self.norm is not None:
|
||||
intermediate = self.norm(intermediate)
|
||||
|
||||
# Only hand back the KV cache when caching was actually requested; SDClipModel reads
|
||||
# outputs[2] as the pooled output.
|
||||
if past_key_values is not None and len(next_key_values) > 0:
|
||||
return x, intermediate, next_key_values
|
||||
return x, intermediate
|
||||
|
||||
|
|
@ -404,6 +480,8 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module):
|
|||
cap = self.model.config.final_logit_softcapping
|
||||
if cap:
|
||||
logits = cap * torch.tanh(logits / cap)
|
||||
if self.model.config.suppress_tokens:
|
||||
logits[..., self.model.config.suppress_tokens] = torch.finfo(logits.dtype).min
|
||||
return logits
|
||||
|
||||
def init_kv_cache(self, batch, max_cache_len, device, execution_dtype):
|
||||
|
|
@ -441,6 +519,28 @@ class Gemma4AudioMixin:
|
|||
return None, None
|
||||
|
||||
|
||||
class Gemma4UnifiedBase(Gemma4Base):
|
||||
"""Encoder-free multimodal Gemma4 (gemma4_unified, e.g. 12B): raw image patches and audio frames projected directly into LM space."""
|
||||
def _init_model(self, config, dtype, device, operations):
|
||||
self.num_layers = config.num_hidden_layers
|
||||
self.model = Gemma4Transformer(config, device=device, dtype=dtype, ops=operations)
|
||||
self.dtype = dtype
|
||||
self.vision_model = Gemma4UnifiedVisionEmbedder(config.vision_config, device=device, dtype=dtype, ops=operations)
|
||||
self.multi_modal_projector = Gemma4RMSNormProjector(config.vision_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
|
||||
self.audio_projector = Gemma4RMSNormProjector(config.audio_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
|
||||
|
||||
def preprocess_embed(self, embed, device):
|
||||
if embed["type"] == "image":
|
||||
pixels = embed.pop("data").movedim(-1, 1).to(device, dtype=self.dtype) # [B, H, W, C] -> [B, C, H, W], [0,1]
|
||||
patches, positions = self.vision_model.patchify(pixels)
|
||||
vision_out = self.vision_model(patches, positions)
|
||||
return self.multi_modal_projector(vision_out), None
|
||||
if embed["type"] == "audio":
|
||||
audio = embed.pop("data").to(device, dtype=self.dtype) # [1, T, audio_samples_per_token]
|
||||
return self.audio_projector(audio), None
|
||||
return None, None
|
||||
|
||||
|
||||
# Vision Encoder
|
||||
|
||||
def _compute_vision_2d_rope(head_dim, pixel_position_ids, theta=100.0, device=None):
|
||||
|
|
@ -713,6 +813,73 @@ class Gemma4MultiModalProjector(Gemma4RMSNormProjector):
|
|||
super().__init__(config.vision_config["hidden_size"], config.hidden_size, dtype=dtype, device=device, ops=ops)
|
||||
|
||||
|
||||
# Encoder-free vision (gemma4_unified): raw merged pixel patches projected directly into LM space.
|
||||
|
||||
def _patches_merge(patches, positions_xy, length):
|
||||
patch_size = math.isqrt(patches.shape[-1] // 3)
|
||||
k = math.isqrt(patches.shape[-2] // length)
|
||||
batch = patches.shape[:-2]
|
||||
|
||||
max_x = positions_xy[..., 0].max(dim=-1, keepdim=True)[0] + 1
|
||||
kidx = torch.div(positions_xy, k, rounding_mode="floor")
|
||||
rem = torch.remainder(positions_xy, k)
|
||||
order = rem[..., 0] + rem[..., 1] * k + k * k * kidx[..., 0] + k * max_x * kidx[..., 1]
|
||||
perm = order.long().argsort(dim=-1)
|
||||
|
||||
merged = patches.gather(-2, perm.unsqueeze(-1).expand_as(patches))
|
||||
merged = merged.reshape(*batch, length, k, k, patch_size, patch_size, 3)
|
||||
merged = merged.permute(*range(len(batch)), -6, -5, -3, -4, -2, -1).reshape(*batch, length, (k * patch_size) ** 2 * 3)
|
||||
|
||||
pos = positions_xy.gather(-2, perm.unsqueeze(-1).expand_as(positions_xy))
|
||||
pad = (positions_xy == -1).all(dim=-1, keepdim=True)
|
||||
pos = torch.where(pad, positions_xy, pos).reshape(*batch, length, k * k, 2)
|
||||
pos = torch.div(pos, k, rounding_mode="floor").min(dim=-2)[0]
|
||||
return merged, pos
|
||||
|
||||
|
||||
class Gemma4UnifiedVisionEmbedder(nn.Module):
|
||||
"""Encoder-free patch embedder (LN -> Dense -> LN -> +2D posemb -> LN); projection to text space is the separate multi_modal_projector."""
|
||||
def __init__(self, config, device=None, dtype=None, ops=None):
|
||||
super().__init__()
|
||||
self.patch_size = config["patch_size"]
|
||||
self.pooling_kernel_size = config["pooling_kernel_size"]
|
||||
patch_dim = config["model_patch_size"] ** 2 * 3
|
||||
mm_embed_dim = config["mm_embed_dim"]
|
||||
self.patch_ln1 = ops.LayerNorm(patch_dim, device=device, dtype=dtype)
|
||||
self.patch_dense = ops.Linear(patch_dim, mm_embed_dim, device=device, dtype=dtype)
|
||||
self.patch_ln2 = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
|
||||
self.pos_embedding = nn.Parameter(torch.empty(config["mm_posemb_size"], 2, mm_embed_dim, device=device, dtype=dtype))
|
||||
self.pos_norm = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
|
||||
|
||||
def patchify(self, pixels):
|
||||
"""pixels: [B, C, H, W] in [0,1] -> merged patches [B, N, 6912], positions [B, N, 2]."""
|
||||
ps, k = self.patch_size, self.pooling_kernel_size
|
||||
out_patches, out_positions = [], []
|
||||
for img in pixels:
|
||||
ph, pw = img.shape[-2] // ps, img.shape[-1] // ps
|
||||
teacher = img.reshape(img.shape[0], ph, ps, pw, ps).permute(1, 3, 2, 4, 0).reshape(ph * pw, -1)
|
||||
grid = torch.meshgrid(torch.arange(pw, device=img.device), torch.arange(ph, device=img.device), indexing="xy")
|
||||
tpos = torch.stack(grid, dim=-1).reshape(teacher.shape[0], 2)
|
||||
n_model = teacher.shape[0] // (k * k)
|
||||
mp, mpos = _patches_merge(teacher.unsqueeze(0), tpos.unsqueeze(0), n_model)
|
||||
out_patches.append(mp.squeeze(0))
|
||||
out_positions.append(mpos.squeeze(0))
|
||||
return torch.stack(out_patches), torch.stack(out_positions)
|
||||
|
||||
def forward(self, pixel_values, image_position_ids):
|
||||
x = self.patch_ln1(pixel_values)
|
||||
x = self.patch_dense(x)
|
||||
x = self.patch_ln2(x)
|
||||
|
||||
clamped = image_position_ids.clamp(min=0).long()
|
||||
valid = (image_position_ids != -1).to(x.dtype).unsqueeze(-1)
|
||||
axes = torch.arange(2, device=image_position_ids.device)
|
||||
pos = comfy.model_management.cast_to_device(self.pos_embedding, x.device, x.dtype)
|
||||
pos_embs = (pos[clamped, axes] * valid).sum(-2)
|
||||
x = x + pos_embs
|
||||
return self.pos_norm(x)
|
||||
|
||||
|
||||
# Audio Encoder
|
||||
|
||||
class Gemma4AudioConvSubsampler(nn.Module):
|
||||
|
|
@ -990,6 +1157,30 @@ class Gemma4AudioProjector(Gemma4RMSNormProjector):
|
|||
|
||||
# Tokenizer and Wrappers
|
||||
|
||||
def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, pooling_kernel_size):
|
||||
target_px = max_patches * patch_size ** 2
|
||||
factor = math.sqrt(target_px / (height * width))
|
||||
side_mult = pooling_kernel_size * patch_size
|
||||
target_height = math.floor(factor * height / side_mult) * side_mult
|
||||
target_width = math.floor(factor * width / side_mult) * side_mult
|
||||
|
||||
if target_height == 0 and target_width == 0:
|
||||
raise ValueError(f"Attempting to resize to a 0 x 0 image. Resized height should be divisible by {side_mult}.")
|
||||
|
||||
max_side_length = (max_patches // pooling_kernel_size ** 2) * side_mult
|
||||
if target_height == 0:
|
||||
target_height = side_mult
|
||||
target_width = min(math.floor(width / height) * side_mult, max_side_length)
|
||||
elif target_width == 0:
|
||||
target_width = side_mult
|
||||
target_height = min(math.floor(height / width) * side_mult, max_side_length)
|
||||
|
||||
if target_height * target_width > target_px:
|
||||
raise ValueError(f"Resizing [{height}x{width}] to [{target_height}x{target_width}] exceeds the patch budget.")
|
||||
|
||||
return target_height, target_width
|
||||
|
||||
|
||||
class Gemma4_Tokenizer():
|
||||
tokenizer_json_data = None
|
||||
|
||||
|
|
@ -998,25 +1189,35 @@ class Gemma4_Tokenizer():
|
|||
return {"tokenizer_json": self.tokenizer_json_data}
|
||||
return {}
|
||||
|
||||
def _extract_mel_spectrogram(self, waveform, sample_rate):
|
||||
"""Extract 128-bin log mel spectrogram.
|
||||
Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
|
||||
"""
|
||||
# Mix to mono first, then resample to 16kHz
|
||||
def _audio_token_count(self, num_samples):
|
||||
# Default (E2B/E4B): mel frames after two stride-2 conv subsamples.
|
||||
_fl = 320 # int(round(16000 * 20.0 / 1000.0))
|
||||
_hl = 160 # int(round(16000 * 10.0 / 1000.0))
|
||||
_nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
|
||||
_t = _nmel
|
||||
for _ in range(2):
|
||||
_t = (_t + 2 - 3) // 2 + 1
|
||||
return min(_t, 750)
|
||||
|
||||
@staticmethod
|
||||
def _resample_16k(waveform, sample_rate):
|
||||
"""Mix to mono and resample to 16kHz. Kaiser params reproduce the reference (transformers
|
||||
load_audio -> librosa/soxr_hq) to ~1e-12 MSE using only torchaudio."""
|
||||
if waveform.dim() > 1 and waveform.shape[0] > 1:
|
||||
waveform = waveform.mean(dim=0, keepdim=True)
|
||||
if waveform.dim() == 1:
|
||||
waveform = waveform.unsqueeze(0)
|
||||
audio = waveform.squeeze(0).float().numpy()
|
||||
audio = waveform.float()
|
||||
if sample_rate != 16000:
|
||||
# Use scipy's resample_poly with a high-quality FIR filter to get as close as possible to librosa's resampling (while still not full match)
|
||||
from scipy.signal import resample_poly, firwin
|
||||
from math import gcd
|
||||
g = gcd(sample_rate, 16000)
|
||||
up, down = 16000 // g, sample_rate // g
|
||||
L = max(up, down)
|
||||
h = firwin(160 * L + 1, 0.96 / L, window=('kaiser', 6.5))
|
||||
audio = resample_poly(audio, up, down, window=h).astype(np.float32)
|
||||
audio = AF.resample(audio, sample_rate, 16000, resampling_method="sinc_interp_kaiser",
|
||||
lowpass_filter_width=121, rolloff=0.9568384289091556, beta=21.01531462440614)
|
||||
return audio.squeeze(0).contiguous()
|
||||
|
||||
def _extract_audio_features(self, waveform, sample_rate):
|
||||
"""Default (E2B/E4B): 128-bin log mel spectrogram for the conformer audio encoder.
|
||||
Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
|
||||
"""
|
||||
audio = self._resample_16k(waveform, sample_rate).numpy()
|
||||
n = len(audio)
|
||||
|
||||
# Pad to multiple of 128, build sample-level mask
|
||||
|
|
@ -1064,8 +1265,8 @@ class Gemma4_Tokenizer():
|
|||
if audio is not None:
|
||||
waveform = audio["waveform"].squeeze(0) if hasattr(audio, "__getitem__") else audio
|
||||
sample_rate = audio.get("sample_rate", 16000) if hasattr(audio, "get") else 16000
|
||||
mel, mel_mask = self._extract_mel_spectrogram(waveform, sample_rate)
|
||||
audio_features = [(mel.unsqueeze(0), mel_mask.unsqueeze(0))] # ([1, T, 128], [1, T])
|
||||
feat, feat_mask = self._extract_audio_features(waveform, sample_rate)
|
||||
audio_features = [(feat.unsqueeze(0), feat_mask.unsqueeze(0))] # ([1, T, D], [1, T])
|
||||
|
||||
# Process image/video frames
|
||||
is_video = video is not None
|
||||
|
|
@ -1090,13 +1291,8 @@ class Gemma4_Tokenizer():
|
|||
pooling_k = 3
|
||||
max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280)
|
||||
max_patches = max_soft_tokens * pooling_k * pooling_k
|
||||
target_px = max_patches * patch_size * patch_size
|
||||
factor = (target_px / (h * w)) ** 0.5
|
||||
side_mult = pooling_k * patch_size
|
||||
target_h = max(int(factor * h // side_mult) * side_mult, side_mult)
|
||||
target_w = max(int(factor * w // side_mult) * side_mult, side_mult)
|
||||
target_h, target_w = _get_aspect_ratio_preserving_size(h, w, patch_size, max_patches, pooling_k)
|
||||
|
||||
import torchvision.transforms.functional as TVF
|
||||
for i in range(num_frames):
|
||||
# rescaling to match reference code
|
||||
s = (samples[i].clamp(0, 1) * 255).to(torch.uint8) # [C, H, W] uint8
|
||||
|
|
@ -1115,7 +1311,7 @@ class Gemma4_Tokenizer():
|
|||
llama_text = llama_template.format(text)
|
||||
else:
|
||||
# Build template from modalities present
|
||||
system = "<|turn>system\n<|think|><turn|>\n" if thinking else ""
|
||||
system = "<|turn>system\n<|think|>\n<turn|>\n" if thinking else ""
|
||||
media = ""
|
||||
if len(images) > 0:
|
||||
if is_video:
|
||||
|
|
@ -1135,15 +1331,11 @@ class Gemma4_Tokenizer():
|
|||
if len(audio_features) > 0:
|
||||
# Compute audio token count (always at 16kHz)
|
||||
num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1]
|
||||
_fl = 320 # int(round(16000 * 20.0 / 1000.0))
|
||||
_hl = 160 # int(round(16000 * 10.0 / 1000.0))
|
||||
_nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
|
||||
_t = _nmel
|
||||
for _ in range(2):
|
||||
_t = (_t + 2 - 3) // 2 + 1
|
||||
n_audio_tokens = min(_t, 750)
|
||||
n_audio_tokens = self._audio_token_count(num_samples)
|
||||
media += "<|audio>" + "<|audio|>" * n_audio_tokens + "<audio|>"
|
||||
llama_text = f"{system}<|turn>user\n{media}{text}<turn|>\n<|turn>model\n"
|
||||
# Non-thinking mode primes an empty thought channel so the model answers directly.
|
||||
model_open = "" if thinking else "<|channel>thought\n<channel|>"
|
||||
llama_text = f"{system}<|turn>user\n{text}{media}<turn|>\n<|turn>model\n{model_open}"
|
||||
|
||||
text_tokens = super().tokenize_with_weights(llama_text, return_word_ids)
|
||||
|
||||
|
|
@ -1178,7 +1370,6 @@ class Gemma4_Tokenizer():
|
|||
class _Gemma4Tokenizer:
|
||||
"""Tokenizer using the tokenizers (Gemma4 doesn't come with sentencepiece model)"""
|
||||
def __init__(self, tokenizer_json_bytes=None, **kwargs):
|
||||
from tokenizers import Tokenizer
|
||||
if isinstance(tokenizer_json_bytes, torch.Tensor):
|
||||
tokenizer_json_bytes = bytes(tokenizer_json_bytes.tolist())
|
||||
self.tokenizer = Tokenizer.from_str(tokenizer_json_bytes.decode("utf-8"))
|
||||
|
|
@ -1224,6 +1415,30 @@ class Gemma4Tokenizer(sd1_clip.SD1Tokenizer):
|
|||
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma4", tokenizer=self.tokenizer_class)
|
||||
|
||||
|
||||
class Gemma4UnifiedSDTokenizer(Gemma4SDTokenizer):
|
||||
"""Encoder-free (gemma4_unified) audio: raw 16kHz waveform frames instead of mel spectrogram."""
|
||||
embedding_size = 3840
|
||||
|
||||
def _extract_audio_features(self, waveform, sample_rate):
|
||||
audio = self._resample_16k(waveform, sample_rate)
|
||||
spt = 640 # audio_samples_per_token (40ms at 16kHz)
|
||||
pad = (-audio.shape[0]) % spt
|
||||
if pad:
|
||||
audio = torch.nn.functional.pad(audio, (0, pad))
|
||||
num_tokens = audio.shape[0] // spt
|
||||
feats = audio[:num_tokens * spt].reshape(num_tokens, spt)
|
||||
feats = feats[:750] # audio_seq_length cap (matches reference truncation, ~30s)
|
||||
mask = torch.ones(feats.shape[0], dtype=torch.bool)
|
||||
return feats, mask
|
||||
|
||||
def _audio_token_count(self, num_samples):
|
||||
return min((num_samples + 639) // 640, 750)
|
||||
|
||||
|
||||
class Gemma4UnifiedTokenizer(Gemma4Tokenizer):
|
||||
tokenizer_class = Gemma4UnifiedSDTokenizer
|
||||
|
||||
|
||||
# Model wrappers
|
||||
class Gemma4Model(sd1_clip.SDClipModel):
|
||||
model_class = None
|
||||
|
|
@ -1256,7 +1471,7 @@ class Gemma4Model(sd1_clip.SDClipModel):
|
|||
expanded_idx += 1
|
||||
initial_token_ids = [ids]
|
||||
input_ids = torch.tensor(initial_token_ids, device=self.execution_device)
|
||||
return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids)
|
||||
return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids, embeds_info=embeds_info)
|
||||
|
||||
|
||||
def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None):
|
||||
|
|
@ -1296,3 +1511,11 @@ def _make_variant(config_cls):
|
|||
Gemma4_E4B = _make_variant(Gemma4Config)
|
||||
Gemma4_E2B = _make_variant(Gemma4_E2B_Config)
|
||||
Gemma4_31B = _make_variant(Gemma4_31B_Config)
|
||||
|
||||
|
||||
# Gemma4 12B Unified: encoder-free multimodal, distinct base/tokenizer (not via _make_variant).
|
||||
class Gemma4_12B(Gemma4UnifiedBase):
|
||||
def __init__(self, config_dict, dtype, device, operations):
|
||||
super().__init__()
|
||||
self._init_model(Gemma4_12B_Config(**config_dict), dtype, device, operations)
|
||||
Gemma4_12B.tokenizer = Gemma4UnifiedTokenizer
|
||||
|
|
|
|||
|
|
@ -876,7 +876,7 @@ class BaseGenerate:
|
|||
torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), 0))
|
||||
return past_key_values
|
||||
|
||||
def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None):
|
||||
def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None, embeds_info=None):
|
||||
device = embeds.device
|
||||
|
||||
if stop_tokens is None:
|
||||
|
|
@ -911,7 +911,7 @@ class BaseGenerate:
|
|||
if step == 0 and deepstack_embeds is not None:
|
||||
extra["deepstack_embeds"] = deepstack_embeds
|
||||
extra["visual_pos_masks"] = visual_pos_masks
|
||||
x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra)
|
||||
x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra, embeds_info=(embeds_info if step == 0 else None))
|
||||
logits = self.logits(x)[:, -1]
|
||||
next_token = self.sample_token(logits, temperature, top_k, top_p, min_p, repetition_penalty, initial_tokens + generated_token_ids, generator, do_sample=do_sample, presence_penalty=presence_penalty)
|
||||
token_id = next_token[0].item()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from datetime import date
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
|
@ -242,3 +242,60 @@ class GeminiGenerateContentResponse(BaseModel):
|
|||
promptFeedback: GeminiPromptFeedback | None = Field(None)
|
||||
usageMetadata: GeminiUsageMetadata | None = Field(None)
|
||||
modelVersion: str | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionTextPart(BaseModel):
|
||||
type: Literal["text"] = "text"
|
||||
text: str = Field(...)
|
||||
|
||||
|
||||
class GeminiInteractionMediaPart(BaseModel):
|
||||
type: str = Field(..., description="One of: image, video, audio, document.")
|
||||
data: str | None = Field(None, description="Base64-encoded media bytes.")
|
||||
uri: str | None = Field(None, description="URI of the media, as an alternative to inline data.")
|
||||
mime_type: str | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionGenerationConfig(BaseModel):
|
||||
temperature: float | None = Field(None, ge=0.0, le=2.0)
|
||||
top_p: float | None = Field(None, ge=0.0, le=1.0)
|
||||
|
||||
|
||||
class GeminiInteractionRequest(BaseModel):
|
||||
model: str = Field(...)
|
||||
input: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = Field(...)
|
||||
generation_config: GeminiInteractionGenerationConfig | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionModalityTokens(BaseModel):
|
||||
modality: str | None = Field(None, description="One of: text, image, audio, video, document.")
|
||||
tokens: int | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionUsage(BaseModel):
|
||||
input_tokens_by_modality: list[GeminiInteractionModalityTokens] | None = Field(None)
|
||||
output_tokens_by_modality: list[GeminiInteractionModalityTokens] | None = Field(None)
|
||||
total_thought_tokens: int | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionContent(BaseModel):
|
||||
type: str | None = Field(None)
|
||||
text: str | None = Field(None)
|
||||
data: str | None = Field(None)
|
||||
uri: str | None = Field(None)
|
||||
mime_type: str | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionStep(BaseModel):
|
||||
type: str | None = Field(None)
|
||||
content: list[GeminiInteractionContent] | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteraction(BaseModel):
|
||||
id: str | None = Field(None)
|
||||
status: str | None = Field(
|
||||
None,
|
||||
description="One of: in_progress, requires_action, completed, failed, cancelled, incomplete.",
|
||||
)
|
||||
steps: list[GeminiInteractionStep] | None = Field(None)
|
||||
usage: GeminiInteractionUsage | None = Field(None)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,452 @@
|
|||
# (label, avatar_id, avatar_type, supported engines)
|
||||
HEYGEN_AVATAR_LOOKS: list[tuple[str, str, str, tuple[str, ...]]] = [
|
||||
(
|
||||
"Annie Lounge Standing Side",
|
||||
"Annie_Lounge_Standing_Side_public",
|
||||
"studio_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Yara Modern Lecture Hall",
|
||||
"fd6814ecc5e143cd899e615a80eaa2dc",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Brandon Business Sitting Front",
|
||||
"Brandon_Business_Sitting_Front_public",
|
||||
"studio_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Caroline Business Sitting Side",
|
||||
"Caroline_Business_Sitting_Side_public",
|
||||
"studio_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Ursula Lawyer Angle 4",
|
||||
"f7173d2bb8584c00bfec6905c5e9a492",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Sofia Corporate Presenter 01 Angle 3",
|
||||
"fe563971fd2d438e957372dac9e2be8c",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Seoyeon Health Nutrition Coach Angle 3",
|
||||
"fe3c5d5028d941398d064b8fc64a2dea",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Sanne Fitness Coach Angle 4",
|
||||
"d967f935a8bf4a0c8f0bccfd66c501d2",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
("Sander", "f5cd7b94056f495ca0610602d64a9aa3", "photo_avatar", ("avatar_v", "avatar_iv", "avatar_iii")),
|
||||
(
|
||||
"Rupert Personal Development Coach Angle 4",
|
||||
"f57b3e626adb4bc997b38f64884adce4",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Olivier Professor Angle 2",
|
||||
"f6659bbb094b459c87c967edbb9ee481",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Obi Health Nutrition Coach Angle 5",
|
||||
"f3dc2c38201d414382f506d2d8e8d029",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Matilda Modern Office Setting",
|
||||
"fda889ac354a440da8dbecc410981273",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Mateo Traditional Law Office",
|
||||
"ff172d6c499c4e47ba6fcc5de631e9fc",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Marlon Inviting Armchair Setting",
|
||||
"f5a57db099ab462daa3e7c604a05dacc",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Margaret Professor Angle 1",
|
||||
"fb472bc29ab04bcca576e3703978fecb",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Marek Therapy Coach Angle 3",
|
||||
"e197768703f1463a93dc25ada1f421fb",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Maeve Warm, Professional Setting",
|
||||
"faf66681d8cc48dc82c4283200b3e782",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Lorenzo Professor Angle 5",
|
||||
"fc268dc244bb40d7a554663ce723dcf0",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
("Luca", "Luca_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Bruce", "Bruce_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Nico", "Nico_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Lisa", "Lisa_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Sophie", "Sophie_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Aiko", "Aiko_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Rebecca (portrait)", "Rebecca_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Daphne in Grey blazer (portrait)", "Daphne_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Bryce in Black t-shirt", "Bryce_public_5", "studio_avatar", ("avatar_iii",)),
|
||||
("Diora in White shirt", "Diora_public_3", "studio_avatar", ("avatar_iii",)),
|
||||
("Freja in White blazer", "Freja_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Albert in Blue blazer", "Albert_public_2", "studio_avatar", ("avatar_iii",)),
|
||||
("Emery in Red blazer", "Emery_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Minho in Blue shirt", "Minho_public_6", "studio_avatar", ("avatar_iii",)),
|
||||
("Aditya in Brown blazer", "Aditya_public_4", "studio_avatar", ("avatar_iii",)),
|
||||
("Nadim in Blue blazer", "Nadim_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Iker in Black blazer", "Iker_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Nour in Black blazer", "Nour_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Saskia in Blue blazer", "Saskia_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Lucien in Blue blazer", "Lucien_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Esmond in Blue suit", "Esmond_public_3", "studio_avatar", ("avatar_iii",)),
|
||||
("Jinwoo in Blue suit", "Jinwoo_public_5", "studio_avatar", ("avatar_iii",)),
|
||||
("Annelore in Red sweater (portrait)", "Annelore_public_3", "studio_avatar", ("avatar_iii",)),
|
||||
("Bastien in Blue shirt", "Bastien_public_4", "studio_avatar", ("avatar_iii",)),
|
||||
("Zosia in Khaki blazer", "Zosia_public_3", "studio_avatar", ("avatar_iii",)),
|
||||
("Tahlia in Dark blue suit", "Tahlia_public_4", "studio_avatar", ("avatar_iii",)),
|
||||
]
|
||||
HEYGEN_AVATAR_OPTIONS = [x[0] for x in HEYGEN_AVATAR_LOOKS]
|
||||
HEYGEN_AVATAR_MAP = {x[0]: (x[1], x[2], x[3]) for x in HEYGEN_AVATAR_LOOKS}
|
||||
|
||||
# (label, voice_id) — Starfish-compatible voices for the TTS endpoint
|
||||
HEYGEN_VOICE_TTS: list[tuple[str, str]] = [
|
||||
("Chill Brian (English, male)", "d2f4f24783d04e22ab49ee8fdc3715e0"),
|
||||
("Zain (English, female)", "0047732240584155b1588455313e78ec"),
|
||||
("Narrator Mateo - Excited 🤩 (Spanish, male)", "0077225a877e457db4572ccaf245910b"),
|
||||
("Aria (English, female)", "007e1378fc454a9f976db570ba6164a7"),
|
||||
("Caryns (English, female)", "0082e70326864107823605db0d77f5e0"),
|
||||
("Klara (English, female)", "01209fdcd1c24a109c86dc24ee0f71c0"),
|
||||
("Bold Kasia - Excited 🤩 (Polish, female)", "015482a78b9a46ebae74bd0beb17765b"),
|
||||
("Shaun (English, male)", "01c42cddcfdc4665a57b8d89cba8ffc1"),
|
||||
("Senthil (English, male)", "01d674cfd32b4728a3fddd21b7e7d543"),
|
||||
("Cody (English, male)", "01f98ed43e6140349f47dbd37a416827"),
|
||||
("Saffron (English, female)", "0258bbc2cd8648cfa357adfb833f6d7b"),
|
||||
("Blanka - Lifelike (English, female)", "02880d1c6fd94b7799d91135581ed810"),
|
||||
("Rami (English, male)", "02d5366a90af4c7a87157808ff352e33"),
|
||||
("Rhodes (English, female)", "02dce0a169b3460084b6c914d18fb2c8"),
|
||||
("Michelle - Voice 1 (English, female)", "02X8sHnuxFpsq1caYWN0"),
|
||||
("Autumn - UGC 3 (English, female)", "03dca9ebfca441dba55fb14afa0791b7"),
|
||||
("Reassuring Rupert (English, male)", "03fcf8ecb0a94b6b94e9007edb7c35f8"),
|
||||
("Rose - UGC -2 (English, female)", "0495e14c2bd74eb3aeeef03583e0bce5"),
|
||||
("Derya - Lifelike - Broadcaster 🎙️ (English, female)", "04d0ae1d0af2489ca7d3bb402a39a890"),
|
||||
("Dynamic Derek (English, male)", "0516c2d857eb425c94e90b068241914e"),
|
||||
("Lotte (English, female)", "052fcfb83d1a4c2f8d0368c226fea4b9"),
|
||||
("Thanos - Broadcaster 🎙️ (English, male)", "054af44a167344d0af2722fdfef08d17"),
|
||||
("Marcia (English, female)", "05f19352e8f74b0392a8f411eba40de1"),
|
||||
("Camden (English, male)", "06468055edd4458aa131a1dfd813c1e9"),
|
||||
("Rumi (English, female)", "06672207805f41a9ad0af6797f8aa14b"),
|
||||
("Pippa (English, female)", "06b68c4dbb544935b9af984e80efa4fb"),
|
||||
("William Prescott - Broadcaster 🎙️ (English, male)", "06c816b952f14fa9b3a6c42aa151f731"),
|
||||
("Sammy (English, female)", "06e6facd99654b9dbb9308f67bf3a31c"),
|
||||
("Breezy Bagus (Indonesian, male)", "06e81a5d7c8b41818d3f0b38f7cf15a1"),
|
||||
("Ben (English, male)", "07ca39b243184dbcb82e7e0f0e524b21"),
|
||||
("Smooth Dev (English, male)", "07d2ba65847541feb97abc9b60181555"),
|
||||
("Daran inside booth (English, male)", "080d9383c0314056aef392892e009806"),
|
||||
("Peppy Stella (English, female)", "084760b4922a44599575c770070ec2d7"),
|
||||
("Silas (English, male)", "08f561403ec846dbbd8c691cc448f45a"),
|
||||
("Aditya (English, male)", "09c3d65e44e247dd8b78a97a903feb58"),
|
||||
("Christy (English, female)", "09d88c036bf449fa905900c08b235a37"),
|
||||
("Elio (English, male)", "0a0b38624ac64ec6afcd5842a977ca10"),
|
||||
("Luminous Laksh (Hindi, male)", "0adc547b76a5401c856274c379904eb7"),
|
||||
("Jeff (English, male)", "0add542e349f4ccaba6ecb3b7ced6034"),
|
||||
("Tahlia Brooks - Excited 🤩 (English, female)", "0b440d1ac2454d69a73302fc806522b1"),
|
||||
("Riya Mehta (Hindi, female)", "0b464b2f4e2249a4b5a05e60eaf41e7e"),
|
||||
("Ben Hart (English, male)", "0b47b5a637e944f9bfd49913999b344b"),
|
||||
("Skylar (English, female)", "0bbfbda5aa924a68a9d1da7b8496052a"),
|
||||
("Relaxed Reece (English, male)", "0c2151d538844c70a8b096de533f2828"),
|
||||
("Daniel (English, male)", "0c23804af39a4946ac6fda42bfff2738"),
|
||||
("Melani (English, female)", "0c54c6399ad64551a304e1a346677723"),
|
||||
("Clover (English, female)", "0ccb0bea067d4449ad367baeed7ea2e9"),
|
||||
("Pedro Lima - Serious 😐 (Portuguese, male)", "0d0e23e8170446e38b18a7380b2d30a8"),
|
||||
("Ana Carvalho (Portuguese, female)", "0d23c5b2f6004e909802a2e8bfcd52c2"),
|
||||
("Confident Connor - Excited 🤩 (English, male)", "0dd34c3eb79247238219eea35aeb58cd"),
|
||||
("Vibrant Victor (Spanish, male)", "1062976ea8bf42f4adc27c7e868b8fde"),
|
||||
("Young Olivier (French, male)", "1c5dc9a8f8cf4de0932f91d75f43a15d"),
|
||||
("Émile Noir (French, male)", "25a6a67280574d3da78e97b1935ebfc7"),
|
||||
("Steadfast Stefan (German, male)", "0eb85e6e8710473b82f7e88609ba3053"),
|
||||
("Deep Dieter (German, male)", "118949676b0a46629d1ad52981c3ef84"),
|
||||
("Serene Marco (Italian, male)", "72e922488a614041b5ab5f6ee07e3deb"),
|
||||
("Murmuring Matteo (Italian, male)", "755902b751654f30a6ef49e8bbcacfec"),
|
||||
("Gail in car (Multilingual, female)", "0214ac51f93e420f8711d568dcfbc50e"),
|
||||
("Daran outside walking (Multilingual, male)", "0ac81e725f4948dfa9638ceca216bcfa"),
|
||||
("BOB - Voice 1 (Chinese, unknown)", "dMkR1XwIkarpNqWUJLnX"),
|
||||
("Hakeem Hassan (Arabic, male)", "61a4359785664d01a59664ceb87ce6d4"),
|
||||
("Rami Idris (Arabic, male)", "a0bd2e5d41a74643be47ac75ca9171a2"),
|
||||
("Bold Kasia - Friendly 😊 (Polish, female)", "331624aec8b24a6c9287b8e16bdf54e8"),
|
||||
("Tranquil Tulin (Turkish, female)", "61646c861eb64e2d9036d8db51385356"),
|
||||
("Dynamic Derya (Turkish, female)", "664b73058b784aa89ddb2924c141d441"),
|
||||
("Quiet Dewa (Indonesian, male)", "1fa1193cf1d74f27ba58531c07ef9862"),
|
||||
("Cuong (Vietnamese, male)", "8af68d7ea38f4e7ca05cf46c3f7a590b"),
|
||||
]
|
||||
HEYGEN_VOICE_TTS_OPTIONS = [x[0] for x in HEYGEN_VOICE_TTS]
|
||||
HEYGEN_VOICE_TTS_MAP = dict(HEYGEN_VOICE_TTS)
|
||||
|
||||
# (label, voice_id) — top-ranked voices for video narration (any engine)
|
||||
HEYGEN_VOICE_GENERAL: list[tuple[str, str]] = [
|
||||
("Cassidy (English, female)", "16a09e4706f74997ba4ed05ea11470f6"),
|
||||
("Hope (English, female)", "42d00d4aac5441279d8536cd6b52c53c"),
|
||||
("Archer (English, male)", "453c20e1525a429080e2ad9e4b26f2cd"),
|
||||
("Brittney (English, female)", "4754e1ec667544b0bd18cdf4bec7d6a7"),
|
||||
("Mark (English, male)", "5d8c378ba8c3434586081a52ac368738"),
|
||||
("Andrew (English, male)", "6be73833ef9a4eb0aeee399b8fe9d62b"),
|
||||
("Spuds Oxley (English, male)", "76940a9adcd0490a9ce2cfe9a64a2664"),
|
||||
("Patrick (English, male)", "7e157ec62c9c45f1adca12faae72c86f"),
|
||||
("David Castlemore (English, male)", "828b59f834fd4c7188da322b6d9b6c75"),
|
||||
("Michael C (English, male)", "8661cd40d6c44c709e2d0031c0186ada"),
|
||||
("Adam Stone (English, male)", "88bb9ee1c81b466eb2a08fdde86d3619"),
|
||||
("Alex (English, male)", "897d6a9b2c844f56aa077238768fe10a"),
|
||||
("Monika Sogam (English, female)", "97dd67ab8ce242b6a9e7689cb00c6414"),
|
||||
("Jessica Anne Bogart (English, female)", "b966c31caf124c2a99f19ff1479c964f"),
|
||||
("John Doe (English, male)", "c4a8ceb7a2954500bc047fb092bcff3f"),
|
||||
("Ivy (English, female)", "cef3bc4e0a84424cafcde6f2cf466c97"),
|
||||
("Chill Brian (English, male)", "d2f4f24783d04e22ab49ee8fdc3715e0"),
|
||||
("Allison (English, female)", "f8c69e517f424cafaecde32dde57096b"),
|
||||
("Mia Starset (Norwegian, female)", "000466f8ac6d47a49f5743d50b3778de"),
|
||||
("William Shanks (Spanish, male)", "001248bb63f847888d37b766ee8b3a47"),
|
||||
("Zain (English, female)", "0047732240584155b1588455313e78ec"),
|
||||
("Jora Slobod (Romanian, male)", "00631519159a402ab5d8f719e51532bb"),
|
||||
("Narrator Mateo - Excited 🤩 (Spanish, male)", "0077225a877e457db4572ccaf245910b"),
|
||||
("Aria (English, female)", "007e1378fc454a9f976db570ba6164a7"),
|
||||
("Caryns (English, female)", "0082e70326864107823605db0d77f5e0"),
|
||||
("Klara (English, female)", "01209fdcd1c24a109c86dc24ee0f71c0"),
|
||||
("Son Tran (Vietnamese, male)", "0132f85950a94d11ba180f885101bf84"),
|
||||
("Bold Kasia - Excited 🤩 (Polish, female)", "015482a78b9a46ebae74bd0beb17765b"),
|
||||
("Marc Aurèle (French, male)", "018a94cf15574491a0bab7f6799ac15b"),
|
||||
("Shaun (English, male)", "01c42cddcfdc4665a57b8d89cba8ffc1"),
|
||||
("Senthil (English, male)", "01d674cfd32b4728a3fddd21b7e7d543"),
|
||||
("Cody (English, male)", "01f98ed43e6140349f47dbd37a416827"),
|
||||
("Saffron (English, female)", "0258bbc2cd8648cfa357adfb833f6d7b"),
|
||||
("Blanka - Lifelike (English, female)", "02880d1c6fd94b7799d91135581ed810"),
|
||||
("Rami (English, male)", "02d5366a90af4c7a87157808ff352e33"),
|
||||
("Rhodes (English, female)", "02dce0a169b3460084b6c914d18fb2c8"),
|
||||
("Michelle - Voice 1 (English, female)", "02X8sHnuxFpsq1caYWN0"),
|
||||
("Tuba (, female)", "034ca0c32b6542028748d6d365d90d6a"),
|
||||
("Autumn - UGC 3 (English, female)", "03dca9ebfca441dba55fb14afa0791b7"),
|
||||
("Reassuring Rupert (English, male)", "03fcf8ecb0a94b6b94e9007edb7c35f8"),
|
||||
]
|
||||
HEYGEN_VOICE_GENERAL_OPTIONS = [x[0] for x in HEYGEN_VOICE_GENERAL]
|
||||
HEYGEN_VOICE_GENERAL_MAP = dict(HEYGEN_VOICE_GENERAL)
|
||||
|
||||
HEYGEN_TRANSLATE_LANGUAGES = [
|
||||
"English",
|
||||
"Spanish",
|
||||
"Spanish (Spain)",
|
||||
"Spanish (Mexico)",
|
||||
"French",
|
||||
"French (France)",
|
||||
"German",
|
||||
"German (Germany)",
|
||||
"Portuguese",
|
||||
"Portuguese (Brazil)",
|
||||
"Italian",
|
||||
"Italian (Italy)",
|
||||
"Japanese",
|
||||
"Japanese (Japan)",
|
||||
"Korean",
|
||||
"Chinese (Mandarin, Simplified)",
|
||||
"Arabic",
|
||||
"Hindi",
|
||||
"Hindi (India)",
|
||||
"Russian",
|
||||
"Russian (Russia)",
|
||||
"Dutch",
|
||||
"Polish",
|
||||
"Turkish",
|
||||
"Indonesian",
|
||||
"Vietnamese",
|
||||
"Ukrainian",
|
||||
"Afrikaans (South Africa)",
|
||||
"Albanian (Albania)",
|
||||
"Amharic (Ethiopia)",
|
||||
"Arabic (Algeria)",
|
||||
"Arabic (Bahrain)",
|
||||
"Arabic (Egypt)",
|
||||
"Arabic (Iraq)",
|
||||
"Arabic (Jordan)",
|
||||
"Arabic (Kuwait)",
|
||||
"Arabic (Lebanon)",
|
||||
"Arabic (Libya)",
|
||||
"Arabic (Morocco)",
|
||||
"Arabic (Oman)",
|
||||
"Arabic (Qatar)",
|
||||
"Arabic (Saudi Arabia)",
|
||||
"Arabic (Syria)",
|
||||
"Arabic (Tunisia)",
|
||||
"Arabic (United Arab Emirates)",
|
||||
"Arabic (World)",
|
||||
"Arabic (Yemen)",
|
||||
"Armenian (Armenia)",
|
||||
"Azerbaijani (Latin, Azerbaijan)",
|
||||
"Bangla (Bangladesh)",
|
||||
"Basque",
|
||||
"Belarusian (Belarus)",
|
||||
"Bengali (India)",
|
||||
"Bosnian (Bosnia and Herzegovina)",
|
||||
"Bulgarian",
|
||||
"Bulgarian (Bulgaria)",
|
||||
"Burmese (Myanmar)",
|
||||
"Catalan",
|
||||
"Chinese (Cantonese, Traditional)",
|
||||
"Chinese (Jilu Mandarin, Simplified)",
|
||||
"Chinese (Northeastern Mandarin, Simplified)",
|
||||
"Chinese (Southwestern Mandarin, Simplified)",
|
||||
"Chinese (Taiwanese Mandarin, Traditional)",
|
||||
"Chinese (Wu, Simplified)",
|
||||
"Chinese (Zhongyuan Mandarin Henan, Simplified)",
|
||||
"Chinese (Zhongyuan Mandarin Shaanxi, Simplified)",
|
||||
"Croatian",
|
||||
"Croatian (Croatia)",
|
||||
"Czech",
|
||||
"Czech (Czechia)",
|
||||
"Danish",
|
||||
"Danish (Denmark)",
|
||||
"Dutch (Belgium)",
|
||||
"Dutch (Netherlands)",
|
||||
"English (Australia)",
|
||||
"English (Canada)",
|
||||
"English (Hong Kong SAR)",
|
||||
"English (India)",
|
||||
"English (Ireland)",
|
||||
"English (Kenya)",
|
||||
"English (New Zealand)",
|
||||
"English (Nigeria)",
|
||||
"English (Philippines)",
|
||||
"English (Singapore)",
|
||||
"English (South Africa)",
|
||||
"English (Tanzania)",
|
||||
"English (UK)",
|
||||
"English (United States)",
|
||||
"Estonian (Estonia)",
|
||||
"Filipino",
|
||||
"Filipino (Cebuano)",
|
||||
"Filipino (Philippines)",
|
||||
"Finnish",
|
||||
"Finnish (Finland)",
|
||||
"French (Belgium)",
|
||||
"French (Canada)",
|
||||
"French (Switzerland)",
|
||||
"Galician",
|
||||
"Georgian (Georgia)",
|
||||
"German (Austria)",
|
||||
"German (Switzerland)",
|
||||
"Greek",
|
||||
"Greek (Greece)",
|
||||
"Gujarati (India)",
|
||||
"Haitian Creole (Haiti)",
|
||||
"Hebrew (Israel)",
|
||||
"Hungarian (Hungary)",
|
||||
"Icelandic (Iceland)",
|
||||
"Indonesian (Indonesia)",
|
||||
"Irish (Ireland)",
|
||||
"Javanese (Latin, Indonesia)",
|
||||
"Kannada (India)",
|
||||
"Kazakh (Kazakhstan)",
|
||||
"Khmer (Cambodia)",
|
||||
"Konkani (India)",
|
||||
"Korean (Korea)",
|
||||
"Lao (Laos)",
|
||||
"Latin (Vatican City)",
|
||||
"Latvian (Latvia)",
|
||||
"Lithuanian (Lithuania)",
|
||||
"Luxembourgish (Luxembourg)",
|
||||
"Macedonian (North Macedonia)",
|
||||
"Maithili (India)",
|
||||
"Malagasy (Madagascar)",
|
||||
"Malay",
|
||||
"Malay (Malaysia)",
|
||||
"Malayalam (India)",
|
||||
"Maltese (Malta)",
|
||||
"Mandarin",
|
||||
"Marathi (India)",
|
||||
"Mongolian (Mongolia)",
|
||||
"Nepali (Nepal)",
|
||||
"Norwegian Bokmål (Norway)",
|
||||
"Norwegian Nynorsk (Norway)",
|
||||
"Odia (India)",
|
||||
"Pashto (Afghanistan)",
|
||||
"Persian (Iran)",
|
||||
"Polish (Poland)",
|
||||
"Portuguese (Portugal)",
|
||||
"Punjabi (India)",
|
||||
"Romanian",
|
||||
"Romanian (Romania)",
|
||||
"Serbian (Latin, Serbia)",
|
||||
"Sindhi (India)",
|
||||
"Sinhala (Sri Lanka)",
|
||||
"Slovak",
|
||||
"Slovak (Slovakia)",
|
||||
"Slovenian (Slovenia)",
|
||||
"Somali (Somalia)",
|
||||
"Spanish (Argentina)",
|
||||
"Spanish (Bolivia)",
|
||||
"Spanish (Chile)",
|
||||
"Spanish (Colombia)",
|
||||
"Spanish (Costa Rica)",
|
||||
"Spanish (Cuba)",
|
||||
"Spanish (Dominican Republic)",
|
||||
"Spanish (Ecuador)",
|
||||
"Spanish (El Salvador)",
|
||||
"Spanish (Equatorial Guinea)",
|
||||
"Spanish (Guatemala)",
|
||||
"Spanish (Honduras)",
|
||||
"Spanish (Latin America)",
|
||||
"Spanish (Nicaragua)",
|
||||
"Spanish (Panama)",
|
||||
"Spanish (Paraguay)",
|
||||
"Spanish (Peru)",
|
||||
"Spanish (Puerto Rico)",
|
||||
"Spanish (United States)",
|
||||
"Spanish (Uruguay)",
|
||||
"Spanish (Venezuela)",
|
||||
"Sundanese (Indonesia)",
|
||||
"Swahili (Kenya)",
|
||||
"Swahili (Tanzania)",
|
||||
"Swedish",
|
||||
"Swedish (Sweden)",
|
||||
"Tamil",
|
||||
"Tamil (India)",
|
||||
"Tamil (Malaysia)",
|
||||
"Tamil (Singapore)",
|
||||
"Tamil (Sri Lanka)",
|
||||
"Telugu (India)",
|
||||
"Thai (Thailand)",
|
||||
"Turkish (Türkiye)",
|
||||
"Ukrainian (Ukraine)",
|
||||
"Urdu (India)",
|
||||
"Urdu (Pakistan)",
|
||||
"Uzbek (Latin, Uzbekistan)",
|
||||
"Vietnamese (Vietnam)",
|
||||
"Welsh (United Kingdom)",
|
||||
"Zulu (South Africa)",
|
||||
]
|
||||
|
|
@ -28,6 +28,9 @@ ANTHROPIC_IMAGE_MAX_PIXELS = 1568 * 1568
|
|||
CLAUDE_MAX_IMAGES = 20
|
||||
|
||||
CLAUDE_MODELS: dict[str, str] = {
|
||||
"Opus 4.8": "claude-opus-4-8",
|
||||
"Fable 5": "claude-fable-5",
|
||||
"Sonnet 5": "claude-sonnet-5",
|
||||
"Opus 4.7": "claude-opus-4-7",
|
||||
"Opus 4.6": "claude-opus-4-6",
|
||||
"Sonnet 4.6": "claude-sonnet-4-6",
|
||||
|
|
@ -36,9 +39,12 @@ CLAUDE_MODELS: dict[str, str] = {
|
|||
}
|
||||
|
||||
_THINKING_UNSUPPORTED = {"Haiku 4.5"}
|
||||
# Models that use the newer "adaptive" thinking mode (Opus 4.7 requires it; older models keep the explicit budget API).
|
||||
# Models that use the newer "adaptive" thinking mode (Opus 4.7+ require it; older models keep the explicit budget API).
|
||||
# Anthropic decides the actual budget when adaptive is used, based on the `output_config.effort` hint.
|
||||
_ADAPTIVE_THINKING_MODELS = {"Opus 4.7", "Opus 4.6", "Sonnet 4.6"}
|
||||
_ADAPTIVE_THINKING_MODELS = {"Opus 4.8", "Sonnet 5", "Opus 4.7", "Opus 4.6", "Sonnet 4.6"}
|
||||
_ALWAYS_THINKING_MODELS = {"Fable 5"}
|
||||
_EXPLICIT_THINKING_OFF_MODELS = {"Sonnet 5"}
|
||||
_NO_TEMPERATURE_MODELS = {"Opus 4.8", "Fable 5", "Sonnet 5"}
|
||||
|
||||
# Budget mode (Sonnet 4.5): effort -> reasoning budget in tokens. Must be < max_tokens.
|
||||
# Sized so even the "high" budget fits comfortably under the default max_tokens=32768.
|
||||
|
|
@ -60,20 +66,33 @@ def _claude_model_inputs(model_label: str):
|
|||
tooltip="Maximum number of tokens to generate (includes reasoning tokens when enabled).",
|
||||
advanced=True,
|
||||
),
|
||||
IO.Float.Input(
|
||||
"temperature",
|
||||
default=1.0,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Controls randomness. 0.0 is deterministic, 1.0 is most random. "
|
||||
"Ignored for Opus 4.7 and any model when reasoning_effort is set."
|
||||
),
|
||||
advanced=True,
|
||||
),
|
||||
]
|
||||
if model_label not in _THINKING_UNSUPPORTED:
|
||||
if model_label not in _NO_TEMPERATURE_MODELS:
|
||||
inputs.append(
|
||||
IO.Float.Input(
|
||||
"temperature",
|
||||
default=1.0,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Controls randomness. 0.0 is deterministic, 1.0 is most random. "
|
||||
"Ignored for Opus 4.7 and any model when reasoning_effort is set."
|
||||
),
|
||||
advanced=True,
|
||||
)
|
||||
)
|
||||
if model_label in _ALWAYS_THINKING_MODELS:
|
||||
inputs.append(
|
||||
IO.Combo.Input(
|
||||
"reasoning_effort",
|
||||
options=[e for e in _REASONING_EFFORTS if e != "off"],
|
||||
default="high",
|
||||
tooltip="Extended thinking effort. Reasoning is always enabled for this model.",
|
||||
advanced=True,
|
||||
)
|
||||
)
|
||||
elif model_label not in _THINKING_UNSUPPORTED:
|
||||
inputs.append(
|
||||
IO.Combo.Input(
|
||||
"reasoning_effort",
|
||||
|
|
@ -88,6 +107,12 @@ def _claude_model_inputs(model_label: str):
|
|||
|
||||
def _model_price_per_million(model: str) -> tuple[float, float] | None:
|
||||
"""Return (input_per_1M, output_per_1M) USD for a Claude model, or None if unknown."""
|
||||
if "fable-5" in model:
|
||||
return 14.30, 71.50
|
||||
if "opus-4-8" in model:
|
||||
return 7.15, 35.75
|
||||
if "sonnet-5" in model:
|
||||
return 2.86, 14.30
|
||||
if "opus-4-7" in model or "opus-4-6" in model or "opus-4-5" in model:
|
||||
return 5.0, 25.0
|
||||
if "sonnet-4" in model:
|
||||
|
|
@ -213,7 +238,22 @@ class ClaudeNode(IO.ComfyNode):
|
|||
expr="""
|
||||
(
|
||||
$m := widgets.model;
|
||||
$contains($m, "opus") ? {
|
||||
$contains($m, "fable") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.0143, 0.0715],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "opus 4.8") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.00715, 0.03575],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "sonnet 5") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.00286, 0.0143],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "opus") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.005, 0.025],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
|
|
@ -247,18 +287,23 @@ class ClaudeNode(IO.ComfyNode):
|
|||
model_label = model["model"]
|
||||
max_tokens = model.get("max_tokens", 32768)
|
||||
reasoning_effort = model.get("reasoning_effort", "off")
|
||||
thinking_enabled = reasoning_effort not in ("off", None) and model_label not in _THINKING_UNSUPPORTED
|
||||
always_thinking = model_label in _ALWAYS_THINKING_MODELS
|
||||
thinking_enabled = always_thinking or (
|
||||
reasoning_effort not in ("off", None) and model_label not in _THINKING_UNSUPPORTED
|
||||
)
|
||||
|
||||
# Anthropic requires temperature to be unset (defaults to 1.0) when thinking is enabled.
|
||||
# Opus 4.7 also rejects user-supplied temperature.
|
||||
if thinking_enabled or model_label == "Opus 4.7":
|
||||
if model_label in _NO_TEMPERATURE_MODELS or thinking_enabled or model_label == "Opus 4.7":
|
||||
temperature = None
|
||||
else:
|
||||
temperature = model.get("temperature", 1.0)
|
||||
|
||||
thinking_cfg: AnthropicThinkingConfig | None = None
|
||||
output_cfg: AnthropicOutputConfig | None = None
|
||||
if thinking_enabled:
|
||||
if always_thinking:
|
||||
output_cfg = AnthropicOutputConfig(effort=reasoning_effort)
|
||||
elif thinking_enabled:
|
||||
if model_label in _ADAPTIVE_THINKING_MODELS:
|
||||
# Adaptive mode - Anthropic chooses the budget based on effort hint
|
||||
thinking_cfg = AnthropicThinkingConfig(type="adaptive")
|
||||
|
|
@ -268,6 +313,8 @@ class ClaudeNode(IO.ComfyNode):
|
|||
budget = _REASONING_BUDGET[reasoning_effort]
|
||||
budget = min(budget, max(1024, max_tokens - 1024))
|
||||
thinking_cfg = AnthropicThinkingConfig(type="enabled", budget_tokens=budget)
|
||||
elif model_label in _EXPLICIT_THINKING_OFF_MODELS:
|
||||
thinking_cfg = AnthropicThinkingConfig(type="disabled")
|
||||
|
||||
image_tensors: list[Input.Image] = [t for t in (images or {}).values() if t is not None]
|
||||
if sum(get_number_of_images(t) for t in image_tensors) > CLAUDE_MAX_IMAGES:
|
||||
|
|
@ -293,6 +340,11 @@ class ClaudeNode(IO.ComfyNode):
|
|||
),
|
||||
price_extractor=calculate_tokens_price,
|
||||
)
|
||||
if response.stop_reason == "refusal":
|
||||
raise ValueError(
|
||||
"Claude declined to answer this request for safety reasons. "
|
||||
"Rephrase the prompt or try a different model."
|
||||
)
|
||||
return IO.NodeOutput(_get_text_from_response(response) or "Empty response from Claude model.")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2690,7 +2690,8 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
|
|||
"with ByteDance Seed Audio 1.0. Describe the voice(s), emotion, ambience, background music "
|
||||
"and sound effects in the prompt, and include the lines to speak. Optionally pick a built-in "
|
||||
"preset voice, clone voices from up to 3 reference clips (tagged @Audio1-3 in the prompt), "
|
||||
"or derive a voice from a character image. Up to 2 minutes of audio per run."
|
||||
"or derive a voice from a character image. Up to 2 minutes of audio per run. "
|
||||
"The multilingual model supports 20 languages and timestamp-based timing control."
|
||||
),
|
||||
inputs=[
|
||||
IO.String.Input(
|
||||
|
|
@ -2701,7 +2702,9 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
|
|||
"Describe the voice(s), emotion, pacing, ambience, background music and sound "
|
||||
"effects, and include the lines to speak (name characters inline for dialogue). "
|
||||
"In 'audio reference' mode, refer to connected clips by order as @Audio1, @Audio2, "
|
||||
"@Audio3. Maximum 3000 characters."
|
||||
"@Audio3. With the multilingual model, a quoted line can start with a timestamp "
|
||||
'range that controls when and how long it is spoken, e.g. "[5.5s:8.0s] Wait for me!". '
|
||||
"Write the prompt in the same language as the lines to speak. Maximum 3000 characters."
|
||||
),
|
||||
),
|
||||
IO.DynamicCombo.Input(
|
||||
|
|
@ -2796,6 +2799,19 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
|
|||
tooltip="Seed controls whether the node should re-run; "
|
||||
"results are non-deterministic regardless of seed.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"model",
|
||||
options=["seed-audio-1.0-multilingual", "seed-audio-1.0"],
|
||||
default="seed-audio-1.0-multilingual",
|
||||
optional=True,
|
||||
tooltip=(
|
||||
"seed-audio-1.0-multilingual: 20 languages (English, Chinese, Japanese, Korean, "
|
||||
"Mexican & Castilian Spanish, Indonesian, German, Brazilian Portuguese, French, "
|
||||
"Thai, Vietnamese, Malay, Filipino, Italian, Russian, Dutch, Polish, Turkish, "
|
||||
'Swedish) plus per-sentence timing control via "[5.5s:8.0s] ..." timestamps. '
|
||||
"seed-audio-1.0: English and Chinese only, no timing control."
|
||||
),
|
||||
),
|
||||
],
|
||||
outputs=[IO.Audio.Output()],
|
||||
hidden=[
|
||||
|
|
@ -2819,6 +2835,7 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
|
|||
loudness_rate: int,
|
||||
pitch_rate: int,
|
||||
seed: int,
|
||||
model: str = "seed-audio-1.0-multilingual",
|
||||
) -> IO.NodeOutput:
|
||||
mode = reference_mode["reference_mode"]
|
||||
audio_indices = connected_audio_indices(reference_mode)
|
||||
|
|
@ -2845,6 +2862,7 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
|
|||
ApiEndpoint(path="/proxy/byteplus/api/v3/tts/create", method="POST"),
|
||||
response_model=SeedAudioResponse,
|
||||
data=SeedAudioRequest(
|
||||
model=model,
|
||||
text_prompt=text_prompt,
|
||||
references=references,
|
||||
audio_config=SeedAudioConfig(
|
||||
|
|
|
|||
|
|
@ -24,6 +24,11 @@ from comfy_api_nodes.apis.gemini import (
|
|||
GeminiImageGenerateContentRequest,
|
||||
GeminiImageGenerationConfig,
|
||||
GeminiInlineData,
|
||||
GeminiInteraction,
|
||||
GeminiInteractionGenerationConfig,
|
||||
GeminiInteractionMediaPart,
|
||||
GeminiInteractionRequest,
|
||||
GeminiInteractionTextPart,
|
||||
GeminiMimeType,
|
||||
GeminiPart,
|
||||
GeminiRole,
|
||||
|
|
@ -51,9 +56,11 @@ from comfy_api_nodes.util import (
|
|||
)
|
||||
|
||||
GEMINI_BASE_ENDPOINT = "/proxy/vertexai/gemini"
|
||||
GEMINI_INTERACTIONS_ENDPOINT = "/proxy/gemini-interactions"
|
||||
GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB
|
||||
GEMINI_URL_INPUT_BUDGET = 10
|
||||
GEMINI_MAX_INLINE_BYTES = 18 * 1024 * 1024
|
||||
GEMINI_INTERACTIONS_MAX_INLINE_BYTES = 90 * 1024 * 1024 # the Interactions API rejects requests over ~100MiB
|
||||
GEMINI_IMAGE_SYS_PROMPT = (
|
||||
"You are an expert image-generation engine. You must ALWAYS produce an image.\n"
|
||||
"Interpret all user input—regardless of "
|
||||
|
|
@ -231,29 +238,10 @@ async def get_image_from_response(response: GeminiGenerateContentResponse, thoug
|
|||
return torch.cat(image_tensors, dim=0)
|
||||
|
||||
|
||||
async def get_video_from_response(
|
||||
response: GeminiGenerateContentResponse, cls: type[IO.ComfyNode] | None = None
|
||||
) -> InputImpl.VideoFromFile:
|
||||
parts = get_parts_by_type(response, "video/*")
|
||||
for part in parts:
|
||||
if part.inlineData and part.inlineData.data:
|
||||
return InputImpl.VideoFromFile(BytesIO(base64.b64decode(part.inlineData.data)))
|
||||
if part.fileData and part.fileData.fileUri:
|
||||
return await download_url_to_video_output(part.fileData.fileUri, cls=cls)
|
||||
model_message = get_text_from_response(response).strip()
|
||||
if model_message:
|
||||
raise ValueError(f"Gemini did not generate a video. Model response: {model_message}")
|
||||
raise ValueError(
|
||||
"Gemini did not generate a video. Try rephrasing your prompt, "
|
||||
"shortening the requested duration, or reducing the number of input images/videos."
|
||||
)
|
||||
|
||||
|
||||
def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | None:
|
||||
if not response.modelVersion:
|
||||
return None
|
||||
# Define prices (Cost per 1,000,000 tokens), see https://cloud.google.com/vertex-ai/generative-ai/pricing
|
||||
output_video_tokens_price = 0.0
|
||||
if response.modelVersion == "gemini-2.5-pro":
|
||||
input_tokens_price = 1.25
|
||||
output_text_tokens_price = 10.0
|
||||
|
|
@ -274,6 +262,10 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
|
|||
input_tokens_price = 0.25
|
||||
output_text_tokens_price = 1.50
|
||||
output_image_tokens_price = 0.0
|
||||
elif response.modelVersion == "gemini-3.5-flash":
|
||||
input_tokens_price = 1.50
|
||||
output_text_tokens_price = 9.0
|
||||
output_image_tokens_price = 0.0
|
||||
elif response.modelVersion in ("gemini-3-pro-image-preview", "gemini-3-pro-image"):
|
||||
input_tokens_price = 2
|
||||
output_text_tokens_price = 12.0
|
||||
|
|
@ -286,11 +278,6 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
|
|||
input_tokens_price = 0.25
|
||||
output_text_tokens_price = 1.50
|
||||
output_image_tokens_price = 30.0
|
||||
elif response.modelVersion == "gemini-omni-flash-preview":
|
||||
input_tokens_price = 2.145
|
||||
output_text_tokens_price = 12.87
|
||||
output_image_tokens_price = 0.0
|
||||
output_video_tokens_price = 25.025
|
||||
else:
|
||||
return None
|
||||
final_price = response.usageMetadata.promptTokenCount * input_tokens_price
|
||||
|
|
@ -298,8 +285,6 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
|
|||
for i in response.usageMetadata.candidatesTokensDetails:
|
||||
if i.modality == Modality.IMAGE:
|
||||
final_price += output_image_tokens_price * i.tokenCount # for Nano Banana models
|
||||
elif i.modality == Modality.VIDEO:
|
||||
final_price += output_video_tokens_price * i.tokenCount # for Omni Flash
|
||||
else:
|
||||
final_price += output_text_tokens_price * i.tokenCount
|
||||
if response.usageMetadata.thoughtsTokenCount:
|
||||
|
|
@ -307,6 +292,58 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
|
|||
return final_price / 1_000_000.0
|
||||
|
||||
|
||||
def get_text_from_interaction(interaction: GeminiInteraction) -> str:
|
||||
"""Extract and concatenate all model output text from an Interactions API response."""
|
||||
texts = []
|
||||
for step in interaction.steps or []:
|
||||
if step.type != "model_output":
|
||||
continue
|
||||
for content in step.content or []:
|
||||
if content.type == "text" and content.text:
|
||||
texts.append(content.text)
|
||||
return "\n".join(texts)
|
||||
|
||||
|
||||
async def get_video_from_interaction(
|
||||
interaction: GeminiInteraction, cls: type[IO.ComfyNode] | None = None
|
||||
) -> InputImpl.VideoFromFile:
|
||||
for step in interaction.steps or []:
|
||||
if step.type != "model_output":
|
||||
continue
|
||||
for content in step.content or []:
|
||||
if content.type != "video":
|
||||
continue
|
||||
if content.data:
|
||||
return InputImpl.VideoFromFile(BytesIO(base64.b64decode(content.data)))
|
||||
if content.uri:
|
||||
return await download_url_to_video_output(content.uri, cls=cls)
|
||||
model_message = get_text_from_interaction(interaction).strip()
|
||||
if model_message:
|
||||
raise ValueError(f"Gemini did not generate a video. Model response: {model_message}")
|
||||
raise ValueError(
|
||||
"Gemini did not generate a video. Try rephrasing your prompt, "
|
||||
"shortening the requested duration, or reducing the number of input images/videos."
|
||||
)
|
||||
|
||||
|
||||
def calculate_interaction_tokens_price(interaction: GeminiInteraction) -> float | None:
|
||||
if interaction.usage is None:
|
||||
return None
|
||||
input_tokens_price = 1.5
|
||||
output_tokens_prices = {"text": 9.0, "video": 17.5}
|
||||
thoughts_tokens_price = 9.0
|
||||
final_price = 0.0
|
||||
for i in interaction.usage.input_tokens_by_modality or []:
|
||||
if i.tokens:
|
||||
final_price += input_tokens_price * i.tokens
|
||||
for i in interaction.usage.output_tokens_by_modality or []:
|
||||
if i.tokens and i.modality in output_tokens_prices:
|
||||
final_price += output_tokens_prices[i.modality] * i.tokens
|
||||
if interaction.usage.total_thought_tokens:
|
||||
final_price += thoughts_tokens_price * interaction.usage.total_thought_tokens
|
||||
return final_price / 1_000_000.0
|
||||
|
||||
|
||||
def create_video_parts(video_input: Input.Video) -> list[GeminiPart]:
|
||||
"""Convert a single video input to Gemini API compatible parts (inline MP4/H.264)."""
|
||||
base_64_string = video_to_base64_string(
|
||||
|
|
@ -433,14 +470,24 @@ async def build_gemini_media_parts(
|
|||
part, nbytes = _media_inline_part(kind, payload)
|
||||
inline_bytes += nbytes
|
||||
if inline_bytes > max_inline_bytes:
|
||||
detail = f" after the first {url_budget} inputs are uploaded as URLs" if url_budget else ""
|
||||
raise ValueError(
|
||||
f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB after the first "
|
||||
f"{url_budget} inputs are uploaded as URLs). Reduce the number or size of attached media."
|
||||
f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB{detail}). "
|
||||
"Reduce the number or size of attached media."
|
||||
)
|
||||
parts.append(part)
|
||||
return parts
|
||||
|
||||
|
||||
def to_interaction_media_part(part: GeminiPart) -> GeminiInteractionMediaPart:
|
||||
"""Convert a fileData/inlineData GeminiPart into an Interactions API media part."""
|
||||
if part.fileData:
|
||||
mime = part.fileData.mimeType.value
|
||||
return GeminiInteractionMediaPart(type=mime.split("/")[0], uri=part.fileData.fileUri, mime_type=mime)
|
||||
mime = part.inlineData.mimeType.value
|
||||
return GeminiInteractionMediaPart(type=mime.split("/")[0], data=part.inlineData.data, mime_type=mime)
|
||||
|
||||
|
||||
class GeminiNode(IO.ComfyNode):
|
||||
"""
|
||||
Node to generate text responses from a Gemini model.
|
||||
|
|
@ -619,11 +666,12 @@ class GeminiNode(IO.ComfyNode):
|
|||
|
||||
GEMINI_V2_MODELS: dict[str, str] = {
|
||||
"Gemini 3.1 Pro": "gemini-3.1-pro-preview",
|
||||
"Gemini 3.5 Flash": "gemini-3.5-flash",
|
||||
"Gemini 3.1 Flash-Lite": "gemini-3.1-flash-lite-preview",
|
||||
}
|
||||
|
||||
|
||||
def _gemini_text_model_inputs(thinking_default: str) -> list[Input]:
|
||||
def _gemini_text_model_inputs(thinking_default: str, thinking_options: list[str] | None = None) -> list[Input]:
|
||||
"""Per-model inputs revealed by the model DynamicCombo (shared media + sampling controls)."""
|
||||
return [
|
||||
IO.Autogrow.Input(
|
||||
|
|
@ -661,7 +709,7 @@ def _gemini_text_model_inputs(thinking_default: str) -> list[Input]:
|
|||
),
|
||||
IO.Combo.Input(
|
||||
"thinking_level",
|
||||
options=["LOW", "HIGH"],
|
||||
options=thinking_options or ["LOW", "HIGH"],
|
||||
default=thinking_default,
|
||||
tooltip="How hard the model reasons internally before answering. "
|
||||
"HIGH improves quality on difficult tasks but costs more (thinking) tokens and is slower.",
|
||||
|
|
@ -719,6 +767,10 @@ class GeminiNodeV2(IO.ComfyNode):
|
|||
IO.DynamicCombo.Input(
|
||||
"model",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"Gemini 3.5 Flash",
|
||||
_gemini_text_model_inputs("MEDIUM", ["MINIMAL", "LOW", "MEDIUM", "HIGH"]),
|
||||
),
|
||||
IO.DynamicCombo.Option("Gemini 3.1 Pro", _gemini_text_model_inputs("HIGH")),
|
||||
IO.DynamicCombo.Option("Gemini 3.1 Flash-Lite", _gemini_text_model_inputs("LOW")),
|
||||
],
|
||||
|
|
@ -759,7 +811,13 @@ class GeminiNodeV2(IO.ComfyNode):
|
|||
"type": "list_usd",
|
||||
"usd": [0.00025, 0.0015],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
} : {
|
||||
}
|
||||
: $contains($m, "3.5 flash") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.0015, 0.009],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: {
|
||||
"type": "list_usd",
|
||||
"usd": [0.002, 0.012],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
|
|
@ -1661,7 +1719,7 @@ class GeminiVideoOmni(IO.ComfyNode):
|
|||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr='{"type":"usd","usd":0.146,"format":{"suffix":"/second","approximate":true}}'
|
||||
expr='{"type":"usd","usd":0.101,"format":{"suffix":"/second","approximate":true}}'
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -1680,27 +1738,41 @@ class GeminiVideoOmni(IO.ComfyNode):
|
|||
for video in videos:
|
||||
validate_video_duration(video, max_duration=10)
|
||||
|
||||
parts: list[GeminiPart] = []
|
||||
parts: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = []
|
||||
if images or videos:
|
||||
parts.extend(await build_gemini_media_parts(cls, images, [], videos))
|
||||
parts.append(GeminiPart(text=prompt))
|
||||
response = await sync_op(
|
||||
# The Interactions API accepts video only inline or as a Files API URI, not as an HTTP URL.
|
||||
media_parts = await build_gemini_media_parts(
|
||||
cls, [], [], videos, url_budget=0, max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES
|
||||
)
|
||||
video_inline_bytes = sum(len(p.inlineData.data) for p in media_parts)
|
||||
media_parts += await build_gemini_media_parts(
|
||||
cls, images, [], [], max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES - video_inline_bytes
|
||||
)
|
||||
parts.extend(to_interaction_media_part(p) for p in media_parts)
|
||||
parts.append(GeminiInteractionTextPart(text=prompt))
|
||||
interaction = await sync_op(
|
||||
cls,
|
||||
ApiEndpoint(path=f"{GEMINI_BASE_ENDPOINT}/{model_id}", method="POST"),
|
||||
data=GeminiGenerateContentRequest(
|
||||
contents=[GeminiContent(role=GeminiRole.user, parts=parts)],
|
||||
generationConfig=GeminiGenerationConfig(
|
||||
responseModalities=["TEXT", "VIDEO"],
|
||||
ApiEndpoint(path=GEMINI_INTERACTIONS_ENDPOINT, method="POST"),
|
||||
data=GeminiInteractionRequest(
|
||||
model=model_id,
|
||||
input=parts,
|
||||
generation_config=GeminiInteractionGenerationConfig(
|
||||
temperature=model.get("temperature", 1.0),
|
||||
topP=model.get("top_p", 0.95),
|
||||
top_p=model.get("top_p", 0.95),
|
||||
),
|
||||
),
|
||||
response_model=GeminiGenerateContentResponse,
|
||||
price_extractor=calculate_tokens_price,
|
||||
response_model=GeminiInteraction,
|
||||
price_extractor=calculate_interaction_tokens_price,
|
||||
)
|
||||
if interaction.status != "completed":
|
||||
model_message = get_text_from_interaction(interaction).strip()
|
||||
raise ValueError(
|
||||
f"Gemini interaction did not complete (status: {interaction.status})."
|
||||
+ (f" Model response: {model_message}" if model_message else "")
|
||||
)
|
||||
return IO.NodeOutput(
|
||||
await get_video_from_response(response, cls=cls),
|
||||
get_text_from_response(response),
|
||||
await get_video_from_interaction(interaction, cls=cls),
|
||||
get_text_from_interaction(interaction),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,799 @@
|
|||
import uuid
|
||||
|
||||
import torch
|
||||
from typing_extensions import override
|
||||
|
||||
from comfy_api.latest import IO, ComfyExtension, Input
|
||||
from comfy_api_nodes.apis.heygen import (
|
||||
HEYGEN_AVATAR_MAP,
|
||||
HEYGEN_AVATAR_OPTIONS,
|
||||
HEYGEN_TRANSLATE_LANGUAGES,
|
||||
HEYGEN_VOICE_GENERAL_MAP,
|
||||
HEYGEN_VOICE_GENERAL_OPTIONS,
|
||||
HEYGEN_VOICE_TTS_MAP,
|
||||
HEYGEN_VOICE_TTS_OPTIONS,
|
||||
)
|
||||
from comfy_api_nodes.util import (
|
||||
ApiEndpoint,
|
||||
audio_bytes_to_audio_input,
|
||||
download_url_as_bytesio,
|
||||
download_url_to_image_tensor,
|
||||
download_url_to_video_output,
|
||||
downscale_image_tensor_by_max_side,
|
||||
get_number_of_images,
|
||||
poll_op_raw,
|
||||
sync_op_raw,
|
||||
upload_audio_to_comfyapi,
|
||||
upload_image_to_comfyapi,
|
||||
upload_images_to_comfyapi,
|
||||
upload_video_to_comfyapi,
|
||||
validate_string,
|
||||
)
|
||||
from server import PromptServer
|
||||
|
||||
_VIDEOS_PATH = "/proxy/heygen/v3/videos"
|
||||
_TRANSLATIONS_PATH = "/proxy/heygen/v3/video-translations"
|
||||
_SPEECH_PATH = "/proxy/heygen/v3/voices/speech"
|
||||
_AVATARS_PATH = "/proxy/heygen/v3/avatars"
|
||||
_LOOKS_PATH = "/proxy/heygen/v3/avatars/looks"
|
||||
|
||||
_DEFAULT_VOICE_OPTION = "(avatar's default voice)"
|
||||
|
||||
_AVATARS_BY_ENGINE = {
|
||||
e: [label for label, (_aid, _atype, engines) in HEYGEN_AVATAR_MAP.items() if e in engines]
|
||||
for e in ("avatar_iv", "avatar_iii", "avatar_v")
|
||||
}
|
||||
|
||||
|
||||
async def _apply_speech_source(cls: type[IO.ComfyNode], payload: dict, speech: dict, require_voice: bool) -> None:
|
||||
"""Fill script/audio speech fields of a /v3/videos payload from the DynamicCombo dict."""
|
||||
if speech["speech"] == "audio":
|
||||
payload["audio_url"] = await upload_audio_to_comfyapi(
|
||||
cls, speech["audio"], container_format="mp3", codec_name="libmp3lame", mime_type="audio/mpeg"
|
||||
)
|
||||
elif speech["speech"] == "script":
|
||||
validate_string(speech["text"], strip_whitespace=True, min_length=1, max_length=5000)
|
||||
payload["script"] = speech["text"]
|
||||
voice_id = speech.get("custom_voice_id", "").strip()
|
||||
if not voice_id and speech["voice"] != _DEFAULT_VOICE_OPTION:
|
||||
voice_id = HEYGEN_VOICE_GENERAL_MAP[speech["voice"]]
|
||||
if voice_id:
|
||||
payload["voice_id"] = voice_id
|
||||
elif require_voice:
|
||||
raise ValueError("A voice is required when driving the video with a text script.")
|
||||
speed = speech.get("voice_speed", 1.0)
|
||||
if speed != 1.0:
|
||||
payload["voice_settings"] = {"speed": round(speed, 2)}
|
||||
|
||||
|
||||
async def _create_and_poll_video(cls: type[IO.ComfyNode], payload: dict) -> dict:
|
||||
"""POST a /v3/videos payload, poll until terminal, and return the final video data."""
|
||||
created = await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=_VIDEOS_PATH, method="POST", headers={"Idempotency-Key": uuid.uuid4().hex}),
|
||||
data=payload,
|
||||
)
|
||||
video_id = (created.get("data") or {}).get("video_id")
|
||||
if not video_id:
|
||||
raise ValueError(f"HeyGen did not return a video_id: {created}")
|
||||
final = await poll_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=f"{_VIDEOS_PATH}/{video_id}"),
|
||||
status_extractor=lambda r: (r.get("data") or {}).get("status"),
|
||||
queued_statuses=["pending", "waiting"],
|
||||
poll_interval=5.0,
|
||||
)
|
||||
data = final["data"]
|
||||
if not data.get("video_url"):
|
||||
raise ValueError(f"HeyGen returned no video_url for video {video_id}.")
|
||||
return data
|
||||
|
||||
|
||||
async def _resolve_avatar(
|
||||
cls: type[IO.ComfyNode], avatar_label: str, custom_avatar_id: str, engine_choice: str
|
||||
) -> tuple[str, str | None]:
|
||||
"""Resolve (avatar_id, engine_type) from the combo/override + engine widgets."""
|
||||
custom_avatar_id = custom_avatar_id.strip()
|
||||
if custom_avatar_id:
|
||||
look = (
|
||||
await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=f"{_LOOKS_PATH}/{custom_avatar_id}"),
|
||||
final_label_on_success=None,
|
||||
)
|
||||
).get("data") or {}
|
||||
avatar_id = custom_avatar_id
|
||||
avatar_label = look.get("name") or custom_avatar_id
|
||||
supported = look.get("supported_api_engines") or []
|
||||
else:
|
||||
avatar_id, avatar_type, supported = HEYGEN_AVATAR_MAP[avatar_label]
|
||||
|
||||
if engine_choice == "auto":
|
||||
engine = next((e for e in ("avatar_iv", "avatar_iii", "avatar_v") if e in supported), None)
|
||||
else:
|
||||
engine = engine_choice
|
||||
if supported and engine not in supported:
|
||||
raise ValueError(
|
||||
f"Avatar '{avatar_label}' does not support the {engine} engine "
|
||||
f"(supported: {', '.join(supported)}). Set engine to 'auto' to pick "
|
||||
"a compatible engine automatically."
|
||||
)
|
||||
return avatar_id, engine
|
||||
|
||||
|
||||
class HeyGenTalkingPhotoNode(IO.ComfyNode):
|
||||
"""Animate a still image of a person into a lip-synced talking video."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenTalkingPhotoNode",
|
||||
display_name="HeyGen Talking Photo",
|
||||
category="partner/video/HeyGen",
|
||||
description="Animate any image of a person into a lip-synced talking video "
|
||||
"(HeyGen Avatar IV). Drive it with a text script or your own audio.",
|
||||
inputs=[
|
||||
IO.Image.Input(
|
||||
"image",
|
||||
tooltip="Image of a person to animate. Downscaled automatically if larger than 2K.",
|
||||
),
|
||||
IO.DynamicCombo.Input(
|
||||
"speech",
|
||||
display_name="speech source",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"script",
|
||||
[
|
||||
IO.String.Input(
|
||||
"text",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Text for the avatar to speak (up to 5000 characters). "
|
||||
"The generated speech must be at least 1 second long.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"voice",
|
||||
options=HEYGEN_VOICE_GENERAL_OPTIONS,
|
||||
tooltip="Voice for the script (HeyGen's most popular voices).",
|
||||
),
|
||||
IO.String.Input(
|
||||
"custom_voice_id",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. "
|
||||
"Any voice from HeyGen's library (2000+) can be used.",
|
||||
),
|
||||
IO.Float.Input(
|
||||
"voice_speed",
|
||||
default=1.0,
|
||||
min=0.5,
|
||||
max=1.5,
|
||||
step=0.05,
|
||||
optional=True,
|
||||
tooltip="Speech speed multiplier.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"audio",
|
||||
[
|
||||
IO.Audio.Input(
|
||||
"audio",
|
||||
tooltip="Audio for the avatar to lip-sync, up to 10 minutes.",
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
tooltip="Drive the avatar with a text script (HeyGen text-to-speech) or your own audio.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"resolution",
|
||||
options=["720p", "1080p"],
|
||||
default="1080p",
|
||||
optional=True,
|
||||
tooltip="Output video resolution.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"aspect_ratio",
|
||||
options=["auto", "16:9", "9:16", "1:1", "4:5", "5:4"],
|
||||
default="auto",
|
||||
optional=True,
|
||||
tooltip="Output aspect ratio. 'auto' follows the input image.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"expressiveness",
|
||||
options=["low", "medium", "high"],
|
||||
default="low",
|
||||
optional=True,
|
||||
tooltip="How expressive the animated face and gestures are.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
optional=True,
|
||||
tooltip="Not sent to HeyGen; change it to force a re-run.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Video.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr="""{"type":"usd","usd":0.0715,"format":{"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
image: Input.Image,
|
||||
speech: dict,
|
||||
resolution: str = "1080p",
|
||||
aspect_ratio: str = "auto",
|
||||
expressiveness: str = "low",
|
||||
seed: int = 0,
|
||||
) -> IO.NodeOutput:
|
||||
image = downscale_image_tensor_by_max_side(image, max_side=2000)
|
||||
image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None)
|
||||
payload = {
|
||||
"type": "image",
|
||||
"image": {"type": "url", "url": image_url},
|
||||
"resolution": resolution,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"expressiveness": expressiveness,
|
||||
"title": "ComfyUI Talking Photo",
|
||||
}
|
||||
await _apply_speech_source(cls, payload, speech, require_voice=True)
|
||||
video = await _create_and_poll_video(cls, payload)
|
||||
return IO.NodeOutput(await download_url_to_video_output(video["video_url"]))
|
||||
|
||||
|
||||
class HeyGenAvatarVideoNode(IO.ComfyNode):
|
||||
"""Generate a presenter video from a HeyGen avatar look."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenAvatarVideoNode",
|
||||
display_name="HeyGen Avatar Video",
|
||||
category="partner/video/HeyGen",
|
||||
description="Generate a talking-presenter video from a HeyGen avatar. "
|
||||
"Includes HeyGen's most popular public avatars; any look ID can be supplied as an override.",
|
||||
inputs=[
|
||||
IO.DynamicCombo.Input(
|
||||
"engine",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"auto",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"avatar",
|
||||
options=HEYGEN_AVATAR_OPTIONS,
|
||||
tooltip="Avatar look to present the video (curated from HeyGen's "
|
||||
"public library). The best engine the look supports is chosen "
|
||||
"automatically.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"avatar_iv",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"avatar",
|
||||
options=_AVATARS_BY_ENGINE["avatar_iv"],
|
||||
tooltip="Avatar looks that support the Avatar IV engine.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"avatar_iii",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"avatar",
|
||||
options=_AVATARS_BY_ENGINE["avatar_iii"],
|
||||
tooltip="Avatar looks that support the Avatar III engine.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"avatar_v",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"avatar",
|
||||
options=_AVATARS_BY_ENGINE["avatar_v"],
|
||||
tooltip="Avatar looks that support the Avatar V engine.",
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
tooltip="Rendering engine; each choice lists only the avatars that support it. "
|
||||
"'auto' offers every avatar and picks its best engine (Avatar IV preferred). "
|
||||
"Avatar V is highest fidelity, Avatar III is the most affordable.",
|
||||
),
|
||||
IO.String.Input(
|
||||
"custom_avatar_id",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional HeyGen avatar look ID. When set, overrides the avatar selected above. "
|
||||
"Any of HeyGen's 3000+ public looks (or your private avatars) can be used.",
|
||||
),
|
||||
IO.DynamicCombo.Input(
|
||||
"speech",
|
||||
display_name="speech source",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"script",
|
||||
[
|
||||
IO.String.Input(
|
||||
"text",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Text for the avatar to speak (up to 5000 characters). "
|
||||
"The generated speech must be at least 1 second long.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"voice",
|
||||
options=[_DEFAULT_VOICE_OPTION] + HEYGEN_VOICE_GENERAL_OPTIONS,
|
||||
tooltip="Voice for the script. The default option uses the voice HeyGen assigned to the avatar.",
|
||||
),
|
||||
IO.String.Input(
|
||||
"custom_voice_id",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. "
|
||||
"Any voice from HeyGen's library (2000+) can be used.",
|
||||
),
|
||||
IO.Float.Input(
|
||||
"voice_speed",
|
||||
default=1.0,
|
||||
min=0.5,
|
||||
max=1.5,
|
||||
step=0.05,
|
||||
optional=True,
|
||||
tooltip="Speech speed multiplier.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"audio",
|
||||
[
|
||||
IO.Audio.Input(
|
||||
"audio",
|
||||
tooltip="Audio for the avatar to lip-sync, up to 10 minutes.",
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
tooltip="Drive the avatar with a text script (HeyGen text-to-speech) or your own audio.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"resolution",
|
||||
options=["720p", "1080p"],
|
||||
default="1080p",
|
||||
optional=True,
|
||||
tooltip="Output video resolution.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"aspect_ratio",
|
||||
options=["auto", "16:9", "9:16", "1:1", "4:5", "5:4"],
|
||||
default="auto",
|
||||
optional=True,
|
||||
tooltip="Output aspect ratio. 'auto' follows the avatar's source footage.",
|
||||
),
|
||||
IO.String.Input(
|
||||
"background_color",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional solid background color as a hex code (e.g. '#00ff00'). "
|
||||
"Leave empty for the avatar's own background.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
optional=True,
|
||||
tooltip="Not sent to HeyGen; change it to force a re-run.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Video.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
depends_on=IO.PriceBadgeDepends(widgets=["engine"]),
|
||||
expr="""
|
||||
widgets.engine = "avatar_iii"
|
||||
? {"type":"range_usd","min_usd":0.023881,"max_usd":0.061919,"format":{"suffix":"/second"}}
|
||||
: widgets.engine = "avatar_v"
|
||||
? {"type":"usd","usd":0.095381,"format":{"suffix":"/second"}}
|
||||
: widgets.engine = "avatar_iv"
|
||||
? {"type":"range_usd","min_usd":0.0715,"max_usd":0.095381,"format":{"suffix":"/second"}}
|
||||
: {"type":"range_usd","min_usd":0.023881,"max_usd":0.095381,"format":{"suffix":"/second"}}
|
||||
""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
engine: dict,
|
||||
speech: dict,
|
||||
custom_avatar_id: str = "",
|
||||
resolution: str = "1080p",
|
||||
aspect_ratio: str = "auto",
|
||||
background_color: str = "",
|
||||
seed: int = 0,
|
||||
) -> IO.NodeOutput:
|
||||
avatar_id, engine_type = await _resolve_avatar(cls, engine["avatar"], custom_avatar_id, engine["engine"])
|
||||
payload = {
|
||||
"type": "avatar",
|
||||
"avatar_id": avatar_id,
|
||||
"resolution": resolution,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"title": "ComfyUI Avatar Video",
|
||||
}
|
||||
if engine_type:
|
||||
payload["engine"] = {"type": engine_type}
|
||||
background_color = background_color.strip()
|
||||
if background_color:
|
||||
if not background_color.startswith("#"):
|
||||
raise ValueError("background_color must be a hex color code like '#00ff00'.")
|
||||
payload["background"] = {"type": "color", "value": background_color}
|
||||
await _apply_speech_source(cls, payload, speech, require_voice=False)
|
||||
video = await _create_and_poll_video(cls, payload)
|
||||
return IO.NodeOutput(await download_url_to_video_output(video["video_url"]))
|
||||
|
||||
|
||||
class HeyGenCreateAvatarNode(IO.ComfyNode):
|
||||
"""Create a reusable HeyGen avatar from a photo or a text prompt."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenCreateAvatarNode",
|
||||
display_name="HeyGen Create Avatar",
|
||||
category="partner/video/HeyGen",
|
||||
description="Create your own reusable HeyGen avatar from a photo of a person or "
|
||||
"from a text prompt (a generated character). Feed the resulting avatar_id into "
|
||||
"HeyGen Avatar Video's custom_avatar_id — and save the ID somewhere to reuse the "
|
||||
"avatar in future workflows.",
|
||||
inputs=[
|
||||
IO.DynamicCombo.Input(
|
||||
"source",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"prompt",
|
||||
[
|
||||
IO.String.Input(
|
||||
"prompt",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Description of the avatar to generate (up to 1000 characters).",
|
||||
),
|
||||
IO.Autogrow.Input(
|
||||
"reference_images",
|
||||
template=IO.Autogrow.TemplateNames(
|
||||
IO.Image.Input("ref_image"),
|
||||
names=[f"ref_image_{i}" for i in range(1, 4)],
|
||||
min=0,
|
||||
),
|
||||
tooltip="Up to 3 reference images guiding the generated look.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"photo",
|
||||
[
|
||||
IO.Image.Input(
|
||||
"identity_photo",
|
||||
tooltip="Photo of the person to turn into an avatar. "
|
||||
"Downscaled automatically if larger than 2K.",
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
tooltip="Generate a new character from a text prompt, or create the avatar "
|
||||
"from a connected photo of a person.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
IO.String.Output(
|
||||
display_name="avatar_id",
|
||||
tooltip="Avatar look ID. Pass it to HeyGen Avatar Video's custom_avatar_id; "
|
||||
"save it to reuse the avatar later.",
|
||||
),
|
||||
IO.Image.Output(display_name="preview"),
|
||||
],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr="""{"type":"usd","usd":1.43}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
source: dict,
|
||||
) -> IO.NodeOutput:
|
||||
payload: dict = {"name": "ComfyUI Avatar"}
|
||||
if source["source"] == "photo":
|
||||
image = downscale_image_tensor_by_max_side(source["identity_photo"], max_side=2000)
|
||||
image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None)
|
||||
payload["type"] = "photo"
|
||||
payload["file"] = {"type": "url", "url": image_url}
|
||||
else:
|
||||
validate_string(source["prompt"], strip_whitespace=True, min_length=1, max_length=1000)
|
||||
payload["type"] = "prompt"
|
||||
payload["prompt"] = source["prompt"]
|
||||
ref_tensors = [t for t in (source.get("reference_images") or {}).values() if t is not None]
|
||||
if ref_tensors:
|
||||
n_images = sum(get_number_of_images(t) for t in ref_tensors)
|
||||
if n_images > 3:
|
||||
raise ValueError(f"HeyGen accepts at most 3 reference images; got {n_images}.")
|
||||
scaled = [downscale_image_tensor_by_max_side(t, max_side=2000) for t in ref_tensors]
|
||||
ref_urls = await upload_images_to_comfyapi(
|
||||
cls, scaled, max_images=3, mime_type="image/png", total_pixels=None
|
||||
)
|
||||
payload["reference_images"] = [{"type": "url", "url": u} for u in ref_urls]
|
||||
created = await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=_AVATARS_PATH, method="POST"),
|
||||
data=payload,
|
||||
)
|
||||
look_id = ((created.get("data") or {}).get("avatar_item") or {}).get("id")
|
||||
if not look_id:
|
||||
raise ValueError(f"HeyGen did not return an avatar: {created}")
|
||||
final = await poll_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=f"{_LOOKS_PATH}/{look_id}"),
|
||||
# A missing status means the look needed no training and is ready.
|
||||
status_extractor=lambda r: (r.get("data") or {}).get("status") or "completed",
|
||||
failed_statuses=["failed", "pending_consent"],
|
||||
poll_interval=5.0,
|
||||
)
|
||||
data = final["data"]
|
||||
if data.get("preview_image_url"):
|
||||
preview = await download_url_to_image_tensor(data["preview_image_url"])
|
||||
else:
|
||||
preview = torch.zeros(1, 64, 64, 3)
|
||||
PromptServer.instance.send_progress_text(
|
||||
f"Please save the avatar_id for reuse.\n\navatar_id: {look_id}",
|
||||
cls.hidden.unique_id,
|
||||
)
|
||||
return IO.NodeOutput(look_id, preview)
|
||||
|
||||
|
||||
class HeyGenVideoTranslateNode(IO.ComfyNode):
|
||||
"""Translate a spoken video into another language with voice cloning and lip sync."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenVideoTranslateNode",
|
||||
display_name="HeyGen Video Translate",
|
||||
category="partner/video/HeyGen",
|
||||
description="Translate a spoken video into another language. Clones the original "
|
||||
"speaker's voice and re-animates the mouth to match the translated speech.",
|
||||
inputs=[
|
||||
IO.Video.Input(
|
||||
"video",
|
||||
tooltip="Video with speech to translate.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"output_language",
|
||||
options=HEYGEN_TRANSLATE_LANGUAGES,
|
||||
tooltip="Target language for the translated video.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"mode",
|
||||
options=["speed", "precision"],
|
||||
default="speed",
|
||||
tooltip="'speed' is faster; 'precision' produces higher-quality lip sync at twice the price.",
|
||||
),
|
||||
IO.Boolean.Input(
|
||||
"translate_audio_only",
|
||||
default=False,
|
||||
optional=True,
|
||||
tooltip="Only swap the audio track, keeping the original mouth movements (no lip sync).",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_count",
|
||||
default=0,
|
||||
min=0,
|
||||
max=10,
|
||||
optional=True,
|
||||
tooltip="Number of speakers in the video. 0 = detect automatically.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
optional=True,
|
||||
tooltip="Not sent to HeyGen; change it to force a re-run.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Video.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
depends_on=IO.PriceBadgeDepends(widgets=["mode"]),
|
||||
expr="""{"type":"usd","usd": widgets.mode = "precision" ? 0.095381 : 0.047619,"""
|
||||
""""format":{"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
video: Input.Video,
|
||||
output_language: str,
|
||||
mode: str,
|
||||
translate_audio_only: bool = False,
|
||||
speaker_count: int = 0,
|
||||
seed: int = 0,
|
||||
) -> IO.NodeOutput:
|
||||
video_url = await upload_video_to_comfyapi(cls, video)
|
||||
payload = {
|
||||
"video": {"type": "url", "url": video_url},
|
||||
"output_languages": [output_language],
|
||||
"mode": mode,
|
||||
"translate_audio_only": translate_audio_only,
|
||||
"title": "ComfyUI Video Translate",
|
||||
}
|
||||
if speaker_count > 0:
|
||||
payload["speaker_num"] = speaker_count
|
||||
created = await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=_TRANSLATIONS_PATH, method="POST"),
|
||||
data=payload,
|
||||
)
|
||||
translation_ids = (created.get("data") or {}).get("video_translation_ids") or []
|
||||
if not translation_ids:
|
||||
raise ValueError(f"HeyGen did not return a translation ID: {created}")
|
||||
final = await poll_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=f"{_TRANSLATIONS_PATH}/{translation_ids[0]}"),
|
||||
status_extractor=lambda r: (r.get("data") or {}).get("status"),
|
||||
queued_statuses=["pending"],
|
||||
poll_interval=5.0,
|
||||
)
|
||||
data = final["data"]
|
||||
if not data.get("video_url"):
|
||||
raise ValueError(f"HeyGen returned no video_url for translation {translation_ids[0]}.")
|
||||
return IO.NodeOutput(await download_url_to_video_output(data["video_url"]))
|
||||
|
||||
|
||||
class HeyGenTextToSpeechNode(IO.ComfyNode):
|
||||
"""Synthesize speech audio from text with HeyGen's Starfish TTS engine."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenTextToSpeechNode",
|
||||
display_name="HeyGen Text to Speech",
|
||||
category="partner/audio/HeyGen",
|
||||
description="Generate speech audio from text using HeyGen's Starfish TTS engine. "
|
||||
"Includes HeyGen's most popular voices across 17 languages.",
|
||||
inputs=[
|
||||
IO.String.Input(
|
||||
"text",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Text to synthesize (up to 5000 characters). The generated speech "
|
||||
"must be at least 1 second long.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"voice",
|
||||
options=HEYGEN_VOICE_TTS_OPTIONS,
|
||||
tooltip="Voice to use (curated from HeyGen's most popular Starfish-compatible voices).",
|
||||
),
|
||||
IO.String.Input(
|
||||
"custom_voice_id",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. "
|
||||
"The voice must support the Starfish engine.",
|
||||
),
|
||||
IO.Float.Input(
|
||||
"speed",
|
||||
default=1.0,
|
||||
min=0.5,
|
||||
max=2.0,
|
||||
step=0.05,
|
||||
optional=True,
|
||||
tooltip="Speech speed multiplier.",
|
||||
),
|
||||
IO.Boolean.Input(
|
||||
"ssml",
|
||||
default=False,
|
||||
optional=True,
|
||||
tooltip="Treat the text as SSML markup (for pauses, emphasis, and pronunciation control).",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
optional=True,
|
||||
tooltip="Not sent to HeyGen; change it to force a re-run.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Audio.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr="""{"type":"usd","usd":0.00095381,"format":{"approximate":true,"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
text: str,
|
||||
voice: str,
|
||||
custom_voice_id: str = "",
|
||||
speed: float = 1.0,
|
||||
ssml: bool = False,
|
||||
seed: int = 0,
|
||||
) -> IO.NodeOutput:
|
||||
validate_string(text, strip_whitespace=True, min_length=1, max_length=5000)
|
||||
payload = {
|
||||
"text": text,
|
||||
"voice_id": custom_voice_id.strip() or HEYGEN_VOICE_TTS_MAP[voice],
|
||||
"speed": round(speed, 2),
|
||||
}
|
||||
if ssml:
|
||||
payload["input_type"] = "ssml"
|
||||
response = await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=_SPEECH_PATH, method="POST"),
|
||||
data=payload,
|
||||
)
|
||||
audio_url = (response.get("data") or {}).get("audio_url")
|
||||
if not audio_url:
|
||||
raise ValueError(f"HeyGen did not return an audio_url: {response}")
|
||||
audio_bytes = await download_url_as_bytesio(audio_url)
|
||||
return IO.NodeOutput(audio_bytes_to_audio_input(audio_bytes.getvalue()))
|
||||
|
||||
|
||||
class HeyGenExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[IO.ComfyNode]]:
|
||||
return [
|
||||
HeyGenTalkingPhotoNode,
|
||||
HeyGenAvatarVideoNode,
|
||||
HeyGenCreateAvatarNode,
|
||||
HeyGenVideoTranslateNode,
|
||||
HeyGenTextToSpeechNode,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> HeyGenExtension:
|
||||
return HeyGenExtension()
|
||||
|
|
@ -45,27 +45,39 @@ class _ModelSpec:
|
|||
|
||||
|
||||
MODELS: list[_ModelSpec] = [
|
||||
_ModelSpec("anthropic/claude-opus-4.7", "frontier_reasoning", 0.000005, 0.000025, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.5-pro", "frontier_reasoning", 0.00003, 0.00018, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.5", "frontier_reasoning", 0.000005, 0.00003, max_images=20),
|
||||
_ModelSpec("google/gemini-3.5-flash", "reasoning", 0.0000015, 0.000009, max_images=20, max_videos=4),
|
||||
_ModelSpec("x-ai/grok-4.20", "reasoning", 0.00000125, 0.0000025, max_images=20),
|
||||
_ModelSpec("x-ai/grok-4.3", "reasoning", 0.00000125, 0.0000025, max_images=20),
|
||||
_ModelSpec("deepseek/deepseek-v4-pro", "reasoning", 0.000000435, 0.00000087),
|
||||
_ModelSpec("deepseek/deepseek-v4-flash", "reasoning", 0.000000112, 0.000000224),
|
||||
_ModelSpec("deepseek/deepseek-v3.2", "reasoning", 0.000000252, 0.000000378),
|
||||
_ModelSpec("qwen/qwen3.6-max-preview", "reasoning", 0.00000104, 0.00000624),
|
||||
_ModelSpec("qwen/qwen3.6-plus", "reasoning", 0.000000325, 0.00000195, max_images=10, max_videos=4),
|
||||
_ModelSpec("qwen/qwen3.6-flash", "reasoning", 0.0000001875, 0.000001125, max_images=10, max_videos=4),
|
||||
_ModelSpec("mistralai/mistral-large-2512", "standard", 0.0000005, 0.0000015, max_images=8),
|
||||
_ModelSpec("mistralai/mistral-medium-3-5", "reasoning", 0.0000015, 0.0000075, max_images=8),
|
||||
_ModelSpec("z-ai/glm-4.6", "reasoning", 0.00000043, 0.00000174),
|
||||
_ModelSpec("z-ai/glm-5", "reasoning", 0.0000006, 0.00000192),
|
||||
_ModelSpec("moonshotai/kimi-k2.6", "reasoning", 0.00000073, 0.00000349, max_images=10),
|
||||
_ModelSpec("moonshotai/kimi-k2-thinking", "reasoning", 0.0000006, 0.0000025),
|
||||
_ModelSpec("perplexity/sonar-pro", "perplexity", 0.000003, 0.000015),
|
||||
_ModelSpec("perplexity/sonar-reasoning-pro", "perplexity_reasoning", 0.000002, 0.000008),
|
||||
_ModelSpec("perplexity/sonar-deep-research", "perplexity_reasoning", 0.000002, 0.000008),
|
||||
_ModelSpec("anthropic/claude-opus-4.8", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20),
|
||||
_ModelSpec("anthropic/claude-opus-4.7", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20),
|
||||
_ModelSpec("anthropic/claude-fable-5", "frontier_reasoning", 0.0000143, 0.0000715, max_images=20),
|
||||
_ModelSpec("anthropic/claude-sonnet-5", "frontier_reasoning", 0.00000286, 0.0000143, max_images=20),
|
||||
_ModelSpec("anthropic/claude-haiku-4.5", "frontier_reasoning", 0.00000143, 0.00000715, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.6-sol-pro", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.6-sol", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.6-terra-pro", "frontier_reasoning", 0.000003575, 0.00002145, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.6-terra", "frontier_reasoning", 0.000003575, 0.00002145, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.6-luna-pro", "frontier_reasoning", 0.00000143, 0.00000858, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.6-luna", "frontier_reasoning", 0.00000143, 0.00000858, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.5-pro", "frontier_reasoning", 0.0000429, 0.0002574, max_images=20),
|
||||
_ModelSpec("openai/gpt-5.5", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20),
|
||||
_ModelSpec("google/gemini-3.5-flash", "reasoning", 0.000002145, 0.00001287, max_images=20, max_videos=4),
|
||||
_ModelSpec("x-ai/grok-4.5", "reasoning", 0.00000286, 0.00000858, max_images=20),
|
||||
_ModelSpec("x-ai/grok-4.20", "reasoning", 0.0000017875, 0.000003575, max_images=20),
|
||||
_ModelSpec("x-ai/grok-4.3", "reasoning", 0.0000017875, 0.000003575, max_images=20),
|
||||
_ModelSpec("deepseek/deepseek-v4-pro", "reasoning", 0.00000062205, 0.0000012441),
|
||||
_ModelSpec("deepseek/deepseek-v4-flash", "reasoning", 0.00000016016, 0.00000032032),
|
||||
_ModelSpec("deepseek/deepseek-v3.2", "reasoning", 0.00000036036, 0.00000054054),
|
||||
_ModelSpec("qwen/qwen3.6-max-preview", "reasoning", 0.0000014872, 0.0000089232),
|
||||
_ModelSpec("qwen/qwen3.6-plus", "reasoning", 0.00000046475, 0.0000027885, max_images=10, max_videos=4),
|
||||
_ModelSpec("qwen/qwen3.6-flash", "reasoning", 0.000000268125, 0.00000160875, max_images=10, max_videos=4),
|
||||
_ModelSpec("mistralai/mistral-large-2512", "standard", 0.000000715, 0.000002145, max_images=8),
|
||||
_ModelSpec("mistralai/mistral-medium-3-5", "reasoning", 0.000002145, 0.000010725, max_images=8),
|
||||
_ModelSpec("z-ai/glm-4.6", "reasoning", 0.0000006149, 0.0000024882),
|
||||
_ModelSpec("z-ai/glm-5", "reasoning", 0.000000858, 0.0000027456),
|
||||
_ModelSpec("moonshotai/kimi-k3", "reasoning", 0.00000429, 0.00002145, max_images=10),
|
||||
_ModelSpec("moonshotai/kimi-k2.6", "reasoning", 0.0000010439, 0.0000049907, max_images=10),
|
||||
_ModelSpec("moonshotai/kimi-k2-thinking", "reasoning", 0.000000858, 0.000003575),
|
||||
_ModelSpec("perplexity/sonar-pro", "perplexity", 0.00000429, 0.00002145),
|
||||
_ModelSpec("perplexity/sonar-reasoning-pro", "perplexity_reasoning", 0.00000286, 0.00001144),
|
||||
_ModelSpec("perplexity/sonar-deep-research", "perplexity_reasoning", 0.00000286, 0.00001144),
|
||||
]
|
||||
|
||||
_MODELS_BY_SLUG: dict[str, _ModelSpec] = {m.slug: m for m in MODELS}
|
||||
|
|
@ -148,7 +160,7 @@ def _build_model_options() -> list[IO.DynamicCombo.Option]:
|
|||
|
||||
def _calculate_price(response: OpenRouterChatResponse) -> float | None:
|
||||
if response.usage and response.usage.cost is not None:
|
||||
return float(response.usage.cost)
|
||||
return float(response.usage.cost) * 1.43
|
||||
return None
|
||||
|
||||
|
||||
|
|
@ -269,8 +281,8 @@ class OpenRouterLLMNode(IO.ComfyNode):
|
|||
essentials_category="Text Generation",
|
||||
description=(
|
||||
"Generate text responses through OpenRouter. Routes to a curated set of popular "
|
||||
"models from xAI, DeepSeek, Qwen, Mistral, Z.AI (GLM), Moonshot (Kimi), and "
|
||||
"Perplexity Sonar."
|
||||
"models from Anthropic (Claude), OpenAI (GPT), Google (Gemini), xAI (Grok), "
|
||||
"DeepSeek, Qwen, Mistral, Z.AI (GLM), Moonshot (Kimi), and Perplexity Sonar."
|
||||
),
|
||||
inputs=[
|
||||
IO.String.Input(
|
||||
|
|
|
|||
|
|
@ -194,6 +194,7 @@ class RunwayImageToVideoNodeGen3a(IO.ComfyNode):
|
|||
depends_on=IO.PriceBadgeDepends(widgets=["duration"]),
|
||||
expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""",
|
||||
),
|
||||
is_deprecated=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
|
|
@ -390,6 +391,7 @@ class RunwayFirstLastFrameNode(IO.ComfyNode):
|
|||
depends_on=IO.PriceBadgeDepends(widgets=["duration"]),
|
||||
expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""",
|
||||
),
|
||||
is_deprecated=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import logging
|
|||
import os
|
||||
import json
|
||||
|
||||
import av
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
|
@ -9,7 +10,7 @@ from typing_extensions import override
|
|||
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
from comfy_api.latest import ComfyExtension, io, Input, InputImpl, Types
|
||||
|
||||
|
||||
def load_and_process_images(image_files, input_dir):
|
||||
|
|
@ -42,6 +43,38 @@ def load_and_process_images(image_files, input_dir):
|
|||
return output_images
|
||||
|
||||
|
||||
VALID_VIDEO_EXTENSIONS = [".mp4", ".avi", ".mov", ".webm", ".mkv", ".flv"]
|
||||
|
||||
|
||||
def _decode_selected_frames(video: Input.Video, indices: list[int]) -> Input.Video:
|
||||
"""Decode only the requested frame indices from a video.
|
||||
|
||||
Opens the underlying container once, decodes frames in presentation order,
|
||||
keeps only the ones whose index is in ``indices``, and returns the result
|
||||
wrapped in a VideoFromComponents so it still satisfies the VideoInput
|
||||
contract for downstream nodes.
|
||||
"""
|
||||
indices_sorted = sorted(set(indices))
|
||||
max_idx = indices_sorted[-1]
|
||||
source = video.get_stream_source()
|
||||
|
||||
frames_by_idx: dict[int, torch.Tensor] = {}
|
||||
with av.open(source, mode="r") as container:
|
||||
stream = container.streams.video[0]
|
||||
wanted = set(indices_sorted)
|
||||
for frame_idx, frame in enumerate(container.decode(stream)):
|
||||
if frame_idx in wanted:
|
||||
img = frame.to_ndarray(format="rgb24")
|
||||
frames_by_idx[frame_idx] = torch.from_numpy(img.copy()).float() / 255.0
|
||||
if frame_idx >= max_idx:
|
||||
break
|
||||
|
||||
stacked = torch.stack([frames_by_idx[i] for i in indices])
|
||||
return InputImpl.VideoFromComponents(
|
||||
Types.VideoComponents(images=stacked, frame_rate=video.get_frame_rate())
|
||||
)
|
||||
|
||||
|
||||
class LoadImageDataSetFromFolderNode(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
|
|
@ -157,6 +190,116 @@ class LoadImageTextDataSetFromFolderNode(io.ComfyNode):
|
|||
return io.NodeOutput(output_tensor, captions)
|
||||
|
||||
|
||||
class LoadVideoDataSetFromFolderNode(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LoadVideoDataSetFromFolder",
|
||||
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
|
||||
display_name="Load Video (from Folder)",
|
||||
category="video",
|
||||
description="Load a dataset of videos from a specified folder and return a list of videos. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Combo.Input(
|
||||
"folder",
|
||||
options=folder_paths.get_input_subfolders(),
|
||||
tooltip="The folder containing video files.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(
|
||||
display_name="videos",
|
||||
is_output_list=True,
|
||||
tooltip="Lazy video references; frames are decoded only when needed downstream.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, folder):
|
||||
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
|
||||
video_files = sorted([
|
||||
f for f in os.listdir(sub_input_dir)
|
||||
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
|
||||
])
|
||||
|
||||
if not video_files:
|
||||
raise ValueError(f"No video files found in {sub_input_dir}")
|
||||
|
||||
videos = [InputImpl.VideoFromFile(os.path.join(sub_input_dir, f)) for f in video_files]
|
||||
logging.info(f"Loaded {len(videos)} lazy video references from {sub_input_dir}")
|
||||
return io.NodeOutput(videos)
|
||||
|
||||
|
||||
class LoadVideoTextDataSetFromFolderNode(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LoadVideoTextDataSetFromFolder",
|
||||
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
|
||||
display_name="Load Video-Text (from Folder)",
|
||||
category="video",
|
||||
description="Load a dataset of pairs of videos and text captions from a specified folder and return them as a list. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Combo.Input(
|
||||
"folder",
|
||||
options=folder_paths.get_input_subfolders(),
|
||||
tooltip="The folder containing video files and .txt captions.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(
|
||||
display_name="videos",
|
||||
is_output_list=True,
|
||||
tooltip="Lazy video references; frames are decoded only when needed downstream.",
|
||||
),
|
||||
io.String.Output(
|
||||
display_name="texts",
|
||||
is_output_list=True,
|
||||
tooltip="List of text captions.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, folder):
|
||||
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
|
||||
|
||||
video_files = []
|
||||
for item in sorted(os.listdir(sub_input_dir)):
|
||||
path = os.path.join(sub_input_dir, item)
|
||||
if any(item.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS):
|
||||
video_files.append(path)
|
||||
elif os.path.isdir(path):
|
||||
# Support kohya-ss/sd-scripts folder structure: {repeat}_{desc}/
|
||||
repeat = 1
|
||||
if item.split("_")[0].isdigit():
|
||||
repeat = int(item.split("_")[0])
|
||||
video_files.extend([
|
||||
os.path.join(path, f)
|
||||
for f in sorted(os.listdir(path))
|
||||
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
|
||||
] * repeat)
|
||||
|
||||
if not video_files:
|
||||
raise ValueError(f"No video files found in {sub_input_dir}")
|
||||
|
||||
captions = []
|
||||
for vf in video_files:
|
||||
caption_path = os.path.splitext(vf)[0] + ".txt"
|
||||
if os.path.exists(caption_path):
|
||||
with open(caption_path, "r", encoding="utf-8") as f:
|
||||
captions.append(f.read().strip())
|
||||
else:
|
||||
captions.append("")
|
||||
|
||||
videos = [InputImpl.VideoFromFile(vf) for vf in video_files]
|
||||
logging.info(f"Loaded {len(videos)} lazy video references with captions from {sub_input_dir}")
|
||||
return io.NodeOutput(videos, captions)
|
||||
|
||||
|
||||
def save_images_to_folder(image_list, output_dir, prefix="image", overwrite=True):
|
||||
"""Utility function to save a list of image tensors to disk.
|
||||
|
||||
|
|
@ -470,7 +613,15 @@ class ImageProcessingNode(io.ComfyNode):
|
|||
|
||||
@classmethod
|
||||
def execute(cls, images, **kwargs):
|
||||
"""Execute the node. Routes to _process or _group_process based on mode."""
|
||||
"""Execute the node. Routes to _process or _group_process based on mode.
|
||||
|
||||
For individual processing (_process), automatically handles multi-frame
|
||||
inputs (video tensors [T, H, W, C]) by applying _process per-frame and
|
||||
concatenating the results. This allows all spatial transform nodes to
|
||||
work with video without modification. Nodes that natively handle batched
|
||||
tensors (e.g. pure tensor math) can set per_frame_process = False to
|
||||
skip the per-frame loop.
|
||||
"""
|
||||
is_group = cls._detect_processing_mode()
|
||||
|
||||
if is_group:
|
||||
|
|
@ -489,7 +640,16 @@ class ImageProcessingNode(io.ComfyNode):
|
|||
result = cls._group_process(images, **params)
|
||||
else:
|
||||
# Individual processing: images is single item, call _process
|
||||
result = cls._process(images, **params)
|
||||
# Auto-loop over frames for multi-frame inputs (video [T, H, W, C])
|
||||
# so that PIL-based spatial transforms work per-frame automatically.
|
||||
if images.shape[0] > 1 and getattr(cls, 'per_frame_process', True):
|
||||
results = []
|
||||
for i in range(images.shape[0]):
|
||||
frame_result = cls._process(images[i:i + 1], **params)
|
||||
results.append(frame_result)
|
||||
result = torch.cat(results, dim=0)
|
||||
else:
|
||||
result = cls._process(images, **params)
|
||||
|
||||
return io.NodeOutput(result)
|
||||
|
||||
|
|
@ -803,6 +963,7 @@ class NormalizeImagesNode(ImageProcessingNode):
|
|||
display_name = "Normalize Image Colors"
|
||||
category = "image/color"
|
||||
description = "Normalize images using mean and standard deviation."
|
||||
per_frame_process = False # Pure tensor math, handles any batch size
|
||||
extra_inputs = [
|
||||
io.Float.Input(
|
||||
"mean",
|
||||
|
|
@ -833,6 +994,7 @@ class AdjustBrightnessNode(ImageProcessingNode):
|
|||
display_name = "Adjust Brightness"
|
||||
category="image/adjustments"
|
||||
description = "Adjust the brightness of an image."
|
||||
per_frame_process = False # Pure tensor math, handles any batch size
|
||||
extra_inputs = [
|
||||
io.Float.Input(
|
||||
"factor",
|
||||
|
|
@ -854,6 +1016,7 @@ class AdjustContrastNode(ImageProcessingNode):
|
|||
display_name = "Adjust Contrast"
|
||||
category="image/adjustments"
|
||||
description = "Adjust the contrast of an image."
|
||||
per_frame_process = False # Pure tensor math, handles any batch size
|
||||
extra_inputs = [
|
||||
io.Float.Input(
|
||||
"factor",
|
||||
|
|
@ -935,6 +1098,261 @@ class ShuffleImageTextDatasetNode(io.ComfyNode):
|
|||
return io.NodeOutput(shuffled_images, shuffled_texts)
|
||||
|
||||
|
||||
# ========== Video Processing Nodes ==========
|
||||
|
||||
|
||||
class VideoFrameSampleNode(io.ComfyNode):
|
||||
"""Sample a fixed number of frames from a video using various strategies.
|
||||
|
||||
For contiguous strategies ("head"/"tail") the result is a fully lazy
|
||||
VideoInput (no frames decoded). For non-contiguous strategies
|
||||
("uniform"/"random") only the selected indices are decoded.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="VideoFrameSample",
|
||||
search_aliases=["sample frames", "extract frames"],
|
||||
display_name="Sample Video Frame",
|
||||
category="video",
|
||||
description="Sample a fixed number of frames from a video using various strategies.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Video.Input("video", tooltip="Input video."),
|
||||
io.Int.Input(
|
||||
"num_frames",
|
||||
default=16,
|
||||
min=1,
|
||||
max=9999,
|
||||
tooltip="Number of frames to sample.",
|
||||
),
|
||||
io.Combo.Input(
|
||||
"strategy",
|
||||
options=["uniform", "head", "tail", "random"],
|
||||
default="uniform",
|
||||
tooltip="uniform: evenly spaced, head: first N, tail: last N, random: random sorted.",
|
||||
),
|
||||
io.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=0xFFFFFFFFFFFFFFFF,
|
||||
tooltip="Random seed (only used with 'random' strategy).",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(display_name="video", tooltip="Sampled video."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, video, num_frames, strategy, seed):
|
||||
total_frames = video.get_frame_count()
|
||||
num_frames = min(num_frames, total_frames)
|
||||
fps = float(video.get_frame_rate())
|
||||
|
||||
if strategy == "head":
|
||||
return io.NodeOutput(
|
||||
video.as_trimmed(0.0, num_frames / fps, strict_duration=False)
|
||||
)
|
||||
if strategy == "tail":
|
||||
start_t = (total_frames - num_frames) / fps
|
||||
return io.NodeOutput(
|
||||
video.as_trimmed(start_t, num_frames / fps, strict_duration=False)
|
||||
)
|
||||
|
||||
if strategy == "uniform":
|
||||
if num_frames == 1:
|
||||
indices = [total_frames // 2]
|
||||
else:
|
||||
indices = [round(i * (total_frames - 1) / (num_frames - 1)) for i in range(num_frames)]
|
||||
elif strategy == "random":
|
||||
rng = np.random.RandomState(seed % (2**32 - 1))
|
||||
indices = sorted(rng.choice(total_frames, size=num_frames, replace=False).tolist())
|
||||
else:
|
||||
raise ValueError(f"Unknown strategy: {strategy}")
|
||||
|
||||
return io.NodeOutput(_decode_selected_frames(video, indices))
|
||||
|
||||
|
||||
class VideoTemporalCropNode(io.ComfyNode):
|
||||
"""Crop a continuous range of frames from a video (fully lazy)."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="VideoTemporalCrop",
|
||||
search_aliases=["crop", "crop video", "temporal crop", "truncate video"],
|
||||
display_name="Crop Video (Temporal)",
|
||||
category="video/transform",
|
||||
description="Crop a continuous range of frames from a video.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Video.Input("video", tooltip="Input video."),
|
||||
io.Int.Input(
|
||||
"start_frame",
|
||||
default=0,
|
||||
min=0,
|
||||
max=99999,
|
||||
tooltip="Starting frame index.",
|
||||
),
|
||||
io.Int.Input(
|
||||
"length",
|
||||
default=16,
|
||||
min=1,
|
||||
max=99999,
|
||||
tooltip="Number of frames to keep.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, video, start_frame, length):
|
||||
total_frames = video.get_frame_count()
|
||||
fps = float(video.get_frame_rate())
|
||||
start_frame = min(start_frame, max(total_frames - 1, 0))
|
||||
length = min(length, total_frames - start_frame)
|
||||
return io.NodeOutput(
|
||||
video.as_trimmed(start_frame / fps, length / fps, strict_duration=False)
|
||||
)
|
||||
|
||||
|
||||
class VideoRandomTemporalCropNode(io.ComfyNode):
|
||||
"""Randomly crop a continuous range of frames from a video (fully lazy)."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="VideoRandomTemporalCrop",
|
||||
search_aliases=["crop", "crop video", "temporal crop", "truncate video", "random crop"],
|
||||
display_name="Crop Video (Temporal Random)",
|
||||
category="video/transform",
|
||||
description="Randomly crop a continuous range of frames from a video.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Video.Input("video", tooltip="Input video."),
|
||||
io.Int.Input(
|
||||
"length",
|
||||
default=16,
|
||||
min=1,
|
||||
max=99999,
|
||||
tooltip="Number of frames to keep.",
|
||||
),
|
||||
io.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=0xFFFFFFFFFFFFFFFF,
|
||||
tooltip="Random seed.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, video, length, seed):
|
||||
total_frames = video.get_frame_count()
|
||||
fps = float(video.get_frame_rate())
|
||||
length = min(length, total_frames)
|
||||
max_start = total_frames - length
|
||||
rng = np.random.RandomState(seed % (2**32 - 1))
|
||||
start = rng.randint(0, max_start + 1) if max_start > 0 else 0
|
||||
return io.NodeOutput(
|
||||
video.as_trimmed(start / fps, length / fps, strict_duration=False)
|
||||
)
|
||||
|
||||
|
||||
class ShuffleVideoDatasetNode(io.ComfyNode):
|
||||
"""Randomly shuffle the order of videos in the dataset."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ShuffleVideoDataset",
|
||||
search_aliases=["shuffle", "randomize", "mix"],
|
||||
display_name="Shuffle Videos List",
|
||||
category="video/batch",
|
||||
description="Randomly shuffle the order of videos in a list.",
|
||||
is_experimental=True,
|
||||
is_input_list=True,
|
||||
inputs=[
|
||||
io.Video.Input("videos", tooltip="List of videos to shuffle."),
|
||||
io.Int.Input(
|
||||
"seed", default=0, min=0, max=0xFFFFFFFFFFFFFFFF, tooltip="Random seed."
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(
|
||||
display_name="videos",
|
||||
is_output_list=True,
|
||||
tooltip="Shuffled videos",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, videos, seed):
|
||||
seed = seed[0] if isinstance(seed, list) else seed
|
||||
np.random.seed(seed % (2**32 - 1))
|
||||
indices = np.random.permutation(len(videos))
|
||||
return io.NodeOutput([videos[i] for i in indices])
|
||||
|
||||
|
||||
class ShuffleVideoTextDatasetNode(io.ComfyNode):
|
||||
"""Shuffle videos and their captions together, preserving pairs."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ShuffleVideoTextDataset",
|
||||
search_aliases=["shuffle", "randomize", "mix"],
|
||||
display_name="Shuffle Pairs of Video-Text",
|
||||
category="dataset/video",
|
||||
description="Randomly shuffle the order of pairs of video-text in a list.",
|
||||
is_experimental=True,
|
||||
is_input_list=True,
|
||||
inputs=[
|
||||
io.Video.Input("videos", tooltip="List of videos to shuffle."),
|
||||
io.String.Input("texts", tooltip="List of texts to shuffle."),
|
||||
io.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=0xFFFFFFFFFFFFFFFF,
|
||||
tooltip="Random seed.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(
|
||||
display_name="videos",
|
||||
is_output_list=True,
|
||||
tooltip="Shuffled videos",
|
||||
),
|
||||
io.String.Output(
|
||||
display_name="texts",
|
||||
is_output_list=True,
|
||||
tooltip="Shuffled texts",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, videos, texts, seed):
|
||||
seed = seed[0] if isinstance(seed, list) else seed
|
||||
np.random.seed(seed % (2**32 - 1))
|
||||
indices = np.random.permutation(len(videos))
|
||||
return io.NodeOutput(
|
||||
[videos[i] for i in indices],
|
||||
[texts[i] for i in indices],
|
||||
)
|
||||
|
||||
|
||||
# ========== Text Transform Nodes ==========
|
||||
|
||||
|
||||
|
|
@ -1608,7 +2026,10 @@ class DatasetExtension(ComfyExtension):
|
|||
LoadImageTextDataSetFromFolderNode,
|
||||
SaveImageDataSetToFolderNode,
|
||||
SaveImageTextDataSetToFolderNode,
|
||||
# Image transform nodes
|
||||
# Video data loading nodes
|
||||
LoadVideoDataSetFromFolderNode,
|
||||
LoadVideoTextDataSetFromFolderNode,
|
||||
# Image transform nodes (auto-handle video via per-frame processing)
|
||||
ResizeImagesByShorterEdgeNode,
|
||||
ResizeImagesByLongerEdgeNode,
|
||||
CenterCropImagesNode,
|
||||
|
|
@ -1618,6 +2039,12 @@ class DatasetExtension(ComfyExtension):
|
|||
AdjustContrastNode,
|
||||
ShuffleDatasetNode,
|
||||
ShuffleImageTextDatasetNode,
|
||||
# Video processing nodes (lazy VideoInput in/out)
|
||||
VideoFrameSampleNode,
|
||||
VideoTemporalCropNode,
|
||||
VideoRandomTemporalCropNode,
|
||||
ShuffleVideoDatasetNode,
|
||||
ShuffleVideoTextDatasetNode,
|
||||
# Text transform nodes
|
||||
TextToLowercaseNode,
|
||||
TextToUppercaseNode,
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
|
|||
Apply frequency-dependent scaling to an image tensor using Fourier transforms.
|
||||
|
||||
Parameters:
|
||||
x: Input tensor of shape (B, C, H, W)
|
||||
x: Input tensor of shape (..., H, W)
|
||||
scale_low: Scaling factor for low-frequency components (default: 1.0)
|
||||
scale_high: Scaling factor for high-frequency components (default: 1.5)
|
||||
freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20)
|
||||
|
|
@ -31,8 +31,8 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
|
|||
# Initialize mask with high-frequency scaling factor
|
||||
mask = torch.ones(x_freq.shape, device=device) * scale_high
|
||||
m = mask
|
||||
for d in range(len(x_freq.shape) - 2):
|
||||
dim = d + 2
|
||||
for d in range(2):
|
||||
dim = len(x_freq.shape) - 2 + d
|
||||
cc = x_freq.shape[dim] // 2
|
||||
f_c = min(freq_cutoff, cc)
|
||||
m = m.narrow(dim, cc - f_c, f_c * 2)
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ import comfy.ldm.common_dit
|
|||
import comfy.latent_formats
|
||||
import comfy.ldm.lumina.controlnet
|
||||
import comfy.ldm.supir.supir_modules
|
||||
import comfy.ldm.anima.lllite
|
||||
import comfy.ldm.wan.uni3c
|
||||
from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel
|
||||
from comfy_api.latest import io
|
||||
from comfy.ldm.supir.supir_patch import SUPIRPatch
|
||||
|
|
@ -236,10 +238,12 @@ class ModelPatchLoader:
|
|||
|
||||
def load_model_patch(self, name):
|
||||
model_patch_path = folder_paths.get_full_path_or_raise("model_patches", name)
|
||||
sd = comfy.utils.load_torch_file(model_patch_path, safe_load=True)
|
||||
sd, metadata = comfy.utils.load_torch_file(model_patch_path, safe_load=True, return_metadata=True)
|
||||
dtype = comfy.utils.weight_dtype(sd)
|
||||
|
||||
if 'controlnet_blocks.0.y_rms.weight' in sd:
|
||||
if 'lllite_conditioning1.conv1.weight' in sd:
|
||||
model = comfy.ldm.anima.lllite.AnimaLLLite(sd, metadata, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast)
|
||||
elif 'controlnet_blocks.0.y_rms.weight' in sd:
|
||||
additional_in_dim = sd["img_in.weight"].shape[1] - 64
|
||||
model = QwenImageBlockWiseControlNet(additional_in_dim=additional_in_dim, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast)
|
||||
elif 'feature_embedder.mid_layer_norm.bias' in sd:
|
||||
|
|
@ -261,6 +265,37 @@ class ModelPatchLoader:
|
|||
if torch.count_nonzero(ref_weight) == 0:
|
||||
config['broken'] = True
|
||||
model = comfy.ldm.lumina.controlnet.ZImage_Control(device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast, **config)
|
||||
elif 'controlnet_patch_embedding.weight' in sd: # Uni3C controlnet for Wan
|
||||
attn_key_replace = {".self_attn.to_q.": ".self_attn.q.",
|
||||
".self_attn.to_k.": ".self_attn.k.",
|
||||
".self_attn.to_v.": ".self_attn.v.",
|
||||
".self_attn.to_out.0.": ".self_attn.o."}
|
||||
converted_sd = {}
|
||||
for k, w in sd.items():
|
||||
for r, rr in attn_key_replace.items():
|
||||
k = k.replace(r, rr)
|
||||
converted_sd[k] = w
|
||||
sd = converted_sd
|
||||
|
||||
num_layers = sum(1 for k in sd if k.startswith("proj_out.") and k.endswith(".weight"))
|
||||
conv_out_dim = sd["controlnet_patch_embedding.weight"].shape[0]
|
||||
if "proj_in.weight" in sd:
|
||||
dim = sd["proj_in.weight"].shape[0]
|
||||
else:
|
||||
dim = conv_out_dim
|
||||
model = comfy.ldm.wan.uni3c.WanUni3CControlnet(
|
||||
in_channels=sd["controlnet_patch_embedding.weight"].shape[1],
|
||||
conv_out_dim=conv_out_dim,
|
||||
dim=dim,
|
||||
ffn_dim=sd["controlnet_blocks.0.ffn.0.bias"].shape[0],
|
||||
num_layers=num_layers,
|
||||
time_embed_dim=sd["controlnet_blocks.0.norm1.linear.weight"].shape[1],
|
||||
out_proj_dim=sd["proj_out.0.weight"].shape[0],
|
||||
add_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[1],
|
||||
mid_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[0],
|
||||
device=comfy.model_management.unet_offload_device(),
|
||||
dtype=dtype,
|
||||
operations=comfy.ops.manual_cast)
|
||||
elif "audio_proj.proj1.weight" in sd:
|
||||
model = MultiTalkModelPatch(
|
||||
audio_window=5, context_tokens=32, vae_scale=4,
|
||||
|
|
@ -296,6 +331,50 @@ class ModelPatchLoader:
|
|||
return (model_patcher,)
|
||||
|
||||
|
||||
class AnimaLLLiteApply:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"model": ("MODEL",),
|
||||
"model_patch": ("MODEL_PATCH",),
|
||||
"image": ("IMAGE",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
},
|
||||
"optional": {"mask": ("MASK",),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "apply_patch"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
CATEGORY = "model_patches/anima"
|
||||
|
||||
def apply_patch(self, model, model_patch, image, strength, start_percent, end_percent, mask=None):
|
||||
image = image[..., :3]
|
||||
|
||||
if model_patch.model.cond_in_channels == 4 and mask is None:
|
||||
mask = torch.zeros_like(image[..., 0])
|
||||
elif model_patch.model.cond_in_channels != 4:
|
||||
mask = None
|
||||
|
||||
model_sampling = model.get_model_object("model_sampling")
|
||||
sigma_start = float(model_sampling.percent_to_sigma(start_percent))
|
||||
sigma_end = float(model_sampling.percent_to_sigma(end_percent))
|
||||
patch = comfy.ldm.anima.lllite.AnimaLLLitePatch(model_patch, image, mask, strength, sigma_start, sigma_end)
|
||||
model_patched = model.clone()
|
||||
model_patched.set_model_post_input_patch(patch)
|
||||
model_patched.set_model_attn1_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch(
|
||||
patch,
|
||||
{"q": "self_attn_q_proj", "k": "self_attn_k_proj", "v": "self_attn_v_proj"},
|
||||
))
|
||||
model_patched.set_model_attn2_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch(
|
||||
patch,
|
||||
{"q": "cross_attn_q_proj"},
|
||||
))
|
||||
model_patched.set_model_patch(comfy.ldm.anima.lllite.AnimaLLLiteMLPPatch(patch), "mlp_patch")
|
||||
return (model_patched,)
|
||||
|
||||
|
||||
class DiffSynthCnetPatch:
|
||||
def __init__(self, model_patch, vae, image, strength, mask=None):
|
||||
self.model_patch = model_patch
|
||||
|
|
@ -514,6 +593,150 @@ class ZImageFunControlnet(QwenImageDiffsynthControlnet):
|
|||
|
||||
CATEGORY = "model/patch/z-image"
|
||||
|
||||
class WanUni3CCnetPatch:
|
||||
def __init__(self, model_patch, render_video, vae, latent_format, strength, sigma_start, sigma_end):
|
||||
self.model_patch = model_patch
|
||||
self.render_video = render_video
|
||||
self.vae = vae
|
||||
self.latent_format = latent_format
|
||||
self.strength = strength
|
||||
self.sigma_start = sigma_start
|
||||
self.sigma_end = sigma_end
|
||||
self.prepared_render = None
|
||||
self.temp_data = None
|
||||
|
||||
def encode_render_video(self, target_latent_shape):
|
||||
t_len, h_len, w_len = target_latent_shape
|
||||
temporal_compression = self.vae.temporal_compression_decode() or 1
|
||||
spatial_compression = self.vae.spacial_compression_encode()
|
||||
target_frames = (t_len - 1) * temporal_compression + 1
|
||||
target_height = h_len * spatial_compression
|
||||
target_width = w_len * spatial_compression
|
||||
|
||||
frames = self.render_video
|
||||
if frames.shape[0] > target_frames:
|
||||
frames = frames[:target_frames]
|
||||
elif frames.shape[0] < target_frames:
|
||||
last_frame = frames[-1:].expand(target_frames - frames.shape[0], -1, -1, -1)
|
||||
frames = torch.cat([frames, last_frame], dim=0)
|
||||
|
||||
if frames.shape[1] != target_height or frames.shape[2] != target_width:
|
||||
frames = comfy.utils.common_upscale(frames.movedim(-1, 1), target_width, target_height, "bilinear", "center").movedim(1, -1)
|
||||
|
||||
loaded_models = comfy.model_management.loaded_models(only_currently_used=True)
|
||||
render_latent = self.vae.encode(frames)
|
||||
comfy.model_management.load_models_gpu(loaded_models)
|
||||
return self.latent_format.process_in(render_latent)
|
||||
|
||||
def build_controlnet_input(self, x, dtype, samples_per_cond):
|
||||
# first 20 channels of the model input: noise latent + I2V mask (zero padded for T2V)
|
||||
hidden = x[:samples_per_cond, :20].to(dtype)
|
||||
if hidden.shape[1] < 20:
|
||||
pad_shape = list(hidden.shape)
|
||||
pad_shape[1] = 20 - hidden.shape[1]
|
||||
hidden = torch.cat([hidden, torch.zeros(pad_shape, dtype=hidden.dtype, device=hidden.device)], dim=1)
|
||||
|
||||
render = self.prepared_render
|
||||
if render is None or render.shape[2:] != hidden.shape[2:]:
|
||||
render = self.encode_render_video(hidden.shape[2:])
|
||||
render = render.to(device=hidden.device, dtype=dtype)
|
||||
self.prepared_render = render
|
||||
if render.shape[0] != hidden.shape[0]:
|
||||
render = render.expand(hidden.shape[0], -1, -1, -1, -1)
|
||||
return torch.cat([hidden, render], dim=1)
|
||||
|
||||
def __call__(self, kwargs):
|
||||
img = kwargs.get("img")
|
||||
block_index = kwargs.get("block_index")
|
||||
transformer_options = kwargs.get("transformer_options", {})
|
||||
|
||||
if block_index == 0:
|
||||
self.temp_data = None
|
||||
active = True
|
||||
sigmas = transformer_options.get("sigmas", None)
|
||||
if sigmas is not None:
|
||||
sigma = sigmas[0].item()
|
||||
if sigma > self.sigma_start or sigma < self.sigma_end:
|
||||
active = False
|
||||
if active:
|
||||
x = kwargs.get("x")
|
||||
# cond and uncond chunks share latents, so we can reuse residuals
|
||||
num_conds = len(transformer_options.get("cond_or_uncond", [0]))
|
||||
samples_per_cond = x.shape[0]
|
||||
if num_conds > 0 and x.shape[0] % num_conds == 0:
|
||||
samples_per_cond = x.shape[0] // num_conds
|
||||
temb = kwargs.get("vec")[:samples_per_cond]
|
||||
if temb.ndim == 3:
|
||||
temb = temb[:, 0]
|
||||
model = self.model_patch.model
|
||||
controlnet_input = self.build_controlnet_input(x, img.dtype, samples_per_cond)
|
||||
hidden, freqs = model.process_input(controlnet_input)
|
||||
self.temp_data = (hidden, temb.to(img.dtype), freqs)
|
||||
|
||||
num_layers = self.model_patch.model.num_layers
|
||||
if self.temp_data is not None and block_index < num_layers:
|
||||
hidden, temb, freqs = self.temp_data
|
||||
hidden, residual = self.model_patch.model.forward_block(block_index, hidden, temb, freqs)
|
||||
residual = residual.to(img.dtype) * self.strength
|
||||
if residual.shape[0] != img.shape[0]:
|
||||
residual = residual.repeat(img.shape[0] // residual.shape[0], 1, 1)
|
||||
img_offset = kwargs.get("img_offset", 0)
|
||||
img[:, img_offset:img_offset + residual.shape[1]] += residual
|
||||
if block_index >= num_layers - 1:
|
||||
self.temp_data = None
|
||||
else:
|
||||
self.temp_data = (hidden, temb, freqs)
|
||||
|
||||
return kwargs
|
||||
|
||||
def to(self, device_or_dtype):
|
||||
if isinstance(device_or_dtype, torch.device):
|
||||
if self.prepared_render is not None:
|
||||
self.prepared_render = self.prepared_render.to(device_or_dtype)
|
||||
self.temp_data = None
|
||||
return self
|
||||
|
||||
def models(self):
|
||||
return [self.model_patch]
|
||||
|
||||
|
||||
class WanUni3CControlnetApply:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "model": ("MODEL",),
|
||||
"model_patch": ("MODEL_PATCH",),
|
||||
"vae": ("VAE",),
|
||||
"render_video": ("IMAGE", {"tooltip": "The guidance video rendered from the camera trajectory, most commonly warped point cloud renders of the input image."}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "apply_patch"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
CATEGORY = "model/patch/wan"
|
||||
|
||||
def apply_patch(self, model, model_patch, vae, render_video, strength, start_percent, end_percent):
|
||||
if not isinstance(model_patch.model, comfy.ldm.wan.uni3c.WanUni3CControlnet):
|
||||
raise ValueError("The connected model patch is not a Uni3C ControlNet.")
|
||||
cnet_dim = model_patch.model.controlnet_blocks[0].norm1.linear.in_features
|
||||
model_dim = getattr(model.get_model_object("diffusion_model"), "dim", None)
|
||||
if model_dim is None:
|
||||
raise ValueError("The Uni3C ControlNet only works with Wan models.")
|
||||
if model_dim != cnet_dim:
|
||||
raise ValueError("This Uni3C ControlNet expects a Wan model with dim {}, the loaded model has dim {}.".format(cnet_dim, model_dim))
|
||||
|
||||
model_patched = model.clone()
|
||||
model_sampling = model.get_model_object("model_sampling")
|
||||
sigma_start = model_sampling.percent_to_sigma(start_percent)
|
||||
sigma_end = model_sampling.percent_to_sigma(end_percent)
|
||||
latent_format = model.get_model_object("latent_format")
|
||||
patch = WanUni3CCnetPatch(model_patch, render_video[:, :, :, :3], vae, latent_format, strength, sigma_start, sigma_end)
|
||||
model_patched.set_model_double_block_patch(patch)
|
||||
return (model_patched,)
|
||||
|
||||
|
||||
class UsoStyleProjectorPatch:
|
||||
def __init__(self, model_patch, encoded_image):
|
||||
self.model_patch = model_patch
|
||||
|
|
@ -672,14 +895,18 @@ NODE_CLASS_MAPPINGS = {
|
|||
"ModelPatchLoader": ModelPatchLoader,
|
||||
"QwenImageDiffsynthControlnet": QwenImageDiffsynthControlnet,
|
||||
"ZImageFunControlnet": ZImageFunControlnet,
|
||||
"WanUni3CControlnetApply": WanUni3CControlnetApply,
|
||||
"USOStyleReference": USOStyleReference,
|
||||
"SUPIRApply": SUPIRApply,
|
||||
"AnimaLLLiteApply": AnimaLLLiteApply,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ModelPatchLoader": "Load Model Patch",
|
||||
"QwenImageDiffsynthControlnet": "Apply Qwen Image DiffSynth ControlNet",
|
||||
"ZImageFunControlnet": "Apply Z-Image Fun ControlNet",
|
||||
"WanUni3CControlnetApply": "Apply Wan Uni3C ControlNet",
|
||||
"USOStyleReference": "Apply USO Style Reference",
|
||||
"SUPIRApply": "Apply SUPIR Patch",
|
||||
"AnimaLLLiteApply": "Apply Anima LLLite",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -920,10 +920,11 @@ def _run_training_loop(
|
|||
"""
|
||||
sigmas = torch.tensor(range(num_images))
|
||||
noise = comfy_extras.nodes_custom_sampler.Noise_RandomNoise(seed)
|
||||
ndim = latents[0].ndim
|
||||
|
||||
if bucket_mode:
|
||||
# Use first bucket's first latent as dummy for guider
|
||||
dummy_latent = latents[0][:1].repeat(num_images, 1, 1, 1)
|
||||
dummy_latent = latents[0][:1].repeat(num_images, *[1]*(ndim-1))
|
||||
guider.sample(
|
||||
noise.generate_noise({"samples": dummy_latent}),
|
||||
dummy_latent,
|
||||
|
|
@ -933,7 +934,7 @@ def _run_training_loop(
|
|||
)
|
||||
elif multi_res:
|
||||
# use first latent as dummy latent if multi_res
|
||||
latents = latents[0].repeat(num_images, 1, 1, 1)
|
||||
latents = latents[0].repeat(num_images, *[1]*(ndim-1))
|
||||
guider.sample(
|
||||
noise.generate_noise({"samples": latents}),
|
||||
latents,
|
||||
|
|
|
|||
51
openapi.yaml
51
openapi.yaml
|
|
@ -530,6 +530,10 @@ components:
|
|||
description: Job creation timestamp (Unix timestamp in milliseconds)
|
||||
format: int64
|
||||
type: integer
|
||||
execution_end_time:
|
||||
description: Workflow execution completion timestamp (Unix milliseconds, only present for terminal states)
|
||||
format: int64
|
||||
type: integer
|
||||
execution_error:
|
||||
allOf:
|
||||
- $ref: '#/components/schemas/ExecutionError'
|
||||
|
|
@ -538,6 +542,10 @@ components:
|
|||
additionalProperties: true
|
||||
description: Node-level execution metadata (only for terminal states)
|
||||
type: object
|
||||
execution_start_time:
|
||||
description: Workflow execution start timestamp (Unix milliseconds, only present once execution has started)
|
||||
format: int64
|
||||
type: integer
|
||||
execution_status:
|
||||
additionalProperties: true
|
||||
description: ComfyUI execution status and timeline (only for terminal states)
|
||||
|
|
@ -570,6 +578,12 @@ components:
|
|||
description: Last update timestamp (Unix timestamp in milliseconds)
|
||||
format: int64
|
||||
type: integer
|
||||
user_id:
|
||||
description: |
|
||||
ID of the user that owns this job (see the `workspace_id`
|
||||
description above for why this is always the caller's own id
|
||||
on a successful response).
|
||||
type: string
|
||||
workflow:
|
||||
additionalProperties: true
|
||||
description: |
|
||||
|
|
@ -583,6 +597,18 @@ components:
|
|||
workflow_id:
|
||||
description: UUID identifying the workflow graph definition
|
||||
type: string
|
||||
workspace_id:
|
||||
description: |
|
||||
ID of the workspace that owns this job. A successful (200)
|
||||
response from this operation is only ever returned for the
|
||||
caller's own job (see this operation's ownership-scoped
|
||||
query), so this is always the caller's own workspace —
|
||||
consumers that also need to correlate this job to its
|
||||
live-progress broadcast channel (workspace+user scoped; see
|
||||
the internal common/gateways/broadcast package) can use this
|
||||
value directly rather than resolving their own identity a
|
||||
second way.
|
||||
type: string
|
||||
required:
|
||||
- id
|
||||
- status
|
||||
|
|
@ -1565,7 +1591,13 @@ paths:
|
|||
schema:
|
||||
default: true
|
||||
type: boolean
|
||||
- description: Filter assets by exact content hash.
|
||||
- description: |
|
||||
Filter assets by content hash, in the canonical `blake3:<hex>`
|
||||
form. Matches regardless of which of this asset store's two
|
||||
internal hash storage formats the matching row was written
|
||||
under (the canonical form used by from-hash-created references,
|
||||
or the raw `<hex>.<ext>`/bare `<hex>` storage key used by direct
|
||||
uploads) — both represent the same content hash.
|
||||
in: query
|
||||
name: hash
|
||||
schema:
|
||||
|
|
@ -2464,6 +2496,23 @@ paths:
|
|||
schema:
|
||||
additionalProperties: true
|
||||
properties:
|
||||
free_tier_balance:
|
||||
description: Free-tier job allowance for an authenticated non-paid (FREE-tier) user in the rollout. Absent for paid users and unauthenticated requests. Synthesized from config before a grant row exists so a brand-new user still sees their full allowance.
|
||||
properties:
|
||||
allowance:
|
||||
description: Total free jobs granted for the current period
|
||||
type: integer
|
||||
remaining:
|
||||
description: Free jobs remaining (allowance - used, floored at 0)
|
||||
type: integer
|
||||
used:
|
||||
description: Free jobs consumed so far
|
||||
type: integer
|
||||
required:
|
||||
- allowance
|
||||
- used
|
||||
- remaining
|
||||
type: object
|
||||
max_upload_size:
|
||||
description: Maximum upload size in bytes
|
||||
type: integer
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
comfyui-frontend-package==1.45.21
|
||||
comfyui-workflow-templates==0.11.9
|
||||
comfyui-frontend-package==1.47.10
|
||||
comfyui-workflow-templates==0.11.15
|
||||
comfyui-embedded-docs==0.5.8
|
||||
torch
|
||||
torchsde
|
||||
|
|
@ -22,7 +22,7 @@ alembic
|
|||
SQLAlchemy>=2.0.0
|
||||
filelock
|
||||
av>=16.0.0
|
||||
comfy-kitchen==0.2.21
|
||||
comfy-kitchen==0.2.22
|
||||
comfy-aimdo==0.4.10
|
||||
requests
|
||||
simpleeval>=1.0.0
|
||||
|
|
|
|||
Loading…
Reference in New Issue