mirror of
https://github.com/comfyanonymous/ComfyUI.git
synced 2026-07-22 15:59:05 +08:00
Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
75a51b6133 | ||
|
|
f024492f81 | ||
|
|
fa585a8660 | ||
|
|
ee7e1cbf3d | ||
|
|
8c70c85f53 | ||
|
|
4cc4d944e7 | ||
|
|
a1aaa1825d | ||
|
|
655fec886e |
@@ -32,11 +32,9 @@ jobs:
|
||||
PR_NUMBER: ${{ github.event.pull_request.number || github.event.issue.number }}
|
||||
PR_AUTHOR: ${{ github.event.pull_request.user.login || github.event.issue.user.login }}
|
||||
BASE_ALLOWLIST: action@github.com,actions-user,ampagent,claude,comfy-pr-bot,GitHub Action,github-actions,github-actions[bot],Glary Bot,Glary-Bot,*[bot]
|
||||
# For each commit emit the GitHub login when the author/committer email resolves to a GitHub account
|
||||
# otherwise fall back to the raw git name.
|
||||
run: |
|
||||
others=$(gh api "repos/${{ github.repository }}/pulls/${PR_NUMBER}/commits" --paginate \
|
||||
--jq '.[] | (.author.login // .commit.author.name // empty), (.committer.login // .commit.committer.name // empty)' \
|
||||
--jq '.[] | (.author.login // empty), (.committer.login // empty)' \
|
||||
| sort -u | grep -vix "${PR_AUTHOR}" | paste -sd, -)
|
||||
if [ -n "$others" ]; then
|
||||
echo "allowlist=${BASE_ALLOWLIST},${others}" >> "$GITHUB_OUTPUT"
|
||||
@@ -45,7 +43,7 @@ jobs:
|
||||
fi
|
||||
|
||||
- name: CLA Assistant
|
||||
# Run on PR events, on "recheck" comment, or when someone posts the signing phrase.
|
||||
# Run on PR events, on "recheck" comment, or when someone posts the exact signing phrase.
|
||||
# IMPORTANT: this phrase must match `custom-pr-sign-comment` below.
|
||||
if: >
|
||||
github.event_name == 'pull_request_target' ||
|
||||
|
||||
@@ -19,9 +19,6 @@
|
||||
better to remove a broken feature path than keep a complicated partial fix.
|
||||
- Preserve existing APIs, node names, model-loading behavior, file layout, and
|
||||
workflow compatibility unless the change is explicitly about replacing them.
|
||||
- When compatibility is explicitly out of scope, remove compatibility-only
|
||||
aliases, duplicate nodes, legacy entry points, and preset wrappers instead of
|
||||
retaining parallel ways to perform the same operation.
|
||||
- Code must look hand-written for this repository. Changes that read like
|
||||
generic AI-generated code will be rejected automatically: unnecessary helper
|
||||
layers, vague names, boilerplate comments, defensive branches without a real
|
||||
@@ -99,13 +96,6 @@
|
||||
unless they are read by current code and change current behavior. Remove
|
||||
pass-through or stored-but-unused values instead of preserving upstream or
|
||||
deprecated API baggage.
|
||||
- Do not add a model-specific option to a shared helper when only one caller
|
||||
needs it. Keep one-off behavior at the model integration boundary, or extend
|
||||
the shared helper only when the option is a coherent reusable capability.
|
||||
- Implementations of shared model interfaces should accept the standard caller
|
||||
contract without model-specific rejection branches for optional capabilities
|
||||
they do not consume. Let supported behavior be determined by implementation
|
||||
paths that actually use those inputs.
|
||||
- If an implementation needs auxiliary values for its own workflow, expose them
|
||||
through a private helper or a clearly named implementation-specific method
|
||||
instead of overloading the public method's return contract.
|
||||
@@ -164,10 +154,6 @@
|
||||
`comfy-kitchen` helpers where they already solve the problem.
|
||||
- Use optimized comfy-kitchen ops in places where they improve performance
|
||||
without changing the expected dtype, device, memory, or interface behavior.
|
||||
- Prefer ComfyUI's shared optimized kernels and backend dispatchers over
|
||||
handwritten implementations of the same operation. Remove duplicate local
|
||||
kernels and adapt inputs to the shared operation's documented layout while
|
||||
preserving the model's original math and output contract.
|
||||
- All models should use the optimized attention function selected by ComfyUI.
|
||||
Treat optimized backend functions, dispatch helpers, and capability-selected
|
||||
callables as opaque. Higher-level code must not inspect function identity,
|
||||
@@ -190,12 +176,6 @@
|
||||
- Model detection code that inspects linear weight shapes should only use the
|
||||
first dimension. The second dimension may be half the original size for
|
||||
NVFP4 or other 4-bit quantized models.
|
||||
- A model-detection signature must guard every state-dict key it dereferences.
|
||||
Do not partially match a format and then raise an incidental `KeyError` while
|
||||
extracting its configuration.
|
||||
- Order model-detection checks from established or more-specific signatures to
|
||||
newer or broader signatures. Put a broad new detector near the generic
|
||||
fallback when giving it higher precedence could steal another model family.
|
||||
- Avoid adding `einops` usage in core inference code. Use native torch tensor
|
||||
ops such as `reshape`, `view`, `permute`, `transpose`, `flatten`, `unflatten`,
|
||||
`unsqueeze`, and `squeeze` instead.
|
||||
@@ -212,23 +192,11 @@
|
||||
methods for scalar or structural calculations.
|
||||
- Avoid unnecessary casts and transfers. Preserve the intended compute dtype,
|
||||
storage dtype, bias dtype, and original tensor shape metadata.
|
||||
- Do not cast the result of an optimized backend operation back to its input
|
||||
dtype unless that backend's documented result contract requires normalization.
|
||||
In particular, trust the selected optimized-attention implementation to honor
|
||||
its dtype contract.
|
||||
- Keep model-native latent layout handling inside the model or latent-format
|
||||
owner, not in helper nodes. Do not collapse, expand, pack, or unpack latent
|
||||
dimensions in nodes or other caller-side adapters just to satisfy a model
|
||||
forward; the model path should consume and return the native latent shape for
|
||||
that model family.
|
||||
- DiT models should accept latent dimensions that are not exact patch-size
|
||||
multiples. Use `comfy.ldm.common_dit.pad_to_patch_size` on every patchified
|
||||
target or reference input, then crop only the target output back to its
|
||||
original dimensions.
|
||||
- Avoid defensive shape and configuration checks that merely replace the clear
|
||||
failure from the tensor operation immediately below them. Add explicit
|
||||
validation only when it provides materially better context at a real boundary
|
||||
or prevents silent incorrect output.
|
||||
- Assume inputs to the main model forward are already in the compute dtype by
|
||||
default, except integer inputs such as some model timestep tensors. Do not add
|
||||
defensive or convenience casts in model code; it is better for invalid dtype
|
||||
@@ -292,15 +260,6 @@
|
||||
- Model implementations should add the minimal number of ComfyUI nodes required
|
||||
to run the model. Reuse existing nodes as much as possible; adapting the model
|
||||
to work with existing nodes is strongly preferred over creating new nodes.
|
||||
- Use `io.Autogrow` for a variable number of repeated inputs instead of a fixed
|
||||
series of numbered optional sockets. Set its minimum to zero when the model
|
||||
has a valid no-item path, and cap it only when the model has a real limit.
|
||||
- Mark inputs optional when execution has a valid path that does not read them.
|
||||
If one optional input is needed only to process another optional input, do not
|
||||
force users on the path that supplies neither to connect it.
|
||||
- Conditioning nodes should normally output conditioning only. Do not expose
|
||||
input or intermediate images as convenience outputs for downstream sizing or
|
||||
routing; use the existing image path or a dedicated image operation instead.
|
||||
- Nodes should output only values they own. Do not add pass-through outputs for
|
||||
workflow convenience unless the node is explicitly an output node. Existing
|
||||
models, latents, conditioning, or other inputs should flow directly to the
|
||||
|
||||
@@ -1,90 +0,0 @@
|
||||
import logging
|
||||
import os
|
||||
import uuid
|
||||
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
import folder_paths
|
||||
from app.assets.database.queries.asset_reference import (
|
||||
get_reference_by_file_path,
|
||||
get_reference_by_id,
|
||||
set_reference_preview,
|
||||
)
|
||||
from app.assets.services.ingest import register_file_in_place
|
||||
from app.database.db import create_session
|
||||
|
||||
PREVIEW_TAG = "preview"
|
||||
|
||||
|
||||
def get_or_create_preview_file(
|
||||
source_abs_path: str, max_size: int, quality: int
|
||||
) -> str:
|
||||
source_hash, source_ref_id = _ensure_source_asset(source_abs_path)
|
||||
hash_hex = source_hash.split(":")[-1]
|
||||
preview_dir = folder_paths.get_system_user_directory("preview_cache")
|
||||
preview_path = os.path.join(
|
||||
preview_dir, f"{hash_hex}_{max_size}_q{quality}.webp"
|
||||
)
|
||||
if os.path.isfile(preview_path):
|
||||
return preview_path
|
||||
|
||||
os.makedirs(preview_dir, exist_ok=True)
|
||||
tmp_path = f"{preview_path}.{uuid.uuid4().hex}.tmp"
|
||||
try:
|
||||
with Image.open(source_abs_path) as img:
|
||||
preview_img = ImageOps.exif_transpose(img)
|
||||
if max(preview_img.size) > max_size:
|
||||
preview_img = ImageOps.contain(
|
||||
preview_img, (max_size, max_size), Image.Resampling.LANCZOS
|
||||
)
|
||||
preview_img.save(tmp_path, format="webp", quality=quality)
|
||||
os.replace(tmp_path, preview_path)
|
||||
except Exception:
|
||||
if os.path.exists(tmp_path):
|
||||
try:
|
||||
os.remove(tmp_path)
|
||||
except OSError:
|
||||
pass
|
||||
raise
|
||||
|
||||
_register_preview_asset(preview_path, source_ref_id)
|
||||
return preview_path
|
||||
|
||||
|
||||
def _ensure_source_asset(abs_path: str) -> tuple[str, str]:
|
||||
abs_path = os.path.abspath(abs_path)
|
||||
mtime_ns = os.stat(abs_path).st_mtime_ns
|
||||
|
||||
with create_session() as session:
|
||||
ref = get_reference_by_file_path(session, abs_path)
|
||||
if (
|
||||
ref is not None
|
||||
and ref.mtime_ns == mtime_ns
|
||||
and ref.asset is not None
|
||||
and ref.asset.hash
|
||||
):
|
||||
return ref.asset.hash, ref.id
|
||||
|
||||
result = register_file_in_place(
|
||||
abs_path=abs_path, name=os.path.basename(abs_path), tags=[]
|
||||
)
|
||||
if not result.asset.hash:
|
||||
raise RuntimeError(f"asset registration produced no hash for {abs_path}")
|
||||
return result.asset.hash, result.ref.id
|
||||
|
||||
|
||||
def _register_preview_asset(preview_path: str, source_ref_id: str) -> None:
|
||||
try:
|
||||
result = register_file_in_place(
|
||||
abs_path=preview_path,
|
||||
name=os.path.basename(preview_path),
|
||||
tags=[PREVIEW_TAG],
|
||||
mime_type="image/webp",
|
||||
)
|
||||
with create_session() as session:
|
||||
source_ref = get_reference_by_id(session, source_ref_id)
|
||||
if source_ref is not None and source_ref.preview_id is None:
|
||||
set_reference_preview(session, source_ref_id, result.ref.id)
|
||||
session.commit()
|
||||
except Exception:
|
||||
logging.warning("Failed to register preview image as asset", exc_info=True)
|
||||
@@ -35,11 +35,7 @@ class ModelFileManager:
|
||||
for folder in model_types:
|
||||
if folder in folder_black_list:
|
||||
continue
|
||||
output_folders.append({
|
||||
"name": folder,
|
||||
"folders": folder_paths.get_folder_paths(folder),
|
||||
"extensions": sorted(folder_paths.folder_names_and_paths[folder][1]),
|
||||
})
|
||||
output_folders.append({"name": folder, "folders": folder_paths.get_folder_paths(folder)})
|
||||
return web.json_response(output_folders)
|
||||
|
||||
# NOTE: This is an experiment to replace `/models/{folder}`
|
||||
|
||||
@@ -92,7 +92,6 @@ parser.add_argument("--directml", type=int, nargs="?", metavar="DIRECTML_DEVICE"
|
||||
parser.add_argument("--oneapi-device-selector", type=str, default=None, metavar="SELECTOR_STRING", help="Sets the oneAPI device(s) this instance will use.")
|
||||
parser.add_argument("--supports-fp8-compute", action="store_true", help="ComfyUI will act like if the device supports fp8 compute.")
|
||||
parser.add_argument("--enable-triton-backend", action="store_true", help="ComfyUI will enable the use of Triton backend in comfy-kitchen. Is disabled at launch by default.")
|
||||
parser.add_argument("--disable-triton-backend", action="store_true", help="Force-disable the comfy-kitchen Triton backend, overriding the automatic ROCm/AMD default and --enable-triton-backend.")
|
||||
|
||||
class LatentPreviewMethod(enum.Enum):
|
||||
NoPreviews = "none"
|
||||
|
||||
@@ -779,10 +779,6 @@ class ACEAudio(LatentFormat):
|
||||
latent_channels = 8
|
||||
latent_dimensions = 2
|
||||
|
||||
class SeedVR2(LatentFormat):
|
||||
latent_channels = 16
|
||||
latent_dimensions = 3
|
||||
|
||||
class ACEAudio15(LatentFormat):
|
||||
latent_channels = 64
|
||||
latent_dimensions = 1
|
||||
|
||||
@@ -15,24 +15,24 @@ def make_two_pass_attention(ar_len: int, transformer_options=None):
|
||||
The AR pass goes through SDPA directand bypasses wrappers, it is only ~1% of T at typical edit sizes.
|
||||
"""
|
||||
|
||||
def two_pass_attention(q, k, v, heads, enable_gqa=False, **kwargs):
|
||||
def two_pass_attention(q, k, v, heads, **kwargs):
|
||||
B, H, T, D = q.shape
|
||||
|
||||
if T < k.shape[2]: # KV-cache hot path: Q is shorter than K/V (cached AR prefix is in K/V only), all fresh Q positions are in the gen region, single full-attention call
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa)
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options)
|
||||
elif ar_len >= T:
|
||||
out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa)
|
||||
out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True)
|
||||
elif ar_len <= 0:
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa)
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options)
|
||||
else:
|
||||
out_ar = comfy.ops.scaled_dot_product_attention(
|
||||
q[:, :, :ar_len], k[:, :, :ar_len], v[:, :, :ar_len],
|
||||
attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa,
|
||||
attn_mask=None, dropout_p=0.0, is_causal=True,
|
||||
)
|
||||
out_gen = optimized_attention(
|
||||
q[:, :, ar_len:], k, v, heads,
|
||||
mask=None, skip_reshape=True, skip_output_reshape=True,
|
||||
transformer_options=transformer_options, enable_gqa=enable_gqa,
|
||||
transformer_options=transformer_options,
|
||||
)
|
||||
out = torch.cat([out_ar, out_gen], dim=2)
|
||||
|
||||
|
||||
@@ -1,445 +0,0 @@
|
||||
# https://github.com/jdopensource/JoyAI-Image-Edit (Apache 2.0)
|
||||
import math
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import comfy_kitchen
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import comfy.ldm.common_dit
|
||||
import comfy.ops
|
||||
import comfy.patcher_extension
|
||||
from comfy.ldm.lightricks.model import GELU_approx, PixArtAlphaTextProjection, TimestepEmbedding, Timesteps
|
||||
from comfy.ldm.modules.attention import optimized_attention
|
||||
|
||||
|
||||
class JoyImageModulate(nn.Module):
|
||||
def __init__(self, hidden_size: int, factor: int, dtype=None, device=None):
|
||||
super().__init__()
|
||||
self.factor = factor
|
||||
self.modulate_table = nn.Parameter(
|
||||
torch.empty(1, factor, hidden_size, dtype=dtype, device=device)
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> list:
|
||||
if x.ndim != 3:
|
||||
x = x.unsqueeze(1)
|
||||
table = comfy.ops.cast_to_input(self.modulate_table, x)
|
||||
return [o.squeeze(1) for o in (table + x).chunk(self.factor, dim=1)]
|
||||
|
||||
|
||||
class JoyImageFeedForward(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
inner_dim: int,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.net = nn.ModuleList([
|
||||
GELU_approx(dim, inner_dim, dtype=dtype, device=device, operations=operations),
|
||||
nn.Identity(),
|
||||
operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device),
|
||||
])
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
for module in self.net:
|
||||
x = module(x)
|
||||
return x
|
||||
|
||||
|
||||
class JoyImageAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
eps: float = 1e-6,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.img_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device)
|
||||
self.img_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
|
||||
self.img_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
|
||||
self.img_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device)
|
||||
|
||||
self.txt_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device)
|
||||
self.txt_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
|
||||
self.txt_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
|
||||
self.txt_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
txt: torch.Tensor,
|
||||
image_rotary_emb: torch.Tensor,
|
||||
transformer_options=None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
heads = self.num_attention_heads
|
||||
|
||||
img_q, img_k, img_v = self.img_attn_qkv(img).chunk(3, dim=-1)
|
||||
txt_q, txt_k, txt_v = self.txt_attn_qkv(txt).chunk(3, dim=-1)
|
||||
|
||||
img_q = img_q.unflatten(-1, (heads, -1))
|
||||
img_k = img_k.unflatten(-1, (heads, -1))
|
||||
img_v = img_v.unflatten(-1, (heads, -1))
|
||||
txt_q = txt_q.unflatten(-1, (heads, -1))
|
||||
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)
|
||||
|
||||
joint_q = torch.cat([img_q, txt_q], dim=1)
|
||||
joint_k = torch.cat([img_k, txt_k], dim=1)
|
||||
joint_v = torch.cat([img_v, txt_v], dim=1)
|
||||
|
||||
joint_q = joint_q.flatten(2, 3)
|
||||
joint_k = joint_k.flatten(2, 3)
|
||||
joint_v = joint_v.flatten(2, 3)
|
||||
|
||||
joint_out = optimized_attention(joint_q, joint_k, joint_v, heads=heads, transformer_options=transformer_options)
|
||||
|
||||
seq_img = img.shape[1]
|
||||
img_out = joint_out[:, :seq_img, :]
|
||||
txt_out = joint_out[:, seq_img:, :]
|
||||
|
||||
img_out = self.img_attn_proj(img_out)
|
||||
txt_out = self.txt_attn_proj(txt_out)
|
||||
return img_out, txt_out
|
||||
|
||||
|
||||
class JoyImageTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
eps: float = 1e-6,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
mlp_hidden_dim = int(dim * mlp_width_ratio)
|
||||
|
||||
self.img_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device)
|
||||
self.img_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
|
||||
self.img_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
|
||||
self.img_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations)
|
||||
|
||||
self.txt_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device)
|
||||
self.txt_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
|
||||
self.txt_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
|
||||
self.txt_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations)
|
||||
|
||||
self.attn = JoyImageAttention(
|
||||
dim=dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
eps=eps,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: torch.Tensor,
|
||||
transformer_options=None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
(
|
||||
img_mod1_shift,
|
||||
img_mod1_scale,
|
||||
img_mod1_gate,
|
||||
img_mod2_shift,
|
||||
img_mod2_scale,
|
||||
img_mod2_gate,
|
||||
) = self.img_mod(temb)
|
||||
(
|
||||
txt_mod1_shift,
|
||||
txt_mod1_scale,
|
||||
txt_mod1_gate,
|
||||
txt_mod2_shift,
|
||||
txt_mod2_scale,
|
||||
txt_mod2_gate,
|
||||
) = self.txt_mod(temb)
|
||||
|
||||
img_normed = self.img_norm1(hidden_states)
|
||||
txt_normed = self.txt_norm1(encoder_hidden_states)
|
||||
img_modulated = img_normed * (1 + img_mod1_scale.unsqueeze(1)) + img_mod1_shift.unsqueeze(1)
|
||||
txt_modulated = txt_normed * (1 + txt_mod1_scale.unsqueeze(1)) + txt_mod1_shift.unsqueeze(1)
|
||||
|
||||
img_attn, txt_attn = self.attn(img_modulated, txt_modulated, image_rotary_emb, transformer_options=transformer_options)
|
||||
|
||||
hidden_states = hidden_states + img_attn * img_mod1_gate.unsqueeze(1)
|
||||
encoder_hidden_states = encoder_hidden_states + txt_attn * txt_mod1_gate.unsqueeze(1)
|
||||
|
||||
img_ffn_normed = self.img_norm2(hidden_states)
|
||||
txt_ffn_normed = self.txt_norm2(encoder_hidden_states)
|
||||
img_ffn_input = img_ffn_normed * (1 + img_mod2_scale.unsqueeze(1)) + img_mod2_shift.unsqueeze(1)
|
||||
txt_ffn_input = txt_ffn_normed * (1 + txt_mod2_scale.unsqueeze(1)) + txt_mod2_shift.unsqueeze(1)
|
||||
hidden_states = hidden_states + self.img_mlp(img_ffn_input) * img_mod2_gate.unsqueeze(1)
|
||||
encoder_hidden_states = encoder_hidden_states + self.txt_mlp(txt_ffn_input) * txt_mod2_gate.unsqueeze(1)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class JoyImageTimeTextImageEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
time_freq_dim: int,
|
||||
time_proj_dim: int,
|
||||
text_embed_dim: int,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
|
||||
self.time_embedder = TimestepEmbedding(
|
||||
in_channels=time_freq_dim,
|
||||
time_embed_dim=dim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
self.act_fn = nn.SiLU()
|
||||
self.time_proj = operations.Linear(dim, time_proj_dim, bias=True, dtype=dtype, device=device)
|
||||
self.text_embedder = PixArtAlphaTextProjection(
|
||||
text_embed_dim, dim, act_fn="gelu_tanh", dtype=dtype, device=device, operations=operations,
|
||||
)
|
||||
|
||||
def forward(self, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor):
|
||||
timestep = self.timesteps_proj(timestep)
|
||||
temb = self.time_embedder(timestep.to(dtype=encoder_hidden_states.dtype)).type_as(encoder_hidden_states)
|
||||
timestep_proj = self.time_proj(self.act_fn(temb))
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
return temb, timestep_proj, encoder_hidden_states
|
||||
|
||||
|
||||
class JoyImageTransformer3DModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: list = [1, 2, 2],
|
||||
in_channels: int = 16,
|
||||
out_channels: Optional[int] = None,
|
||||
hidden_size: int = 3072,
|
||||
num_attention_heads: int = 24,
|
||||
text_dim: int = 4096,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
num_layers: int = 20,
|
||||
rope_dim_list: list = [16, 56, 56],
|
||||
theta: int = 256,
|
||||
image_model=None,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dtype = dtype
|
||||
self.out_channels = out_channels or in_channels
|
||||
self.patch_size = list(patch_size)
|
||||
self.rope_dim_list = list(rope_dim_list)
|
||||
self.theta = theta
|
||||
|
||||
attention_head_dim = hidden_size // num_attention_heads
|
||||
|
||||
self.img_in = operations.Conv3d(
|
||||
in_channels,
|
||||
hidden_size,
|
||||
kernel_size=tuple(self.patch_size),
|
||||
stride=tuple(self.patch_size),
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
self.condition_embedder = JoyImageTimeTextImageEmbedding(
|
||||
dim=hidden_size,
|
||||
time_freq_dim=256,
|
||||
time_proj_dim=hidden_size * 6,
|
||||
text_embed_dim=text_dim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
|
||||
self.double_blocks = nn.ModuleList([
|
||||
JoyImageTransformerBlock(
|
||||
dim=hidden_size,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
])
|
||||
|
||||
self.norm_out = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
|
||||
self.proj_out = operations.Linear(
|
||||
hidden_size,
|
||||
self.out_channels * math.prod(self.patch_size),
|
||||
bias=True,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def _get_rotary_pos_embed_for_range(
|
||||
self,
|
||||
start: Tuple[int, int, int],
|
||||
stop: Tuple[int, int, int],
|
||||
device=None,
|
||||
) -> torch.Tensor:
|
||||
# 3D RoPE for the patch grid range [start, stop) over (t, h, w). Token order after
|
||||
# reshape(-1) is (t, h, w), matching the img_in Conv3d flatten.
|
||||
rope_dim_list = self.rope_dim_list
|
||||
|
||||
grids = [torch.arange(start[i], stop[i], dtype=torch.float32, device=device) for i in range(3)]
|
||||
mesh = torch.stack(torch.meshgrid(*grids, indexing="ij"), dim=0)
|
||||
|
||||
angles_parts = []
|
||||
for i, dim in enumerate(rope_dim_list):
|
||||
pos = mesh[i].reshape(-1)
|
||||
freqs = 1.0 / (self.theta ** (torch.arange(0, dim, 2, dtype=torch.float32, device=device)[: (dim // 2)] / dim))
|
||||
angles_parts.append(torch.outer(pos, freqs))
|
||||
|
||||
angles = torch.cat(angles_parts, dim=1)
|
||||
cos = angles.cos()
|
||||
sin = angles.sin()
|
||||
return torch.stack((cos, -sin, sin, cos), dim=-1).unflatten(-1, (2, 2))
|
||||
|
||||
def get_rotary_pos_embed_for_components(
|
||||
self,
|
||||
component_sizes,
|
||||
device=None,
|
||||
) -> torch.Tensor:
|
||||
# Per-component 3D RoPE. component_sizes is a list of (t, h, w) patch grid sizes in
|
||||
# sequence order [target, ref0, ref1, ...]; h/w restart at 0 for each component while t
|
||||
# continues from the running offset, giving every image its own temporal position band.
|
||||
freqs_parts = []
|
||||
t_offset = 0
|
||||
for (t, h, w) in component_sizes:
|
||||
freqs = self._get_rotary_pos_embed_for_range(
|
||||
start=(t_offset, 0, 0),
|
||||
stop=(t_offset + t, h, w),
|
||||
device=device,
|
||||
)
|
||||
freqs_parts.append(freqs)
|
||||
t_offset += t
|
||||
return torch.cat(freqs_parts, dim=0).unsqueeze(0).unsqueeze(2)
|
||||
|
||||
def unpatchify(self, x: torch.Tensor, t: int, h: int, w: int) -> torch.Tensor:
|
||||
c = self.out_channels
|
||||
pt, ph, pw = self.patch_size
|
||||
x = x.reshape(x.shape[0], t, h, w, pt, ph, pw, c)
|
||||
x = x.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
return x.reshape(x.shape[0], c, t * pt, h * ph, w * pw)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
context: torch.Tensor = None,
|
||||
ref_latents=None,
|
||||
control=None,
|
||||
transformer_options=None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
transformer_options = {} if transformer_options is None else transformer_options.copy()
|
||||
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(hidden_states, timestep, context, ref_latents, transformer_options, **kwargs)
|
||||
|
||||
def _forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
context: torch.Tensor,
|
||||
ref_latents=None,
|
||||
transformer_options=None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
pt, ph, pw = self.patch_size
|
||||
_, _, ot, oh, ow = hidden_states.shape
|
||||
|
||||
components = [hidden_states, *(ref_latents or [])]
|
||||
component_sizes = []
|
||||
img_tokens = []
|
||||
for comp in components:
|
||||
comp = comfy.ldm.common_dit.pad_to_patch_size(comp, self.patch_size)
|
||||
_, _, ct, ch, cw = comp.shape
|
||||
component_sizes.append((ct // pt, ch // ph, cw // pw))
|
||||
tokens = self.img_in(comp).flatten(2).transpose(1, 2) # (B, n_i, D)
|
||||
img_tokens.append(tokens)
|
||||
|
||||
img = torch.cat(img_tokens, dim=1)
|
||||
|
||||
_, vec, txt = self.condition_embedder(timestep, context)
|
||||
vec = vec.unflatten(1, (6, -1))
|
||||
|
||||
image_rotary_emb = self.get_rotary_pos_embed_for_components(
|
||||
component_sizes,
|
||||
device=hidden_states.device,
|
||||
)
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.double_blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
for i, block in enumerate(self.double_blocks):
|
||||
transformer_options["block_index"] = i
|
||||
if ("double_block", i) in blocks_replace:
|
||||
def block_wrap(args):
|
||||
out = {}
|
||||
out["img"], out["txt"] = block(
|
||||
hidden_states=args["img"],
|
||||
encoder_hidden_states=args["txt"],
|
||||
temb=args["vec"],
|
||||
image_rotary_emb=args["pe"],
|
||||
transformer_options=args.get("transformer_options"),
|
||||
)
|
||||
return out
|
||||
|
||||
out = blocks_replace[("double_block", i)]({"img": img,
|
||||
"txt": txt,
|
||||
"vec": vec,
|
||||
"pe": image_rotary_emb,
|
||||
"transformer_options": transformer_options},
|
||||
{"original_block": block_wrap})
|
||||
txt = out["txt"]
|
||||
img = out["img"]
|
||||
else:
|
||||
img, txt = block(
|
||||
hidden_states=img,
|
||||
encoder_hidden_states=txt,
|
||||
temb=vec,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
transformer_options=transformer_options,
|
||||
)
|
||||
|
||||
tt, th, tw = component_sizes[0]
|
||||
target_tokens = tt * th * tw
|
||||
img = img[:, :target_tokens, :]
|
||||
img = self.proj_out(self.norm_out(img))
|
||||
img = self.unpatchify(img, tt, th, tw)
|
||||
return img[:, :, :ot, :oh, :ow]
|
||||
@@ -709,7 +709,7 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
|
||||
return out
|
||||
|
||||
try:
|
||||
@torch.library.custom_op("comfy::flash_attn", mutates_args=())
|
||||
@torch.library.custom_op("flash_attention::flash_attn", mutates_args=())
|
||||
def flash_attn_wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
|
||||
dropout_p: float = 0.0, causal: bool = False, softmax_scale: float = -1.0) -> torch.Tensor:
|
||||
softmax_scale_arg = None if softmax_scale == -1.0 else softmax_scale
|
||||
|
||||
@@ -22,7 +22,7 @@ def torch_cat_if_needed(xl, dim):
|
||||
else:
|
||||
return None
|
||||
|
||||
def get_timestep_embedding(timesteps, embedding_dim, flip_sin_to_cos=False, downscale_freq_shift=1):
|
||||
def get_timestep_embedding(timesteps, embedding_dim):
|
||||
"""
|
||||
This matches the implementation in Denoising Diffusion Probabilistic Models:
|
||||
From Fairseq.
|
||||
@@ -33,13 +33,11 @@ def get_timestep_embedding(timesteps, embedding_dim, flip_sin_to_cos=False, down
|
||||
assert len(timesteps.shape) == 1
|
||||
|
||||
half_dim = embedding_dim // 2
|
||||
emb = math.log(10000) / (half_dim - downscale_freq_shift)
|
||||
emb = math.log(10000) / (half_dim - 1)
|
||||
emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb)
|
||||
emb = emb.to(device=timesteps.device)
|
||||
emb = timesteps.float()[:, None] * emb[None, :]
|
||||
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1)
|
||||
if flip_sin_to_cos:
|
||||
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
|
||||
if embedding_dim % 2 == 1: # zero pad
|
||||
emb = torch.nn.functional.pad(emb, (0,1,0,0))
|
||||
return emb
|
||||
|
||||
@@ -197,9 +197,6 @@ class PixDiT_T2I(nn.Module):
|
||||
"""Hook for subclasses to inject per-block state into the patch stream (e.g. PiD's LQ gate)."""
|
||||
return s
|
||||
|
||||
def _pre_pixel_blocks(self, s, **kwargs):
|
||||
return s
|
||||
|
||||
def _forward(self, x, timesteps, context=None, attention_mask=None, transformer_options={}, **kwargs):
|
||||
H_orig, W_orig = x.shape[2], x.shape[3]
|
||||
x = comfy.ldm.common_dit.pad_to_patch_size(x, (self.patch_size, self.patch_size))
|
||||
@@ -229,7 +226,6 @@ class PixDiT_T2I(nn.Module):
|
||||
s, y_emb = blk(s, y_emb, condition, pos_img, pos_txt, None, transformer_options=transformer_options)
|
||||
s = F.silu(t_emb + s)
|
||||
|
||||
s = self._pre_pixel_blocks(s, **kwargs)
|
||||
s_cond = s.view(B * L, self.hidden_size)
|
||||
x_pixels = self.pixel_embedder(x, patch_size=self.patch_size)
|
||||
for blk in self.pixel_blocks:
|
||||
|
||||
+14
-50
@@ -13,15 +13,15 @@ from .model import PixDiT_T2I
|
||||
from .modules import precompute_freqs_cis_2d
|
||||
|
||||
|
||||
class SigmaAwareGate(nn.Module):
|
||||
class SigmaAwareGatePerTokenPerDim(nn.Module):
|
||||
"""gate = sigmoid(content_proj(cat[x, lq]) - exp(log_alpha) * sigma); out = x + gate * lq.
|
||||
|
||||
Trained init gives ~0.88 gate at sigma=0, ~0.05 at sigma=1.
|
||||
"""
|
||||
|
||||
def __init__(self, dim: int, per_token: bool = False, dtype=None, device=None, operations=None):
|
||||
def __init__(self, dim: int, dtype=None, device=None, operations=None):
|
||||
super().__init__()
|
||||
self.content_proj = operations.Linear(dim * 2, 1 if per_token else dim, dtype=dtype, device=device)
|
||||
self.content_proj = operations.Linear(dim * 2, dim, dtype=dtype, device=device)
|
||||
self.log_alpha = nn.Parameter(torch.empty((), dtype=dtype, device=device))
|
||||
|
||||
def forward(self, x: torch.Tensor, lq: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor:
|
||||
@@ -36,15 +36,15 @@ class SigmaAwareGate(nn.Module):
|
||||
class ResBlock(nn.Module):
|
||||
"""Pre-activation ResNet block: GN -> SiLU -> Conv -> GN -> SiLU -> Conv + skip."""
|
||||
|
||||
def __init__(self, channels: int, num_groups: int = 4, conv_padding_mode: str = "zeros", dtype=None, device=None, operations=None):
|
||||
def __init__(self, channels: int, num_groups: int = 4, dtype=None, device=None, operations=None):
|
||||
super().__init__()
|
||||
self.block = nn.Sequential(
|
||||
operations.GroupNorm(num_groups, channels, dtype=dtype, device=device),
|
||||
nn.SiLU(),
|
||||
operations.Conv2d(channels, channels, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device),
|
||||
operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device),
|
||||
operations.GroupNorm(num_groups, channels, dtype=dtype, device=device),
|
||||
nn.SiLU(),
|
||||
operations.Conv2d(channels, channels, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device),
|
||||
operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device),
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
@@ -62,13 +62,9 @@ class LQProjection2D(nn.Module):
|
||||
patch_size: int = 16,
|
||||
sr_scale: int = 4,
|
||||
latent_spatial_down_factor: int = 8,
|
||||
latent_unpatchify_factor: int = 1,
|
||||
num_res_blocks: int = 4,
|
||||
num_outputs: int = 7,
|
||||
interval: int = 2,
|
||||
conv_padding_mode: str = "zeros",
|
||||
gate_per_token: bool = False,
|
||||
pit_output: bool = False,
|
||||
dtype=None, device=None, operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
@@ -78,38 +74,34 @@ class LQProjection2D(nn.Module):
|
||||
self.patch_size = patch_size
|
||||
self.sr_scale = sr_scale
|
||||
self.latent_spatial_down_factor = latent_spatial_down_factor
|
||||
self.latent_unpatchify_factor = latent_unpatchify_factor
|
||||
self.num_outputs = num_outputs
|
||||
self.interval = interval
|
||||
|
||||
effective_latent_channels = latent_channels // (latent_unpatchify_factor * latent_unpatchify_factor)
|
||||
effective_spatial_down_factor = latent_spatial_down_factor // latent_unpatchify_factor
|
||||
z_to_patch_ratio = (sr_scale * effective_spatial_down_factor) / patch_size
|
||||
z_to_patch_ratio = (sr_scale * latent_spatial_down_factor) / patch_size
|
||||
self.z_to_patch_ratio = z_to_patch_ratio
|
||||
if z_to_patch_ratio >= 1:
|
||||
self.latent_fold_factor = 0
|
||||
latent_proj_in_ch = effective_latent_channels
|
||||
latent_proj_in_ch = latent_channels
|
||||
else:
|
||||
fold_factor = int(1 / z_to_patch_ratio)
|
||||
assert fold_factor * z_to_patch_ratio == 1.0
|
||||
self.latent_fold_factor = fold_factor
|
||||
latent_proj_in_ch = effective_latent_channels * fold_factor * fold_factor
|
||||
latent_proj_in_ch = latent_channels * fold_factor * fold_factor
|
||||
|
||||
layers = [
|
||||
operations.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device),
|
||||
operations.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device),
|
||||
nn.SiLU(),
|
||||
operations.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device),
|
||||
operations.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device),
|
||||
]
|
||||
for _ in range(num_res_blocks):
|
||||
layers.append(ResBlock(hidden_dim, conv_padding_mode=conv_padding_mode, dtype=dtype, device=device, operations=operations))
|
||||
layers.append(ResBlock(hidden_dim, dtype=dtype, device=device, operations=operations))
|
||||
self.latent_proj = nn.Sequential(*layers)
|
||||
|
||||
self.output_heads = nn.ModuleList(
|
||||
[operations.Linear(hidden_dim, out_dim, dtype=dtype, device=device) for _ in range(num_outputs)]
|
||||
)
|
||||
self.pit_head = operations.Linear(hidden_dim, out_dim, dtype=dtype, device=device) if pit_output else None
|
||||
self.gate_modules = nn.ModuleList(
|
||||
[SigmaAwareGate(out_dim, per_token=gate_per_token, dtype=dtype, device=device, operations=operations)
|
||||
[SigmaAwareGatePerTokenPerDim(out_dim, dtype=dtype, device=device, operations=operations)
|
||||
for _ in range(num_outputs)]
|
||||
)
|
||||
|
||||
@@ -123,11 +115,6 @@ class LQProjection2D(nn.Module):
|
||||
return self.gate_modules[out_idx](x, lq_feature, sigma)
|
||||
|
||||
def _align_latent_to_patch_grid(self, lq_latent: torch.Tensor, pH: int, pW: int) -> torch.Tensor:
|
||||
f = self.latent_unpatchify_factor
|
||||
if f > 1:
|
||||
B, C, H, W = lq_latent.shape
|
||||
lq_latent = lq_latent.reshape(B, C // (f * f), f, f, H, W)
|
||||
lq_latent = lq_latent.permute(0, 1, 4, 2, 5, 3).reshape(B, C // (f * f), H * f, W * f)
|
||||
B, z_dim = lq_latent.shape[:2]
|
||||
if self.z_to_patch_ratio >= 1:
|
||||
if lq_latent.shape[2] != pH or lq_latent.shape[3] != pW:
|
||||
@@ -147,10 +134,7 @@ class LQProjection2D(nn.Module):
|
||||
feat = self._align_latent_to_patch_grid(lq_latent, target_pH, target_pW)
|
||||
B, C, H, W = feat.shape
|
||||
tokens = feat.permute(0, 2, 3, 1).contiguous().view(B, H * W, C)
|
||||
outputs = [head(tokens) for head in self.output_heads]
|
||||
if self.pit_head is not None:
|
||||
outputs.append(self.pit_head(tokens))
|
||||
return outputs
|
||||
return [head(tokens) for head in self.output_heads]
|
||||
|
||||
|
||||
class PidNet(PixDiT_T2I):
|
||||
@@ -164,10 +148,6 @@ class PidNet(PixDiT_T2I):
|
||||
lq_interval: int = 2,
|
||||
sr_scale: int = 4,
|
||||
latent_spatial_down_factor: int = 8,
|
||||
lq_latent_unpatchify_factor: int = 1,
|
||||
lq_conv_padding_mode: str = "zeros",
|
||||
lq_gate_per_token: bool = False,
|
||||
pit_lq_inject: bool = False,
|
||||
rope_ref_h: int = 1024, # NTK ref resolution in PIXEL units: 1024px / patch=16 -> grid_ref=64.
|
||||
rope_ref_w: int = 1024,
|
||||
image_model=None,
|
||||
@@ -185,8 +165,6 @@ class PidNet(PixDiT_T2I):
|
||||
for blk in self.pixel_blocks:
|
||||
blk._rope_fn = _pit_rope_fn
|
||||
|
||||
self.pit_lq_inject = pit_lq_inject
|
||||
|
||||
num_lq_outputs = (self.patch_depth + lq_interval - 1) // lq_interval
|
||||
self.lq_proj = LQProjection2D(
|
||||
latent_channels=lq_latent_channels,
|
||||
@@ -195,20 +173,13 @@ class PidNet(PixDiT_T2I):
|
||||
patch_size=self.patch_size,
|
||||
sr_scale=sr_scale,
|
||||
latent_spatial_down_factor=latent_spatial_down_factor,
|
||||
latent_unpatchify_factor=lq_latent_unpatchify_factor,
|
||||
num_res_blocks=lq_num_res_blocks,
|
||||
num_outputs=num_lq_outputs,
|
||||
interval=lq_interval,
|
||||
conv_padding_mode=lq_conv_padding_mode,
|
||||
gate_per_token=lq_gate_per_token,
|
||||
pit_output=pit_lq_inject,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
self.pit_lq_gate = SigmaAwareGate(
|
||||
self.hidden_size, per_token=lq_gate_per_token, dtype=dtype, device=device, operations=operations
|
||||
) if pit_lq_inject else None
|
||||
|
||||
def _fetch_patch_pos(self, height, width, device, dtype, **rope_opts):
|
||||
return precompute_freqs_cis_2d(
|
||||
@@ -226,11 +197,6 @@ class PidNet(PixDiT_T2I):
|
||||
return s
|
||||
return self.lq_proj.gate(s, pid_lq_features[out_idx], pid_degrade_sigma, out_idx)
|
||||
|
||||
def _pre_pixel_blocks(self, s, pid_pit_lq_feature=None, pid_degrade_sigma=None, **kwargs):
|
||||
if pid_pit_lq_feature is None:
|
||||
return s
|
||||
return self.pit_lq_gate(s, pid_pit_lq_feature, pid_degrade_sigma)
|
||||
|
||||
def _forward(self, x, timesteps, context=None, attention_mask=None, transformer_options={}, lq_latent=None, degrade_sigma=None, **kwargs):
|
||||
if lq_latent is None:
|
||||
raise ValueError("PidNet requires lq_latent — attach via PiDConditioning")
|
||||
@@ -250,14 +216,12 @@ class PidNet(PixDiT_T2I):
|
||||
degrade_sigma = degrade_sigma.expand(B).contiguous()
|
||||
|
||||
lq_features = self.lq_proj(lq_latent=lq_latent.to(x), target_pH=Hs, target_pW=Ws)
|
||||
pit_lq_feature = lq_features.pop() if self.pit_lq_inject else None
|
||||
|
||||
return super()._forward(
|
||||
x, timesteps,
|
||||
context=context, attention_mask=attention_mask,
|
||||
transformer_options=transformer_options,
|
||||
pid_lq_features=lq_features,
|
||||
pid_pit_lq_feature=pit_lq_feature,
|
||||
pid_degrade_sigma=degrade_sigma,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -1,51 +0,0 @@
|
||||
import torch
|
||||
|
||||
from comfy.ldm.modules import attention as _attention
|
||||
|
||||
|
||||
def _var_attention_qkv(q, k, v, heads, skip_reshape):
|
||||
if skip_reshape:
|
||||
return q, k, v, q.shape[-1]
|
||||
total_tokens, embed_dim = q.shape
|
||||
head_dim = embed_dim // heads
|
||||
return (
|
||||
q.view(total_tokens, heads, head_dim),
|
||||
k.view(k.shape[0], heads, head_dim),
|
||||
v.view(v.shape[0], heads, head_dim),
|
||||
head_dim,
|
||||
)
|
||||
|
||||
|
||||
def _var_attention_output(out, heads, head_dim, skip_output_reshape):
|
||||
if skip_output_reshape:
|
||||
return out
|
||||
return out.reshape(-1, heads * head_dim)
|
||||
|
||||
|
||||
def var_attention_optimized_split(q, k, v, heads, cu_seqlens_q, cu_seqlens_k, *args, skip_reshape=False, skip_output_reshape=False, **kwargs):
|
||||
q, k, v, head_dim = _var_attention_qkv(q, k, v, heads, skip_reshape)
|
||||
|
||||
q_split_indices = cu_seqlens_q[1:-1]
|
||||
k_split_indices = cu_seqlens_k[1:-1]
|
||||
if k.shape[0] != v.shape[0]:
|
||||
raise ValueError("cu_seqlens_k does not match v token count")
|
||||
|
||||
q_splits = torch.tensor_split(q, q_split_indices, dim=0)
|
||||
k_splits = torch.tensor_split(k, k_split_indices, dim=0)
|
||||
v_splits = torch.tensor_split(v, k_split_indices, dim=0)
|
||||
if len(q_splits) != len(k_splits) or len(q_splits) != len(v_splits):
|
||||
raise ValueError("cu_seqlens_q and cu_seqlens_k must describe the same sequence count")
|
||||
|
||||
out = []
|
||||
for q_i, k_i, v_i in zip(q_splits, k_splits, v_splits):
|
||||
q_i = q_i.permute(1, 0, 2).unsqueeze(0)
|
||||
k_i = k_i.permute(1, 0, 2).unsqueeze(0)
|
||||
v_i = v_i.permute(1, 0, 2).unsqueeze(0)
|
||||
out_i = _attention.optimized_attention(q_i, k_i, v_i, heads, skip_reshape=True, skip_output_reshape=True)
|
||||
out.append(out_i.squeeze(0).permute(1, 0, 2))
|
||||
|
||||
out = torch.cat(out, dim=0)
|
||||
return _var_attention_output(out, heads, head_dim, skip_output_reshape)
|
||||
|
||||
|
||||
optimized_var_attention = var_attention_optimized_split
|
||||
@@ -1,301 +0,0 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from comfy.ldm.seedvr.constants import (
|
||||
CIELAB_DELTA,
|
||||
CIELAB_KAPPA,
|
||||
D65_WHITE_X,
|
||||
D65_WHITE_Z,
|
||||
WAVELET_DECOMP_LEVELS,
|
||||
)
|
||||
|
||||
|
||||
def wavelet_blur(image: Tensor, radius):
|
||||
max_safe_radius = max(1, min(image.shape[-2:]) // 8)
|
||||
if radius > max_safe_radius:
|
||||
radius = max_safe_radius
|
||||
|
||||
num_channels = image.shape[1]
|
||||
|
||||
kernel_vals = [
|
||||
[0.0625, 0.125, 0.0625],
|
||||
[0.125, 0.25, 0.125],
|
||||
[0.0625, 0.125, 0.0625],
|
||||
]
|
||||
kernel = torch.tensor(kernel_vals, dtype=image.dtype, device=image.device)
|
||||
kernel = kernel[None, None].repeat(num_channels, 1, 1, 1)
|
||||
|
||||
image = F.pad(image, (radius, radius, radius, radius), mode='replicate')
|
||||
output = F.conv2d(image, kernel, groups=num_channels, dilation=radius)
|
||||
|
||||
return output
|
||||
|
||||
def wavelet_decomposition(image: Tensor, levels: int = WAVELET_DECOMP_LEVELS):
|
||||
high_freq = torch.zeros_like(image)
|
||||
|
||||
for i in range(levels):
|
||||
radius = 2 ** i
|
||||
low_freq = wavelet_blur(image, radius)
|
||||
high_freq.add_(image).sub_(low_freq)
|
||||
image = low_freq
|
||||
|
||||
return high_freq, low_freq
|
||||
|
||||
def wavelet_reconstruction(content_feat: Tensor, style_feat: Tensor) -> Tensor:
|
||||
|
||||
if content_feat.shape != style_feat.shape:
|
||||
if len(content_feat.shape) >= 3:
|
||||
style_feat = F.interpolate(
|
||||
style_feat,
|
||||
size=content_feat.shape[-2:],
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
content_high_freq, content_low_freq = wavelet_decomposition(content_feat)
|
||||
del content_low_freq
|
||||
|
||||
style_high_freq, style_low_freq = wavelet_decomposition(style_feat)
|
||||
del style_high_freq
|
||||
|
||||
if content_high_freq.shape != style_low_freq.shape:
|
||||
style_low_freq = F.interpolate(
|
||||
style_low_freq,
|
||||
size=content_high_freq.shape[-2:],
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
content_high_freq.add_(style_low_freq)
|
||||
|
||||
return content_high_freq.clamp_(-1.0, 1.0)
|
||||
|
||||
def _histogram_matching_channel(source: Tensor, reference: Tensor) -> Tensor:
|
||||
original_shape = source.shape
|
||||
|
||||
source_flat = source.flatten()
|
||||
reference_flat = reference.flatten()
|
||||
|
||||
source_sorted, source_indices = torch.sort(source_flat)
|
||||
reference_sorted, _ = torch.sort(reference_flat)
|
||||
del reference_flat
|
||||
|
||||
n_source = len(source_sorted)
|
||||
n_reference = len(reference_sorted)
|
||||
|
||||
if n_source == n_reference:
|
||||
matched_sorted = reference_sorted
|
||||
else:
|
||||
source_quantiles = torch.linspace(0, 1, n_source, device=source.device)
|
||||
ref_indices = (source_quantiles * (n_reference - 1)).long()
|
||||
ref_indices.clamp_(0, n_reference - 1)
|
||||
matched_sorted = reference_sorted[ref_indices]
|
||||
del source_quantiles, ref_indices, reference_sorted
|
||||
|
||||
del source_sorted, source_flat
|
||||
|
||||
inverse_indices = torch.argsort(source_indices)
|
||||
del source_indices
|
||||
matched_flat = matched_sorted[inverse_indices]
|
||||
del matched_sorted, inverse_indices
|
||||
|
||||
return matched_flat.reshape(original_shape)
|
||||
|
||||
def _lab_to_rgb_batch(lab: Tensor, matrix_inv: Tensor, epsilon: float, kappa: float) -> Tensor:
|
||||
L, a, b = lab[:, 0], lab[:, 1], lab[:, 2]
|
||||
|
||||
fy = (L + 16.0) / 116.0
|
||||
fx = a.div(500.0).add_(fy)
|
||||
fz = fy - b / 200.0
|
||||
del L, a, b
|
||||
|
||||
x = torch.where(
|
||||
fx > epsilon,
|
||||
torch.pow(fx, 3.0),
|
||||
fx.mul(116.0).sub_(16.0).div_(kappa)
|
||||
)
|
||||
y = torch.where(
|
||||
fy > epsilon,
|
||||
torch.pow(fy, 3.0),
|
||||
fy.mul(116.0).sub_(16.0).div_(kappa)
|
||||
)
|
||||
z = torch.where(
|
||||
fz > epsilon,
|
||||
torch.pow(fz, 3.0),
|
||||
fz.mul(116.0).sub_(16.0).div_(kappa)
|
||||
)
|
||||
del fx, fy, fz
|
||||
|
||||
x.mul_(D65_WHITE_X)
|
||||
z.mul_(D65_WHITE_Z)
|
||||
|
||||
xyz = torch.stack([x, y, z], dim=1)
|
||||
del x, y, z
|
||||
|
||||
B, _, H, W = xyz.shape
|
||||
xyz_flat = xyz.permute(0, 2, 3, 1).reshape(-1, 3)
|
||||
del xyz
|
||||
|
||||
xyz_flat = xyz_flat.to(dtype=matrix_inv.dtype)
|
||||
rgb_linear_flat = torch.matmul(xyz_flat, matrix_inv.T)
|
||||
del xyz_flat
|
||||
|
||||
rgb_linear = rgb_linear_flat.reshape(B, H, W, 3).permute(0, 3, 1, 2)
|
||||
del rgb_linear_flat
|
||||
|
||||
mask = rgb_linear > 0.0031308
|
||||
rgb = torch.where(
|
||||
mask,
|
||||
torch.pow(torch.clamp(rgb_linear, min=0.0), 1.0 / 2.4).mul_(1.055).sub_(0.055),
|
||||
rgb_linear * 12.92
|
||||
)
|
||||
del mask, rgb_linear
|
||||
|
||||
return torch.clamp(rgb, 0.0, 1.0)
|
||||
|
||||
def _rgb_to_lab_batch(rgb: Tensor, matrix: Tensor, epsilon: float, kappa: float) -> Tensor:
|
||||
mask = rgb > 0.04045
|
||||
rgb_linear = torch.where(
|
||||
mask,
|
||||
torch.pow((rgb + 0.055) / 1.055, 2.4),
|
||||
rgb / 12.92
|
||||
)
|
||||
del mask
|
||||
|
||||
B, _, H, W = rgb_linear.shape
|
||||
rgb_flat = rgb_linear.permute(0, 2, 3, 1).reshape(-1, 3)
|
||||
del rgb_linear
|
||||
|
||||
rgb_flat = rgb_flat.to(dtype=matrix.dtype)
|
||||
xyz_flat = torch.matmul(rgb_flat, matrix.T)
|
||||
del rgb_flat
|
||||
|
||||
xyz = xyz_flat.reshape(B, H, W, 3).permute(0, 3, 1, 2)
|
||||
del xyz_flat
|
||||
|
||||
xyz[:, 0].div_(D65_WHITE_X)
|
||||
xyz[:, 2].div_(D65_WHITE_Z)
|
||||
|
||||
epsilon_cubed = epsilon ** 3
|
||||
mask = xyz > epsilon_cubed
|
||||
f_xyz = torch.where(
|
||||
mask,
|
||||
torch.pow(xyz, 1.0 / 3.0),
|
||||
xyz.mul(kappa).add_(16.0).div_(116.0)
|
||||
)
|
||||
del xyz, mask
|
||||
|
||||
L = f_xyz[:, 1].mul(116.0).sub_(16.0)
|
||||
a = (f_xyz[:, 0] - f_xyz[:, 1]).mul_(500.0)
|
||||
b = (f_xyz[:, 1] - f_xyz[:, 2]).mul_(200.0)
|
||||
del f_xyz
|
||||
|
||||
return torch.stack([L, a, b], dim=1)
|
||||
|
||||
def lab_color_transfer(
|
||||
content_feat: Tensor,
|
||||
style_feat: Tensor,
|
||||
luminance_weight: float = 0.8
|
||||
) -> Tensor:
|
||||
content_feat = wavelet_reconstruction(content_feat, style_feat)
|
||||
|
||||
if content_feat.shape != style_feat.shape:
|
||||
style_feat = F.interpolate(
|
||||
style_feat,
|
||||
size=content_feat.shape[-2:],
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
device = content_feat.device
|
||||
original_dtype = content_feat.dtype
|
||||
content_feat = content_feat.float()
|
||||
style_feat = style_feat.float()
|
||||
|
||||
rgb_to_xyz_matrix = torch.tensor([
|
||||
[0.4124564, 0.3575761, 0.1804375],
|
||||
[0.2126729, 0.7151522, 0.0721750],
|
||||
[0.0193339, 0.1191920, 0.9503041]
|
||||
], dtype=torch.float32, device=device)
|
||||
|
||||
xyz_to_rgb_matrix = torch.tensor([
|
||||
[ 3.2404542, -1.5371385, -0.4985314],
|
||||
[-0.9692660, 1.8760108, 0.0415560],
|
||||
[ 0.0556434, -0.2040259, 1.0572252]
|
||||
], dtype=torch.float32, device=device)
|
||||
|
||||
epsilon = CIELAB_DELTA
|
||||
kappa = CIELAB_KAPPA
|
||||
|
||||
content_feat.add_(1.0).mul_(0.5).clamp_(0.0, 1.0)
|
||||
style_feat.add_(1.0).mul_(0.5).clamp_(0.0, 1.0)
|
||||
|
||||
content_lab = _rgb_to_lab_batch(content_feat, rgb_to_xyz_matrix, epsilon, kappa)
|
||||
del content_feat
|
||||
|
||||
style_lab = _rgb_to_lab_batch(style_feat, rgb_to_xyz_matrix, epsilon, kappa)
|
||||
del style_feat, rgb_to_xyz_matrix
|
||||
|
||||
matched_a = _histogram_matching_channel(content_lab[:, 1], style_lab[:, 1])
|
||||
matched_b = _histogram_matching_channel(content_lab[:, 2], style_lab[:, 2])
|
||||
|
||||
if luminance_weight < 1.0:
|
||||
matched_L = _histogram_matching_channel(content_lab[:, 0], style_lab[:, 0])
|
||||
result_L = content_lab[:, 0].mul(luminance_weight).add_(matched_L.mul(1.0 - luminance_weight))
|
||||
del matched_L
|
||||
else:
|
||||
result_L = content_lab[:, 0]
|
||||
|
||||
del content_lab, style_lab
|
||||
|
||||
result_lab = torch.stack([result_L, matched_a, matched_b], dim=1)
|
||||
del result_L, matched_a, matched_b
|
||||
|
||||
result_rgb = _lab_to_rgb_batch(result_lab, xyz_to_rgb_matrix, epsilon, kappa)
|
||||
del result_lab, xyz_to_rgb_matrix
|
||||
|
||||
result = result_rgb.mul_(2.0).sub_(1.0)
|
||||
del result_rgb
|
||||
|
||||
result = result.to(original_dtype)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def wavelet_color_transfer(content_feat: Tensor, style_feat: Tensor) -> Tensor:
|
||||
return wavelet_reconstruction(content_feat, style_feat)
|
||||
|
||||
|
||||
def adain_color_transfer(content_feat: Tensor, style_feat: Tensor, eps: float = 1e-5) -> Tensor:
|
||||
if content_feat.shape != style_feat.shape:
|
||||
style_feat = F.interpolate(
|
||||
style_feat,
|
||||
size=content_feat.shape[-2:],
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
)
|
||||
|
||||
original_dtype = content_feat.dtype
|
||||
content_feat = content_feat.float()
|
||||
style_feat = style_feat.float()
|
||||
|
||||
b, c = content_feat.shape[:2]
|
||||
content_flat = content_feat.reshape(b, c, -1)
|
||||
style_flat = style_feat.reshape(b, c, -1)
|
||||
|
||||
content_mean = content_flat.mean(dim=2).reshape(b, c, 1, 1)
|
||||
content_std = (content_flat.var(dim=2, correction=0) + eps).sqrt().reshape(b, c, 1, 1)
|
||||
style_mean = style_flat.mean(dim=2).reshape(b, c, 1, 1)
|
||||
style_std = (style_flat.var(dim=2, correction=0) + eps).sqrt().reshape(b, c, 1, 1)
|
||||
del content_flat, style_flat
|
||||
|
||||
normalized = (content_feat - content_mean) / content_std
|
||||
del content_mean, content_std
|
||||
result = normalized * style_std + style_mean
|
||||
del normalized, style_mean, style_std
|
||||
|
||||
result = result.clamp_(-1.0, 1.0)
|
||||
if result.dtype != original_dtype:
|
||||
result = result.to(original_dtype)
|
||||
return result
|
||||
@@ -1,48 +0,0 @@
|
||||
"""SeedVR2 constants."""
|
||||
|
||||
# Temporal chunk-size law: the sampler's activation wall is linear in
|
||||
# T_latent * pixel area (17-cell resolution sweep + T bisection, RTX 5090, 3b fp16):
|
||||
# max_latent_frames = (free_GiB - RESERVED - K*SIGMA) / (GIB_PER_MPX_FRAME * megapixels)
|
||||
# RESERVED covers model staging plus fixed CUDA/torch overhead; SIGMA is the measured
|
||||
# run-to-run spread of the wall; K=4 trades ~10% smaller chunks for ~1e-5 OOM odds.
|
||||
SEEDVR2_CHUNK_GIB_PER_MPX_FRAME = 0.55
|
||||
SEEDVR2_CHUNK_RESERVED_GIB = 8.5
|
||||
SEEDVR2_CHUNK_SIGMA_GIB = 0.55
|
||||
SEEDVR2_CHUNK_SIGMA_K = 4
|
||||
|
||||
SEEDVR2_7B_VID_DIM = 3072
|
||||
SEEDVR2_OOM_BACKOFF_DIVISOR = 2
|
||||
SEEDVR2_DTYPE_BYTES_FLOOR = 4
|
||||
SEEDVR2_7B_MLP_CHUNK = 8192
|
||||
SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS = 4096 # partial-RoPE application token-chunk.
|
||||
SEEDVR2_LATENT_CHANNELS = 16
|
||||
|
||||
SEEDVR2_COLOR_MEM_HEADROOM = 0.75
|
||||
SEEDVR2_LAB_SCALE_MULTIPLIER = 13
|
||||
SEEDVR2_WAVELET_SCALE_MULTIPLIER = 10 # per-frame byte multiplier, wavelet path.
|
||||
SEEDVR2_ADAIN_SCALE_MULTIPLIER = 6
|
||||
|
||||
BYTEDANCE_VAE_SCALING_FACTOR = 0.9152 # configs_3b/main.yaml:57.
|
||||
BYTEDANCE_VAE_SHIFTING_FACTOR = 0.0
|
||||
BYTEDANCE_VAE_CONV_MEM_GIB = 0.5
|
||||
BYTEDANCE_VAE_NORM_MEM_GIB = 0.5
|
||||
BYTEDANCE_LOGVAR_CLAMP_MIN = -30.0 # video_vae_v3/modules/types.py:28.
|
||||
BYTEDANCE_LOGVAR_CLAMP_MAX = 20.0 # video_vae_v3/modules/types.py:28.
|
||||
BYTEDANCE_GN_CHUNKS_FP16 = 4 # causal_inflation_lib.py:351 (GroupNorm chunk count, fp16).
|
||||
BYTEDANCE_GN_CHUNKS_FP32 = 2 # causal_inflation_lib.py:351 (GroupNorm chunk count, fp32).
|
||||
BYTEDANCE_BLOCK_OUT_CHANNELS = (128, 256, 512, 512) # s8_c16_t4_inflation_sd3.yaml:7-11.
|
||||
BYTEDANCE_SLICING_SAMPLE_MIN = 4 # s8_c16_t4_inflation_sd3.yaml:22 (slicing_sample_min_size).
|
||||
BYTEDANCE_VAE_TEMPORAL_DOWNSAMPLE = 4 # infer.py:230 (temporal_downsample_factor); the 4n+1 factor.
|
||||
BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE = 8 # infer.py:231 (spatial_downsample_factor).
|
||||
BYTEDANCE_720P_REF_AREA = 45 * 80 # dit_v2/window.py:32 (720p reference area for window scaling).
|
||||
BYTEDANCE_MAX_TEMPORAL_WINDOW = 30 # dit_v2/window.py:35 (max temporal window frames).
|
||||
BYTEDANCE_ROPE_MAX_FREQ = 256 # dit_v2/rope.py:31 (pixel-RoPE max frequency).
|
||||
BYTEDANCE_SINUSOIDAL_DIM = 256 # dit_3b/nadit.py:120 (timestep sinusoidal embed dim).
|
||||
|
||||
ROPE_THETA = 10000 # RoPE base; Su et al., "RoFormer", arXiv:2104.09864.
|
||||
|
||||
CIELAB_DELTA = 6.0 / 29.0 # CIE 15 (delta).
|
||||
CIELAB_KAPPA = (29.0 / 3.0) ** 3 # CIE 15 (kappa).
|
||||
D65_WHITE_X = 0.95047 # CIE D65 standard illuminant Xn (Yn = 1).
|
||||
D65_WHITE_Z = 1.08883 # CIE D65 standard illuminant Zn.
|
||||
WAVELET_DECOMP_LEVELS = 5 # wavelet color-fix decomposition depth (GIMP/Krita; StableSR).
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -55,10 +55,8 @@ import comfy.ldm.pixeldit.model
|
||||
import comfy.ldm.pixeldit.pid
|
||||
import comfy.ldm.ace.model
|
||||
import comfy.ldm.omnigen.omnigen2
|
||||
import comfy.ldm.seedvr.model
|
||||
import comfy.ldm.boogu.model
|
||||
import comfy.ldm.qwen_image.model
|
||||
import comfy.ldm.joyimage.model
|
||||
import comfy.ldm.ideogram4.model
|
||||
import comfy.ldm.krea2.model
|
||||
import comfy.ldm.kandinsky5.model
|
||||
@@ -934,17 +932,6 @@ class HunyuanDiT(BaseModel):
|
||||
out['image_meta_size'] = comfy.conds.CONDRegular(torch.FloatTensor([[height, width, target_height, target_width, 0, 0]]))
|
||||
return out
|
||||
|
||||
class SeedVR2(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.seedvr.model.NaDiT)
|
||||
|
||||
def extra_conds(self, **kwargs):
|
||||
out = super().extra_conds(**kwargs)
|
||||
condition = kwargs.get("condition", None)
|
||||
if condition is not None:
|
||||
out["condition"] = comfy.conds.CONDRegular(condition)
|
||||
return out
|
||||
|
||||
class PixArt(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.EPS, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.pixart.pixartms.PixArtMS)
|
||||
@@ -2277,28 +2264,6 @@ class QwenImage(BaseModel):
|
||||
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
|
||||
return out
|
||||
|
||||
class JoyImage(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.joyimage.model.JoyImageTransformer3DModel)
|
||||
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:
|
||||
out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents])
|
||||
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 Ideogram4(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.ideogram4.model.Ideogram4Transformer2DModel)
|
||||
|
||||
@@ -470,46 +470,15 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
|
||||
# PiD (Pixel Diffusion Decoder). Must check BEFORE plain PixelDiT_T2I.
|
||||
_lq_w_key = '{}lq_proj.latent_proj.0.weight'.format(key_prefix)
|
||||
if _lq_w_key in state_dict_keys:
|
||||
latent_proj_in_channels = int(state_dict[_lq_w_key].shape[1])
|
||||
hidden_dim = int(state_dict[_lq_w_key].shape[0])
|
||||
in_ch = int(state_dict[_lq_w_key].shape[1])
|
||||
_gate_prefix = '{}lq_proj.gate_modules.'.format(key_prefix)
|
||||
num_gates = len({k[len(_gate_prefix):].split('.')[0]
|
||||
for k in state_dict_keys if k.startswith(_gate_prefix)})
|
||||
pid_v1_5 = '{}lq_proj.pit_head.weight'.format(key_prefix) in state_dict_keys
|
||||
dit_config = {"image_model": "pid",
|
||||
"lq_hidden_dim": hidden_dim}
|
||||
"lq_latent_channels": in_ch,
|
||||
"latent_spatial_down_factor": 16 if in_ch >= 64 else 8}
|
||||
if num_gates > 0:
|
||||
dit_config["lq_interval"] = (14 + num_gates - 1) // num_gates
|
||||
if pid_v1_5:
|
||||
pid_v1_5_variants = {
|
||||
16: { # Flux and QwenImage
|
||||
"lq_latent_channels": 16,
|
||||
"latent_spatial_down_factor": 8,
|
||||
"lq_latent_unpatchify_factor": 1,
|
||||
},
|
||||
32: { # Flux2 after 2x latent unpatchify
|
||||
"lq_latent_channels": 128,
|
||||
"latent_spatial_down_factor": 16,
|
||||
"lq_latent_unpatchify_factor": 2,
|
||||
},
|
||||
}
|
||||
variant = pid_v1_5_variants.get(latent_proj_in_channels)
|
||||
if variant is None:
|
||||
raise ValueError(f"Unsupported PiD v1.5 latent projection with {latent_proj_in_channels} input channels")
|
||||
gate_weight = state_dict['{}lq_proj.gate_modules.0.content_proj.weight'.format(key_prefix)]
|
||||
dit_config.update(variant)
|
||||
dit_config.update({
|
||||
"lq_conv_padding_mode": "replicate",
|
||||
"lq_gate_per_token": gate_weight.shape[0] == 1,
|
||||
"pit_lq_inject": True,
|
||||
"rope_ref_h": 2048,
|
||||
"rope_ref_w": 2048,
|
||||
})
|
||||
else:
|
||||
dit_config.update({
|
||||
"lq_latent_channels": latent_proj_in_channels,
|
||||
"latent_spatial_down_factor": 16 if latent_proj_in_channels >= 64 else 8,
|
||||
})
|
||||
return dit_config
|
||||
|
||||
if '{}core.pixel_embedder.proj.weight'.format(key_prefix) in state_dict_keys: # PixelDiT T2I
|
||||
@@ -629,44 +598,6 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
|
||||
|
||||
return dit_config
|
||||
|
||||
seedvr2_7b_separate_key = "{}blocks.35.mlp.vid.proj_out.weight".format(key_prefix)
|
||||
if seedvr2_7b_separate_key in state_dict_keys and state_dict[seedvr2_7b_separate_key].shape[0] == 3072: # seedvr2 7b
|
||||
dit_config = {}
|
||||
dit_config["image_model"] = "seedvr2"
|
||||
dit_config["vid_dim"] = 3072
|
||||
dit_config["heads"] = 24
|
||||
dit_config["num_layers"] = 36
|
||||
# This checkpoint uses separate vid/txt MMModule keys in every block.
|
||||
dit_config["mm_layers"] = 36
|
||||
dit_config["norm_eps"] = 1e-5
|
||||
dit_config["rope_type"] = "rope3d"
|
||||
dit_config["rope_dim"] = 64
|
||||
dit_config["mlp_type"] = "normal"
|
||||
return dit_config
|
||||
if "{}blocks.35.mlp.all.proj_in_gate.weight".format(key_prefix) in state_dict_keys: # seedvr2 7b
|
||||
dit_config = {}
|
||||
dit_config["image_model"] = "seedvr2"
|
||||
dit_config["vid_dim"] = 3072
|
||||
dit_config["heads"] = 24
|
||||
dit_config["num_layers"] = 36
|
||||
# This checkpoint uses shared all.* MMModule keys after the initial blocks.
|
||||
dit_config["mm_layers"] = 10
|
||||
dit_config["norm_eps"] = 1e-5
|
||||
dit_config["rope_type"] = "rope3d"
|
||||
dit_config["rope_dim"] = 64
|
||||
dit_config["mlp_type"] = "swiglu"
|
||||
return dit_config
|
||||
if "{}blocks.31.mlp.all.proj_in_gate.weight".format(key_prefix) in state_dict_keys: # seedvr2 3b
|
||||
dit_config = {}
|
||||
dit_config["image_model"] = "seedvr2"
|
||||
dit_config["vid_dim"] = 2560
|
||||
dit_config["heads"] = 20
|
||||
dit_config["num_layers"] = 32
|
||||
dit_config["norm_eps"] = 1.0e-05
|
||||
dit_config["mlp_type"] = "swiglu"
|
||||
dit_config["vid_out_norm"] = True
|
||||
return dit_config
|
||||
|
||||
if '{}head.modulation'.format(key_prefix) in state_dict_keys: # Wan 2.1
|
||||
dit_config = {}
|
||||
dit_config["image_model"] = "wan2.1"
|
||||
@@ -1058,25 +989,6 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
|
||||
dit_config["image_model"] = "SAM31"
|
||||
return dit_config
|
||||
|
||||
if (
|
||||
'{}double_blocks.0.attn.img_attn_qkv.weight'.format(key_prefix) in state_dict_keys
|
||||
and '{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix) in state_dict_keys
|
||||
and '{}condition_embedder.time_embedder.linear_1.weight'.format(key_prefix) in state_dict_keys
|
||||
and '{}img_in.weight'.format(key_prefix) in state_dict_keys
|
||||
and len(state_dict['{}img_in.weight'.format(key_prefix)].shape) == 5
|
||||
):
|
||||
img_in = state_dict['{}img_in.weight'.format(key_prefix)]
|
||||
head_dim = state_dict['{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix)].shape[0]
|
||||
return {
|
||||
"image_model": "joyimage",
|
||||
"in_channels": img_in.shape[1],
|
||||
"hidden_size": img_in.shape[0],
|
||||
"patch_size": list(img_in.shape[2:]),
|
||||
"num_layers": count_blocks(state_dict_keys, '{}double_blocks.'.format(key_prefix) + '{}.'),
|
||||
"num_attention_heads": img_in.shape[0] // head_dim,
|
||||
"text_dim": 4096,
|
||||
}
|
||||
|
||||
if '{}input_blocks.0.0.weight'.format(key_prefix) not in state_dict_keys:
|
||||
return None
|
||||
|
||||
@@ -1207,10 +1119,9 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
|
||||
|
||||
return unet_config
|
||||
|
||||
|
||||
def model_config_from_unet_config(unet_config, state_dict=None, unet_key_prefix=""):
|
||||
def model_config_from_unet_config(unet_config, state_dict=None):
|
||||
for model_config in comfy.supported_models.models:
|
||||
if model_config.matches(unet_config, state_dict, unet_key_prefix=unet_key_prefix):
|
||||
if model_config.matches(unet_config, state_dict):
|
||||
return model_config(unet_config)
|
||||
|
||||
logging.error("no match {}".format(unet_config))
|
||||
@@ -1220,7 +1131,7 @@ def model_config_from_unet(state_dict, unet_key_prefix, use_base_if_no_match=Fal
|
||||
unet_config = detect_unet_config(state_dict, unet_key_prefix, metadata=metadata)
|
||||
if unet_config is None:
|
||||
return None
|
||||
model_config = model_config_from_unet_config(unet_config, state_dict, unet_key_prefix)
|
||||
model_config = model_config_from_unet_config(unet_config, state_dict)
|
||||
if model_config is None and use_base_if_no_match:
|
||||
model_config = comfy.supported_models_base.BASE(unet_config)
|
||||
|
||||
|
||||
@@ -616,8 +616,6 @@ PIN_PRESSURE_HYSTERESIS = 256 * 1024 * 1024
|
||||
#Freeing registerables on pressure does imply a GPU sync, so go big on
|
||||
#the hysteresis so each expensive sync gives us back a good chunk.
|
||||
REGISTERABLE_PIN_HYSTERESIS = 2048 * 1024 * 1024
|
||||
WINDOWS_PIN_EVICTION_SWAP_PERCENT = 5.0
|
||||
WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE = 512 * 1024 ** 2
|
||||
|
||||
def module_size(module):
|
||||
module_mem = 0
|
||||
@@ -644,15 +642,6 @@ def free_pins(size, evict_active=False):
|
||||
size -= freed
|
||||
return freed_total
|
||||
|
||||
def should_free_pins_for_ram_pressure(shortfall):
|
||||
if shortfall <= 0:
|
||||
return False
|
||||
if not WINDOWS:
|
||||
return True
|
||||
if psutil.virtual_memory().available < WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE:
|
||||
return True
|
||||
return psutil.swap_memory().percent >= WINDOWS_PIN_EVICTION_SWAP_PERCENT
|
||||
|
||||
def ensure_pin_budget(size, evict_active=False):
|
||||
if args.high_ram:
|
||||
return True
|
||||
|
||||
+7
-33
@@ -1104,21 +1104,6 @@ def _load_quantized_module(module, super_load, state_dict, prefix, local_metadat
|
||||
scales["convrot_groupsize"] = int(
|
||||
layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256))
|
||||
)
|
||||
elif module.quant_format == "convrot_w4a4":
|
||||
scale = pop_scale("weight_scale")
|
||||
if scale is None:
|
||||
raise ValueError(f"Missing ConvRot W4A4 weight scale for layer {layer_name}")
|
||||
params_conf = layer_conf.get("params", {})
|
||||
if not isinstance(params_conf, dict):
|
||||
params_conf = {}
|
||||
scales = {
|
||||
"scale": scale,
|
||||
"convrot_groupsize": int(
|
||||
layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256))
|
||||
),
|
||||
"quant_group_size": 64,
|
||||
"linear_dtype": layer_conf.get("linear_dtype", params_conf.get("linear_dtype", "int4")),
|
||||
}
|
||||
else:
|
||||
raise ValueError(f"Unsupported quantization format: {module.quant_format}")
|
||||
|
||||
@@ -1165,11 +1150,6 @@ def _quantized_weight_state_dict(module, sd, prefix, extra_quant_conf=None, extr
|
||||
if module.quant_format == "int8_tensorwise" and getattr(params, "convrot", False):
|
||||
quant_conf["convrot"] = True
|
||||
quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256)
|
||||
elif module.quant_format == "convrot_w4a4":
|
||||
quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256)
|
||||
linear_dtype = getattr(params, "linear_dtype", "int4")
|
||||
if linear_dtype != "int4":
|
||||
quant_conf["linear_dtype"] = linear_dtype
|
||||
if extra_quant_conf:
|
||||
quant_conf.update(extra_quant_conf)
|
||||
sd[f"{prefix}comfy_quant"] = torch.tensor(list(json.dumps(quant_conf).encode("utf-8")), dtype=torch.uint8)
|
||||
@@ -1257,7 +1237,7 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
|
||||
run_every_op()
|
||||
|
||||
input_shape = input.shape
|
||||
reshaped_nd = False
|
||||
reshaped_3d = False
|
||||
#If cast needs to apply lora, it should be done in the compute dtype
|
||||
compute_dtype = input.dtype
|
||||
|
||||
@@ -1294,12 +1274,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
|
||||
# Inference path (unchanged)
|
||||
if _use_quantized and quantize_input:
|
||||
|
||||
# Reshape >=3D tensors to 2D for quantization (needed for NVFP4 and others)
|
||||
input_reshaped = input.reshape(-1, input_shape[-1]) if input.ndim >= 3 else input
|
||||
# Reshape 3D tensors to 2D for quantization (needed for NVFP4 and others)
|
||||
input_reshaped = input.reshape(-1, input_shape[2]) if input.ndim == 3 else input
|
||||
|
||||
# Fall back to non-quantized for non-2D tensors
|
||||
if input_reshaped.ndim == 2:
|
||||
reshaped_nd = input.ndim >= 3
|
||||
reshaped_3d = input.ndim == 3
|
||||
# dtype is now implicit in the layout class
|
||||
scale = getattr(self, 'input_scale', None)
|
||||
if scale is not None:
|
||||
@@ -1314,9 +1294,9 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
|
||||
weight_only_quant=weight_only_quant,
|
||||
)
|
||||
|
||||
# Reshape output back to original rank if input was >2D
|
||||
if reshaped_nd:
|
||||
output = output.reshape((*input_shape[:-1], self.weight.shape[0]))
|
||||
# Reshape output back to 3D if input was 3D
|
||||
if reshaped_3d:
|
||||
output = output.reshape((input_shape[0], input_shape[1], self.weight.shape[0]))
|
||||
|
||||
return output
|
||||
|
||||
@@ -1450,12 +1430,6 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
|
||||
}
|
||||
if hasattr(params, "block_scale"): # NVFP4
|
||||
kwargs["block_scale"] = params.block_scale[i]
|
||||
if hasattr(params, "quant_group_size"):
|
||||
kwargs["quant_group_size"] = params.quant_group_size
|
||||
if hasattr(params, "convrot_groupsize"):
|
||||
kwargs["convrot_groupsize"] = params.convrot_groupsize
|
||||
if hasattr(params, "linear_dtype"):
|
||||
kwargs["linear_dtype"] = params.linear_dtype
|
||||
return QuantizedTensor(weight._qdata[i], weight._layout_cls, type(params)(**kwargs))
|
||||
|
||||
def state_dict(self, *args, destination=None, prefix="", **kwargs):
|
||||
|
||||
+2
-44
@@ -3,22 +3,6 @@ import logging
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
|
||||
def _rocm_kitchen_arch_supported():
|
||||
"""comfy-kitchen's INT8 Triton kernels compile tl.dot to matrix-core instructions.
|
||||
RDNA3/3.5/4 (gfx11xx/gfx12xx) have WMMA and CDNA (gfx9xx) has MFMA; RDNA1/RDNA2
|
||||
(gfx10xx) have neither, so the INT8 path hangs the GPU there. Gates the automatic
|
||||
ROCm default so those cards stay on the eager fallback (an explicit
|
||||
--enable-triton-backend still forces it on any arch)."""
|
||||
try:
|
||||
arch = torch.cuda.get_device_properties(torch.cuda.current_device()).gcnArchName.split(":")[0]
|
||||
except Exception:
|
||||
return False
|
||||
if arch.startswith(("gfx11", "gfx12")):
|
||||
return True
|
||||
return arch in ("gfx908", "gfx90a", "gfx940", "gfx941", "gfx942", "gfx950")
|
||||
|
||||
|
||||
try:
|
||||
import comfy_kitchen as ck
|
||||
from comfy_kitchen.tensor import (
|
||||
@@ -26,7 +10,6 @@ try:
|
||||
QuantizedLayout,
|
||||
TensorCoreFP8Layout as _CKFp8Layout,
|
||||
TensorCoreNVFP4Layout as _CKNvfp4Layout,
|
||||
TensorCoreConvRotW4A4Layout as _CKTensorCoreConvRotW4A4Layout,
|
||||
TensorWiseINT8Layout as _CKTensorWiseINT8Layout,
|
||||
register_layout_op,
|
||||
register_layout_class,
|
||||
@@ -41,22 +24,10 @@ try:
|
||||
ck.registry.disable("cuda")
|
||||
logging.warning("WARNING: You need pytorch with cu130 or higher to use optimized CUDA operations.")
|
||||
|
||||
# On ROCm/AMD the CUDA backend is unavailable, so Triton is the only accelerated
|
||||
# comfy-kitchen backend. Enable it by default there, but only on Triton >= 3.7 AND a
|
||||
# matrix-core GPU (RDNA3+ WMMA gfx11xx/gfx12xx, CDNA MFMA gfx9xx). RDNA1/RDNA2
|
||||
# (gfx10xx) have no WMMA -> the INT8 tl.dot path hangs the GPU, so they stay eager.
|
||||
# older Triton lacks libdevice.rint on the HIP backend and hard-crashes the INT8 path.
|
||||
if args.disable_triton_backend:
|
||||
ck.registry.disable("triton")
|
||||
elif args.enable_triton_backend: # or (torch.version.hip is not None and _rocm_kitchen_arch_supported()):
|
||||
if args.enable_triton_backend:
|
||||
try:
|
||||
import triton
|
||||
triton_version = tuple(int(v) for v in triton.__version__.split(".")[:2])
|
||||
if args.enable_triton_backend or triton_version >= (3, 7):
|
||||
logging.info("Found triton %s. Enabling comfy-kitchen triton backend.", triton.__version__)
|
||||
else:
|
||||
logging.info("Triton %s is too old for the ROCm INT8 path (needs >= 3.7); comfy-kitchen triton backend disabled.", triton.__version__)
|
||||
ck.registry.disable("triton")
|
||||
logging.info("Found triton %s. Enabling comfy-kitchen triton backend.", triton.__version__)
|
||||
except ImportError as e:
|
||||
logging.error(f"Failed to import triton, Error: {e}, the comfy-kitchen triton backend will not be available.")
|
||||
ck.registry.disable("triton")
|
||||
@@ -80,9 +51,6 @@ except ImportError as e:
|
||||
class _CKTensorWiseINT8Layout:
|
||||
pass
|
||||
|
||||
class _CKTensorCoreConvRotW4A4Layout:
|
||||
pass
|
||||
|
||||
def register_layout_class(name, cls):
|
||||
pass
|
||||
|
||||
@@ -211,7 +179,6 @@ class TensorCoreFP8E5M2Layout(_TensorCoreFP8LayoutBase):
|
||||
# Backward compatibility alias - default to E4M3
|
||||
TensorCoreFP8Layout = TensorCoreFP8E4M3Layout
|
||||
TensorWiseINT8Layout = _CKTensorWiseINT8Layout
|
||||
TensorCoreConvRotW4A4Layout = _CKTensorCoreConvRotW4A4Layout
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
@@ -223,7 +190,6 @@ register_layout_class("TensorCoreFP8E4M3Layout", TensorCoreFP8E4M3Layout)
|
||||
register_layout_class("TensorCoreFP8E5M2Layout", TensorCoreFP8E5M2Layout)
|
||||
register_layout_class("TensorCoreNVFP4Layout", TensorCoreNVFP4Layout)
|
||||
register_layout_class("TensorWiseINT8Layout", _CKTensorWiseINT8Layout)
|
||||
register_layout_class("TensorCoreConvRotW4A4Layout", _CKTensorCoreConvRotW4A4Layout)
|
||||
if _CK_MXFP8_AVAILABLE:
|
||||
register_layout_class("TensorCoreMXFP8Layout", TensorCoreMXFP8Layout)
|
||||
|
||||
@@ -261,13 +227,6 @@ QUANT_ALGOS["int8_tensorwise"] = {
|
||||
"quantize_input": False,
|
||||
}
|
||||
|
||||
QUANT_ALGOS["convrot_w4a4"] = {
|
||||
"storage_t": torch.int8,
|
||||
"parameters": {"weight_scale"},
|
||||
"comfy_tensor_layout": "TensorCoreConvRotW4A4Layout",
|
||||
"quantize_input": False,
|
||||
}
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Re-exports for backward compatibility
|
||||
@@ -280,7 +239,6 @@ __all__ = [
|
||||
"TensorCoreFP8E4M3Layout",
|
||||
"TensorCoreFP8E5M2Layout",
|
||||
"TensorCoreNVFP4Layout",
|
||||
"TensorCoreConvRotW4A4Layout",
|
||||
"TensorWiseINT8Layout",
|
||||
"QUANT_ALGOS",
|
||||
"register_layout_op",
|
||||
|
||||
+19
-89
@@ -16,7 +16,6 @@ import comfy.ldm.cosmos.vae
|
||||
import comfy.ldm.wan.vae
|
||||
import comfy.ldm.wan.vae2_2
|
||||
import comfy.ldm.hunyuan3d.vae
|
||||
import comfy.ldm.seedvr.vae
|
||||
import comfy.ldm.triposplat.vae
|
||||
import comfy.ldm.ace.vae.music_dcae_pipeline
|
||||
import comfy.ldm.cogvideo.vae
|
||||
@@ -76,7 +75,6 @@ import comfy.text_encoders.gemma4
|
||||
import comfy.text_encoders.cogvideo
|
||||
import comfy.text_encoders.sa3
|
||||
import comfy.text_encoders.gpt_oss
|
||||
import comfy.text_encoders.joyimage
|
||||
|
||||
import comfy.model_patcher
|
||||
import comfy.lora
|
||||
@@ -475,8 +473,7 @@ class CLIP:
|
||||
|
||||
class VAE:
|
||||
def __init__(self, sd=None, device=None, config=None, dtype=None, metadata=None):
|
||||
is_seedvr2_vae = "decoder.up_blocks.2.upsamplers.0.upscale_conv.weight" in sd
|
||||
if not is_seedvr2_vae and 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format
|
||||
if 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format
|
||||
sd = diffusers_convert.convert_vae_state_dict(sd)
|
||||
|
||||
if model_management.is_amd():
|
||||
@@ -503,8 +500,6 @@ class VAE:
|
||||
self.upscale_index_formula = None
|
||||
self.extra_1d_channel = None
|
||||
self.crop_input = True
|
||||
self.handles_tiling = False
|
||||
self.format_encoded = None
|
||||
|
||||
self.audio_sample_rate = 44100
|
||||
|
||||
@@ -551,22 +546,6 @@ class VAE:
|
||||
self.first_stage_model = StageC_coder()
|
||||
self.downscale_ratio = 32
|
||||
self.latent_channels = 16
|
||||
elif "decoder.up_blocks.2.upsamplers.0.upscale_conv.weight" in sd: # seedvr2
|
||||
self.first_stage_model = comfy.ldm.seedvr.vae.VideoAutoencoderKLWrapper()
|
||||
self.latent_channels = comfy.ldm.seedvr.vae.SEEDVR2_LATENT_CHANNELS
|
||||
self.latent_dim = 3
|
||||
self.disable_offload = True
|
||||
self.memory_used_decode = lambda shape, dtype: self.first_stage_model.comfy_memory_used_decode(shape)
|
||||
self.memory_used_encode = lambda shape, dtype: (max(shape[2], 5) * shape[3] * shape[4] * 64) * model_management.dtype_size(dtype)
|
||||
self.working_dtypes = [torch.float16, torch.bfloat16, torch.float32]
|
||||
self.handles_tiling = True
|
||||
self.format_encoded = self.first_stage_model.comfy_format_encoded
|
||||
self.downscale_ratio = (lambda a: max(0, math.floor((a + 3) / 4)), 8, 8)
|
||||
self.downscale_index_formula = (4, 8, 8)
|
||||
self.upscale_ratio = (lambda a: max(0, a * 4 - 3), 8, 8)
|
||||
self.upscale_index_formula = (4, 8, 8)
|
||||
self.process_input = lambda image: image * 2.0 - 1.0
|
||||
self.crop_input = False
|
||||
elif "decoder.conv_in.weight" in sd:
|
||||
if sd['decoder.conv_in.weight'].shape[1] == 64:
|
||||
ddconfig = {"block_out_channels": [128, 256, 512, 512, 1024, 1024], "in_channels": 3, "out_channels": 3, "num_res_blocks": 2, "ffactor_spatial": 32, "downsample_match_channel": True, "upsample_match_channel": True}
|
||||
@@ -1033,10 +1012,6 @@ class VAE:
|
||||
decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype())
|
||||
return self.process_output(comfy.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, index_formulas=self.upscale_index_formula, output_device=self.output_device))
|
||||
|
||||
def _decode_tiled_owned(self, samples, **kwargs):
|
||||
out = self.first_stage_model.decode_tiled(samples.to(self.vae_dtype).to(self.device), **kwargs)
|
||||
return self.process_output(out.to(device=self.output_device, dtype=self.vae_output_dtype(), copy=True))
|
||||
|
||||
def encode_tiled_(self, pixel_samples, tile_x=512, tile_y=512, overlap = 64):
|
||||
steps = pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x, tile_y, overlap)
|
||||
steps += pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x // 2, tile_y * 2, overlap)
|
||||
@@ -1073,25 +1048,6 @@ class VAE:
|
||||
encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype())
|
||||
return comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.downscale_ratio, out_channels=self.latent_channels, downscale=True, index_formulas=self.downscale_index_formula, output_device=self.output_device)
|
||||
|
||||
def _encode_tiled_owned(self, pixel_samples, **kwargs):
|
||||
x = self.process_input(pixel_samples).to(self.vae_dtype).to(self.device)
|
||||
out = self.first_stage_model.encode_tiled(x, **kwargs)
|
||||
return out.to(device=self.output_device, dtype=self.vae_output_dtype())
|
||||
|
||||
def _owned_tiled_args(self, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
|
||||
args = {}
|
||||
if tile_x is not None:
|
||||
args["tile_x"] = tile_x
|
||||
if tile_y is not None:
|
||||
args["tile_y"] = tile_y
|
||||
if overlap is not None:
|
||||
args["overlap"] = overlap
|
||||
if tile_t is not None:
|
||||
args["tile_t"] = tile_t
|
||||
if overlap_t is not None:
|
||||
args["overlap_t"] = overlap_t
|
||||
return args
|
||||
|
||||
def decode(self, samples_in, vae_options={}):
|
||||
self.throw_exception_if_invalid()
|
||||
pixel_samples = None
|
||||
@@ -1139,19 +1095,11 @@ class VAE:
|
||||
if dims == 1 or self.extra_1d_channel is not None:
|
||||
pixel_samples = self.decode_tiled_1d(samples_in)
|
||||
elif dims == 2:
|
||||
if self.handles_tiling:
|
||||
tile = 256 // self.spacial_compression_decode()
|
||||
overlap = tile // 4
|
||||
pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap)
|
||||
else:
|
||||
pixel_samples = self.decode_tiled_(samples_in)
|
||||
pixel_samples = self.decode_tiled_(samples_in)
|
||||
elif dims == 3:
|
||||
tile = 256 // self.spacial_compression_decode()
|
||||
overlap = tile // 4
|
||||
if self.handles_tiling:
|
||||
pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap)
|
||||
else:
|
||||
pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap))
|
||||
pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap))
|
||||
|
||||
pixel_samples = pixel_samples.to(self.output_device).movedim(1,-1)
|
||||
return pixel_samples
|
||||
@@ -1170,9 +1118,7 @@ class VAE:
|
||||
args["overlap"] = overlap
|
||||
|
||||
with model_management.cuda_device_context(self.device):
|
||||
if self.handles_tiling and dims in (2, 3):
|
||||
output = self._decode_tiled_owned(samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t))
|
||||
elif dims == 1 or self.extra_1d_channel is not None:
|
||||
if dims == 1 or self.extra_1d_channel is not None:
|
||||
args.pop("tile_y")
|
||||
output = self.decode_tiled_1d(samples, **args)
|
||||
elif dims == 2:
|
||||
@@ -1233,17 +1179,12 @@ class VAE:
|
||||
if self.latent_dim == 3:
|
||||
tile = 256
|
||||
overlap = tile // 4
|
||||
if self.handles_tiling:
|
||||
samples = self._encode_tiled_owned(pixel_samples, tile_x=tile, tile_y=tile, overlap=overlap)
|
||||
else:
|
||||
samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap))
|
||||
samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap))
|
||||
elif self.latent_dim == 1 or self.extra_1d_channel is not None:
|
||||
samples = self.encode_tiled_1d(pixel_samples)
|
||||
else:
|
||||
samples = self.encode_tiled_(pixel_samples)
|
||||
|
||||
if self.format_encoded is not None:
|
||||
samples = self.format_encoded(samples)
|
||||
return samples
|
||||
|
||||
def encode_tiled(self, pixel_samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
|
||||
@@ -1251,7 +1192,7 @@ class VAE:
|
||||
pixel_samples = self.vae_encode_crop_pixels(pixel_samples)
|
||||
dims = self.latent_dim
|
||||
pixel_samples = pixel_samples.movedim(-1, 1)
|
||||
if dims == 3 and pixel_samples.ndim < 5:
|
||||
if dims == 3:
|
||||
if not self.not_video:
|
||||
pixel_samples = pixel_samples.movedim(1, 0).unsqueeze(0)
|
||||
else:
|
||||
@@ -1275,27 +1216,21 @@ class VAE:
|
||||
elif dims == 2:
|
||||
samples = self.encode_tiled_(pixel_samples, **args)
|
||||
elif dims == 3:
|
||||
if self.handles_tiling:
|
||||
samples = self._encode_tiled_owned(pixel_samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t))
|
||||
if tile_t is not None:
|
||||
tile_t_latent = max(2, self.downscale_ratio[0](tile_t))
|
||||
else:
|
||||
if tile_t is not None:
|
||||
tile_t_latent = max(2, self.downscale_ratio[0](tile_t))
|
||||
else:
|
||||
tile_t_latent = 9999
|
||||
args["tile_t"] = self.upscale_ratio[0](tile_t_latent)
|
||||
tile_t_latent = 9999
|
||||
args["tile_t"] = self.upscale_ratio[0](tile_t_latent)
|
||||
|
||||
spatial_overlap = overlap if overlap is not None else 64
|
||||
if overlap_t is None:
|
||||
args["overlap"] = (1, spatial_overlap, spatial_overlap)
|
||||
else:
|
||||
args["overlap"] = (self.upscale_ratio[0](max(1, min(tile_t_latent // 2, self.downscale_ratio[0](overlap_t)))), spatial_overlap, spatial_overlap)
|
||||
maximum = pixel_samples.shape[2]
|
||||
maximum = self.upscale_ratio[0](self.downscale_ratio[0](maximum))
|
||||
if overlap_t is None:
|
||||
args["overlap"] = (1, overlap, overlap)
|
||||
else:
|
||||
args["overlap"] = (self.upscale_ratio[0](max(1, min(tile_t_latent // 2, self.downscale_ratio[0](overlap_t)))), overlap, overlap)
|
||||
maximum = pixel_samples.shape[2]
|
||||
maximum = self.upscale_ratio[0](self.downscale_ratio[0](maximum))
|
||||
|
||||
samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args)
|
||||
samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args)
|
||||
|
||||
if self.format_encoded is not None:
|
||||
samples = self.format_encoded(samples)
|
||||
return samples
|
||||
|
||||
def get_sd(self):
|
||||
@@ -1378,7 +1313,6 @@ class CLIPType(Enum):
|
||||
IDEOGRAM4 = 30
|
||||
BOOGU = 31
|
||||
KREA2 = 32
|
||||
JOYIMAGE = 33
|
||||
|
||||
|
||||
|
||||
@@ -1708,10 +1642,6 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
|
||||
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
|
||||
clip_target.clip = comfy.text_encoders.krea2.te(**llama_detect(clip_data))
|
||||
clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer
|
||||
elif clip_type == CLIPType.JOYIMAGE and te_model == TEModel.QWEN3VL_8B: # JoyImageEdit: full Qwen3-VL-8B, edit-conditioning template + drop_idx.
|
||||
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
|
||||
clip_target.clip = comfy.text_encoders.joyimage.te(**llama_detect(clip_data))
|
||||
clip_target.tokenizer = comfy.text_encoders.joyimage.JoyImageTokenizer
|
||||
elif clip_type in (CLIPType.FLUX, CLIPType.FLUX2): # Flux2 Klein reuses the Qwen3-VL LM (3-layer tap -> 12288); visual unused.
|
||||
klein_model_type = "qwen3_8b" if te_model == TEModel.QWEN3VL_8B else "qwen3_4b"
|
||||
clip_target.clip = comfy.text_encoders.flux.klein_te(**llama_detect(clip_data), model_type=klein_model_type)
|
||||
@@ -1968,7 +1898,7 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c
|
||||
manual_cast_dtype = model_management.unet_manual_cast(None, load_device, model_config.supported_inference_dtypes)
|
||||
else:
|
||||
manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes)
|
||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype, device=load_device)
|
||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype)
|
||||
|
||||
if model_config.clip_vision_prefix is not None:
|
||||
if output_clipvision:
|
||||
@@ -2109,7 +2039,7 @@ def load_diffusion_model_state_dict(sd, model_options={}, metadata=None, disable
|
||||
manual_cast_dtype = model_management.unet_manual_cast(None, load_device, model_config.supported_inference_dtypes)
|
||||
else:
|
||||
manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes)
|
||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype, device=load_device)
|
||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype)
|
||||
|
||||
if custom_operations is not None:
|
||||
model_config.custom_operations = custom_operations
|
||||
|
||||
@@ -27,7 +27,6 @@ import comfy.text_encoders.z_image
|
||||
import comfy.text_encoders.ideogram4
|
||||
import comfy.text_encoders.boogu
|
||||
import comfy.text_encoders.krea2
|
||||
import comfy.text_encoders.joyimage
|
||||
import comfy.text_encoders.anima
|
||||
import comfy.text_encoders.ace15
|
||||
import comfy.text_encoders.longcat_image
|
||||
@@ -1686,40 +1685,6 @@ class Chroma(supported_models_base.BASE):
|
||||
t5_detect = comfy.text_encoders.sd3_clip.t5_xxl_detect(state_dict, "{}t5xxl.transformer.".format(pref))
|
||||
return supported_models_base.ClipTarget(comfy.text_encoders.pixart_t5.PixArtTokenizer, comfy.text_encoders.pixart_t5.pixart_te(**t5_detect))
|
||||
|
||||
class SeedVR2(supported_models_base.BASE):
|
||||
unet_config = {
|
||||
"image_model": "seedvr2"
|
||||
}
|
||||
unet_extra_config = {}
|
||||
required_keys = {
|
||||
"{}positive_conditioning",
|
||||
"{}negative_conditioning",
|
||||
}
|
||||
latent_format = comfy.latent_formats.SeedVR2
|
||||
|
||||
vae_key_prefix = ["vae."]
|
||||
text_encoder_key_prefix = ["text_encoders."]
|
||||
supported_inference_dtypes = [torch.bfloat16, torch.float16, torch.float32]
|
||||
sampling_settings = {
|
||||
"shift": 1.0,
|
||||
}
|
||||
|
||||
def set_inference_dtype(self, dtype, manual_cast_dtype, device=None):
|
||||
if (
|
||||
dtype == torch.float16
|
||||
and manual_cast_dtype is None
|
||||
and comfy.model_management.should_use_bf16(device)
|
||||
):
|
||||
manual_cast_dtype = torch.bfloat16
|
||||
super().set_inference_dtype(dtype, manual_cast_dtype, device=device)
|
||||
|
||||
def get_model(self, state_dict, prefix="", device=None):
|
||||
out = model_base.SeedVR2(self, device=device)
|
||||
return out
|
||||
|
||||
def clip_target(self, state_dict={}):
|
||||
return None
|
||||
|
||||
class ChromaRadiance(Chroma):
|
||||
unet_config = {
|
||||
"image_model": "chroma_radiance",
|
||||
@@ -1912,38 +1877,6 @@ class QwenImage(supported_models_base.BASE):
|
||||
hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen25_7b.transformer.".format(pref))
|
||||
return supported_models_base.ClipTarget(comfy.text_encoders.qwen_image.QwenImageTokenizer, comfy.text_encoders.qwen_image.te(**hunyuan_detect))
|
||||
|
||||
class JoyImage(supported_models_base.BASE):
|
||||
unet_config = {
|
||||
"image_model": "joyimage",
|
||||
}
|
||||
|
||||
sampling_settings = {
|
||||
"multiplier": 1000,
|
||||
"shift": 1.5,
|
||||
}
|
||||
|
||||
memory_usage_factor = 1.8
|
||||
|
||||
unet_extra_config = {
|
||||
"theta": 10000,
|
||||
"rope_dim_list": [16, 56, 56],
|
||||
}
|
||||
|
||||
latent_format = latent_formats.Wan21
|
||||
|
||||
supported_inference_dtypes = [torch.bfloat16, torch.float32]
|
||||
|
||||
vae_key_prefix = ["vae."]
|
||||
text_encoder_key_prefix = ["text_encoders."]
|
||||
|
||||
def get_model(self, state_dict, prefix="", device=None):
|
||||
return model_base.JoyImage(self, device=device)
|
||||
|
||||
def clip_target(self, state_dict={}):
|
||||
pref = self.text_encoder_key_prefix[0]
|
||||
qwen3vl_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl.transformer.".format(pref))
|
||||
return supported_models_base.ClipTarget(comfy.text_encoders.joyimage.JoyImageTokenizer, comfy.text_encoders.joyimage.te(**qwen3vl_detect))
|
||||
|
||||
class HunyuanImage21(HunyuanVideo):
|
||||
unet_config = {
|
||||
"image_model": "hunyuan_video",
|
||||
@@ -2415,14 +2348,12 @@ models = [
|
||||
HiDream,
|
||||
HiDreamO1,
|
||||
Chroma,
|
||||
SeedVR2,
|
||||
ChromaRadiance,
|
||||
ACEStep,
|
||||
ACEStep15,
|
||||
Omnigen2,
|
||||
Boogu,
|
||||
QwenImage,
|
||||
JoyImage,
|
||||
Ideogram4,
|
||||
Krea2,
|
||||
Flux2,
|
||||
|
||||
@@ -54,13 +54,13 @@ class BASE:
|
||||
optimizations = {"fp8": False}
|
||||
|
||||
@classmethod
|
||||
def matches(s, unet_config, state_dict=None, unet_key_prefix=""):
|
||||
def matches(s, unet_config, state_dict=None):
|
||||
for k in s.unet_config:
|
||||
if k not in unet_config or s.unet_config[k] != unet_config[k]:
|
||||
return False
|
||||
if state_dict is not None:
|
||||
for k in s.required_keys:
|
||||
if k.format(unet_key_prefix) not in state_dict:
|
||||
if k not in state_dict:
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -115,7 +115,7 @@ class BASE:
|
||||
replace_prefix = {"": self.vae_key_prefix[0]}
|
||||
return utils.state_dict_prefix_replace(state_dict, replace_prefix)
|
||||
|
||||
def set_inference_dtype(self, dtype, manual_cast_dtype, device=None):
|
||||
def set_inference_dtype(self, dtype, manual_cast_dtype):
|
||||
self.unet_config['dtype'] = dtype
|
||||
self.manual_cast_dtype = manual_cast_dtype
|
||||
|
||||
|
||||
@@ -1088,7 +1088,7 @@ class Gemma4_Tokenizer():
|
||||
h, w = samples.shape[2], samples.shape[3]
|
||||
patch_size = 16
|
||||
pooling_k = 3
|
||||
max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280)
|
||||
max_soft_tokens = 70 if is_video else 280 # video uses smaller token budget per frame
|
||||
max_patches = max_soft_tokens * pooling_k * pooling_k
|
||||
target_px = max_patches * patch_size * patch_size
|
||||
factor = (target_px / (h * w)) ** 0.5
|
||||
|
||||
@@ -1,97 +0,0 @@
|
||||
import torch
|
||||
|
||||
from comfy import sd1_clip
|
||||
import comfy.text_encoders.qwen_vl
|
||||
from comfy.text_encoders.qwen3vl import Qwen3VL, Qwen3VLTokenizer
|
||||
|
||||
JOYIMAGE_VISION_BLOCK = "<|vision_start|><|image_pad|><|vision_end|>"
|
||||
JOYIMAGE_TEMPLATE_TEXT = (
|
||||
"<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|im_end|>\n"
|
||||
"<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
|
||||
)
|
||||
JOYIMAGE_TEMPLATE_IMAGE = (
|
||||
"<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|im_end|>\n"
|
||||
f"<|im_start|>user\n{JOYIMAGE_VISION_BLOCK}{{}}<|im_end|>\n<|im_start|>assistant\n"
|
||||
)
|
||||
# The DiT was trained without the leading system-prompt tokens.
|
||||
JOYIMAGE_DROP_IDX = 34
|
||||
PAD_TOKEN = 151643
|
||||
|
||||
|
||||
class Qwen3VL8B_JoyImage(Qwen3VL):
|
||||
model_type = "qwen3vl_8b"
|
||||
|
||||
def preprocess_embed(self, embed, device):
|
||||
if embed["type"] == "image":
|
||||
image, grid = comfy.text_encoders.qwen_vl.process_qwen2vl_images(
|
||||
embed["data"], min_pixels=65536, max_pixels=16777216, patch_size=16,
|
||||
image_mean=[0.5, 0.5, 0.5], image_std=[0.5, 0.5, 0.5],
|
||||
interpolation="bicubic",
|
||||
)
|
||||
merged, deepstack = self.visual(image.to(device, dtype=torch.float32), grid)
|
||||
return merged, {"grid": grid, "deepstack": deepstack}
|
||||
return None, None
|
||||
|
||||
|
||||
class JoyImageTokenizer(Qwen3VLTokenizer):
|
||||
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
||||
super().__init__(
|
||||
embedding_directory=embedding_directory, tokenizer_data=tokenizer_data,
|
||||
model_type="qwen3vl_8b",
|
||||
)
|
||||
self.llama_template = JOYIMAGE_TEMPLATE_TEXT
|
||||
self.llama_template_images = JOYIMAGE_TEMPLATE_IMAGE
|
||||
|
||||
def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=None, **kwargs):
|
||||
kwargs.pop("thinking", None)
|
||||
return super().tokenize_with_weights(
|
||||
text, return_word_ids=return_word_ids, llama_template=llama_template,
|
||||
images=images or [], thinking=True, **kwargs,
|
||||
)
|
||||
|
||||
|
||||
class _JoyImageClipModel(sd1_clip.SDClipModel):
|
||||
def __init__(self, device="cpu", layer="hidden", layer_idx=-1, dtype=None,
|
||||
attention_mask=True, model_options={}):
|
||||
super().__init__(
|
||||
device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={},
|
||||
# JoyImage conditions on the pre-final-norm output of the last decoder layer.
|
||||
dtype=dtype, special_tokens={"pad": PAD_TOKEN}, layer_norm_hidden_state=False,
|
||||
model_class=Qwen3VL8B_JoyImage, enable_attention_masks=attention_mask,
|
||||
return_attention_masks=attention_mask, model_options=model_options,
|
||||
)
|
||||
|
||||
|
||||
class JoyImageTEModel(sd1_clip.SD1ClipModel):
|
||||
def __init__(self, device="cpu", dtype=None, model_options={}):
|
||||
super().__init__(
|
||||
device=device, dtype=dtype, name="qwen3vl_8b",
|
||||
clip_model=_JoyImageClipModel, model_options=model_options,
|
||||
)
|
||||
|
||||
def encode_token_weights(self, token_weight_pairs):
|
||||
out, pooled, extra = super().encode_token_weights(token_weight_pairs)
|
||||
if out.shape[1] <= JOYIMAGE_DROP_IDX:
|
||||
raise ValueError(
|
||||
f"JoyImageTEModel: encoded sequence length {out.shape[1]} is shorter "
|
||||
f"than drop_idx={JOYIMAGE_DROP_IDX}; the prompt did not include the "
|
||||
f"template prefix."
|
||||
)
|
||||
out = out[:, JOYIMAGE_DROP_IDX:]
|
||||
if "attention_mask" in extra:
|
||||
extra["attention_mask"] = extra["attention_mask"][:, JOYIMAGE_DROP_IDX:]
|
||||
return out, pooled, extra
|
||||
|
||||
|
||||
def te(dtype_llama=None, llama_quantization_metadata=None):
|
||||
class JoyImageTEModel_(JoyImageTEModel):
|
||||
def __init__(self, device="cpu", dtype=None, model_options={}):
|
||||
if llama_quantization_metadata is not None:
|
||||
model_options = model_options.copy()
|
||||
model_options["quantization_metadata"] = llama_quantization_metadata
|
||||
if dtype_llama is not None:
|
||||
dtype = dtype_llama
|
||||
super().__init__(device=device, dtype=dtype, model_options=model_options)
|
||||
return JoyImageTEModel_
|
||||
@@ -15,7 +15,6 @@ def process_qwen2vl_images(
|
||||
merge_size: int = 2,
|
||||
image_mean: list = None,
|
||||
image_std: list = None,
|
||||
interpolation: str = "bilinear",
|
||||
):
|
||||
if image_mean is None:
|
||||
image_mean = [0.48145466, 0.4578275, 0.40821073]
|
||||
@@ -48,9 +47,10 @@ def process_qwen2vl_images(
|
||||
img_resized = F.interpolate(
|
||||
img.unsqueeze(0),
|
||||
size=(h_bar, w_bar),
|
||||
mode=interpolation,
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
).squeeze(0)
|
||||
|
||||
normalized = img_resized.clone()
|
||||
for c in range(3):
|
||||
normalized[c] = (img_resized[c] - image_mean[c]) / image_std[c]
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from av.container import InputContainer
|
||||
from av.subtitles.stream import SubtitleStream
|
||||
from av.video.reformatter import ColorRange
|
||||
from fractions import Fraction
|
||||
from typing import Optional
|
||||
from .._input import AudioInput, VideoInput
|
||||
@@ -10,7 +9,6 @@ import itertools
|
||||
import json
|
||||
import numpy as np
|
||||
import math
|
||||
import os
|
||||
import torch
|
||||
from .._util import VideoContainer, VideoCodec, VideoComponents
|
||||
import logging
|
||||
@@ -60,57 +58,6 @@ def video_stream_bit_depth(stream) -> int:
|
||||
return max(component.bits for component in stream.format.components)
|
||||
|
||||
|
||||
def last_decodable_audio_stream(container: InputContainer):
|
||||
"""Streams FFmpeg has no decoder for have no codec context, and decoding their
|
||||
packets crashes the process (e.g. APAC spatial-audio track in iPhone)."""
|
||||
stream = next(
|
||||
(s for s in reversed(container.streams.audio) if s.codec_context is not None),
|
||||
None,
|
||||
)
|
||||
if stream is None and len(container.streams.audio):
|
||||
logging.warning("No decodable audio stream found in video; ignoring audio.")
|
||||
return stream
|
||||
|
||||
|
||||
def probe_audio_params(container: InputContainer, audio_stream, max_packets: int = 200):
|
||||
"""Containers probed only up to a window (mpegts) leave audio codec parameters unset when
|
||||
audio starts beyond it; learn them by decoding ahead. The caller must seek back afterwards.
|
||||
Returns (sample_rate, channels), zeros when the stream never yields a decodable frame."""
|
||||
for i, packet in enumerate(container.demux(audio_stream)):
|
||||
try:
|
||||
frames = packet.decode()
|
||||
except av.error.FFmpegError:
|
||||
frames = ()
|
||||
if frames:
|
||||
return frames[0].sample_rate, frames[0].layout.nb_channels
|
||||
if i >= max_packets:
|
||||
break
|
||||
return 0, 0
|
||||
|
||||
|
||||
def write_output_metadata(container: InputContainer, output, metadata: dict | None):
|
||||
"""Copy the source container's metadata, then overlay the caller's tags."""
|
||||
for key, value in container.metadata.items():
|
||||
if metadata is None or key not in metadata:
|
||||
output.metadata[key] = value
|
||||
if metadata is not None:
|
||||
for key, value in metadata.items():
|
||||
output.metadata[key] = value if isinstance(value, str) else json.dumps(value)
|
||||
|
||||
|
||||
def mp4_output_open_kwargs(path: str | io.BytesIO, format: VideoContainer, codec: VideoCodec) -> dict:
|
||||
if format != VideoContainer.AUTO and format != VideoContainer.MP4:
|
||||
raise ValueError("Only MP4 format is supported for now")
|
||||
if codec != VideoCodec.AUTO and codec != VideoCodec.H264:
|
||||
raise ValueError("Only H264 codec is supported for now")
|
||||
open_kwargs = {"mode": "w", "options": {"movflags": "use_metadata_tags"}}
|
||||
if isinstance(format, VideoContainer) and format != VideoContainer.AUTO:
|
||||
open_kwargs["format"] = format.value
|
||||
elif isinstance(path, io.BytesIO):
|
||||
open_kwargs["format"] = "mp4" # no file extension to infer the format from
|
||||
return open_kwargs
|
||||
|
||||
|
||||
class VideoFromFile(VideoInput):
|
||||
"""
|
||||
Class representing video input from a file.
|
||||
@@ -245,10 +192,13 @@ class VideoFromFile(VideoInput):
|
||||
return estimated_frames
|
||||
|
||||
# 3. Last resort: decode frames and count them (streaming)
|
||||
start_time, duration = self.get_active_trim_window()
|
||||
if self.__start_time < 0:
|
||||
start_time = max(self._get_raw_duration() + self.__start_time, 0)
|
||||
else:
|
||||
start_time = self.__start_time
|
||||
frame_count = 1
|
||||
start_pts = int(start_time / video_stream.time_base)
|
||||
end_pts = int((start_time + duration) / video_stream.time_base)
|
||||
end_pts = int((start_time + self.__duration) / video_stream.time_base)
|
||||
container.seek(start_pts, stream=video_stream)
|
||||
frame_iterator = (
|
||||
container.decode(video_stream)
|
||||
@@ -303,14 +253,17 @@ class VideoFromFile(VideoInput):
|
||||
|
||||
def get_components_internal(self, container: InputContainer) -> VideoComponents:
|
||||
video_stream = self._get_first_video_stream(container)
|
||||
start_time, duration = self.get_active_trim_window()
|
||||
if self.__start_time < 0:
|
||||
start_time = max(self._get_raw_duration() + self.__start_time, 0)
|
||||
else:
|
||||
start_time = self.__start_time
|
||||
|
||||
# Get video frames
|
||||
frames = []
|
||||
audio_frames = []
|
||||
alphas = None
|
||||
start_pts = int(start_time / video_stream.time_base)
|
||||
end_pts = int((start_time + duration) / video_stream.time_base)
|
||||
end_pts = int((start_time + self.__duration) / video_stream.time_base)
|
||||
|
||||
if start_pts != 0:
|
||||
container.seek(start_pts, stream=video_stream)
|
||||
@@ -328,11 +281,18 @@ class VideoFromFile(VideoInput):
|
||||
video_done = False
|
||||
audio_done = True
|
||||
|
||||
audio_stream = last_decodable_audio_stream(container)
|
||||
# Use the last decodable audio stream. Streams FFmpeg has no decoder for have no codec context,
|
||||
# and decoding their packets crashes the process. (e.g. APAC spatial-audio track in iPhone)
|
||||
audio_stream = next(
|
||||
(s for s in reversed(container.streams.audio) if s.codec_context is not None),
|
||||
None,
|
||||
)
|
||||
if audio_stream is not None:
|
||||
streams += [audio_stream]
|
||||
resampler = av.audio.resampler.AudioResampler(format='fltp')
|
||||
audio_done = False
|
||||
elif len(container.streams.audio):
|
||||
logging.warning("No decodable audio stream found in video; ignoring audio.")
|
||||
|
||||
for packet in container.demux(*streams):
|
||||
if video_done and audio_done:
|
||||
@@ -345,7 +305,7 @@ class VideoFromFile(VideoInput):
|
||||
for frame in packet.decode():
|
||||
if frame.pts < start_pts:
|
||||
continue
|
||||
if duration and frame.pts >= end_pts:
|
||||
if self.__duration and frame.pts >= end_pts:
|
||||
video_done = True
|
||||
break
|
||||
|
||||
@@ -412,7 +372,7 @@ class VideoFromFile(VideoInput):
|
||||
map(resampler.resample, packet.decode())
|
||||
)
|
||||
for frame in aframes:
|
||||
if duration and frame.time > start_time + duration:
|
||||
if self.__duration and frame.time > start_time + self.__duration:
|
||||
audio_done = True
|
||||
break
|
||||
|
||||
@@ -434,8 +394,8 @@ class VideoFromFile(VideoInput):
|
||||
|
||||
if len(audio_frames) > 0:
|
||||
audio_data = np.concatenate(audio_frames, axis=1) # shape: (channels, total_samples)
|
||||
if duration:
|
||||
audio_data = audio_data[..., :int(duration * audio_stream.sample_rate)]
|
||||
if self.__duration:
|
||||
audio_data = audio_data[..., :int(self.__duration * audio_stream.sample_rate)]
|
||||
|
||||
audio_tensor = torch.from_numpy(audio_data).unsqueeze(0) # shape: (1, channels, total_samples)
|
||||
audio = AudioInput({
|
||||
@@ -481,14 +441,28 @@ class VideoFromFile(VideoInput):
|
||||
if not reuse_streams:
|
||||
if bit_depth is None:
|
||||
bit_depth = source_bit_depth
|
||||
return self._save_transcoded(container, path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth)
|
||||
components = self.get_components_internal(container)
|
||||
video = VideoFromComponents(components)
|
||||
return video.save_to(
|
||||
path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth,
|
||||
)
|
||||
|
||||
streams = container.streams
|
||||
|
||||
open_kwargs = get_open_write_kwargs(path, container_format, format)
|
||||
with av.open(path, **open_kwargs) as output_container:
|
||||
# Add metadata before writing any streams
|
||||
write_output_metadata(container, output_container, metadata)
|
||||
# Copy over the original metadata
|
||||
for key, value in container.metadata.items():
|
||||
if metadata is None or key not in metadata:
|
||||
output_container.metadata[key] = value
|
||||
|
||||
# Add our new metadata
|
||||
if metadata is not None:
|
||||
for key, value in metadata.items():
|
||||
if isinstance(value, str):
|
||||
output_container.metadata[key] = value
|
||||
else:
|
||||
output_container.metadata[key] = json.dumps(value)
|
||||
|
||||
# Add streams to the new container. Streams with no codec context cannot be used as an output template.
|
||||
stream_map = {}
|
||||
@@ -506,282 +480,6 @@ class VideoFromFile(VideoInput):
|
||||
packet.stream = stream_map[packet.stream]
|
||||
output_container.mux(packet)
|
||||
|
||||
def _save_transcoded(
|
||||
self,
|
||||
container: InputContainer,
|
||||
path: str | io.BytesIO,
|
||||
format: VideoContainer,
|
||||
codec: VideoCodec,
|
||||
metadata: dict | None,
|
||||
bit_depth: int,
|
||||
):
|
||||
"""Re-encode to H.264/AAC one frame at a time; peak memory does not scale with video length."""
|
||||
open_kwargs = mp4_output_open_kwargs(path, format, codec)
|
||||
video_stream = self._get_first_video_stream(container)
|
||||
start_time, duration = self.get_active_trim_window()
|
||||
start_pts = int(start_time / video_stream.time_base)
|
||||
end_pts = int((start_time + duration) / video_stream.time_base) if duration else None
|
||||
stream_end_pts = None
|
||||
if video_stream.duration is not None:
|
||||
stream_end_pts = (video_stream.start_time or 0) + video_stream.duration
|
||||
output_end_pts = end_pts
|
||||
if stream_end_pts is not None and (output_end_pts is None or stream_end_pts < output_end_pts):
|
||||
output_end_pts = stream_end_pts
|
||||
if start_pts != 0:
|
||||
container.seek(start_pts, stream=video_stream)
|
||||
|
||||
audio_stream = last_decodable_audio_stream(container)
|
||||
pix_fmt = "yuv420p10le" if bit_depth >= 10 else "yuv420p"
|
||||
rate = Fraction(video_stream.average_rate) if video_stream.average_rate else Fraction(1)
|
||||
|
||||
resampler = None
|
||||
sample_rate = 0
|
||||
audio_time_base = None
|
||||
duration_cap = None
|
||||
if audio_stream is not None:
|
||||
sample_rate = audio_stream.codec_context.sample_rate
|
||||
channels = audio_stream.codec_context.channels
|
||||
if not sample_rate:
|
||||
sample_rate, channels = probe_audio_params(container, audio_stream)
|
||||
container.seek(start_pts, stream=video_stream)
|
||||
if sample_rate:
|
||||
audio_stream.codec_context.flush_buffers()
|
||||
else:
|
||||
logging.warning("Audio stream parameters could not be determined; ignoring audio.")
|
||||
audio_stream = None
|
||||
if audio_stream is not None:
|
||||
audio_time_base = Fraction(1, sample_rate)
|
||||
layout = {1: "mono", 2: "stereo", 6: "5.1"}.get(channels, "stereo")
|
||||
resampler = av.audio.resampler.AudioResampler(format="fltp", layout=layout, rate=sample_rate)
|
||||
if duration:
|
||||
duration_cap = math.ceil(duration * sample_rate)
|
||||
|
||||
streams = [video_stream] if audio_stream is None else [video_stream, audio_stream]
|
||||
pts_step = max(1, int(round((1 / rate) / video_stream.time_base)))
|
||||
video_done = False
|
||||
audio_done = audio_stream is None
|
||||
video_pts_offset = None
|
||||
last_video_pts = None
|
||||
last_video_end = None
|
||||
# rebased pts -> true display duration: the mp4 muxer pads the last sample with 1/rate otherwise
|
||||
video_frame_durations = {}
|
||||
source_size = None
|
||||
rotation_k = 0
|
||||
rotation_filter = None
|
||||
audio_started = False
|
||||
samples_written = 0
|
||||
pending_audio = []
|
||||
# The output opens lazily on the first kept frame: it decides the geometry (90/270 rotation swaps dims),
|
||||
# and never seeking back keeps webm/mkv leading audio intact.
|
||||
output = None
|
||||
out_video = None
|
||||
out_audio = None
|
||||
|
||||
def audio_frame_from_ndarray(nd_planar):
|
||||
frame = av.AudioFrame.from_ndarray(np.ascontiguousarray(nd_planar), format="fltp", layout=layout)
|
||||
frame.sample_rate = sample_rate
|
||||
return frame
|
||||
|
||||
def drain_audio(final=False):
|
||||
# Audio may cover the pts span of the video written so far, capped by the requested duration
|
||||
nonlocal samples_written, audio_done
|
||||
if last_video_end is None:
|
||||
cap = 0
|
||||
else:
|
||||
cap = math.ceil(last_video_end * video_stream.time_base * sample_rate)
|
||||
if duration_cap is not None:
|
||||
cap = min(cap, duration_cap)
|
||||
while pending_audio and not audio_done:
|
||||
frame = pending_audio[0]
|
||||
if samples_written + frame.samples <= cap:
|
||||
frame.pts = samples_written
|
||||
frame.time_base = audio_time_base
|
||||
output.mux(out_audio.encode(frame))
|
||||
samples_written += frame.samples
|
||||
pending_audio.pop(0)
|
||||
continue
|
||||
if final:
|
||||
keep = frame.to_ndarray()[..., :cap - samples_written]
|
||||
if keep.shape[-1] > 0:
|
||||
tail = audio_frame_from_ndarray(keep)
|
||||
tail.pts = samples_written
|
||||
tail.time_base = audio_time_base
|
||||
output.mux(out_audio.encode(tail))
|
||||
samples_written += keep.shape[-1]
|
||||
pending_audio.clear()
|
||||
break
|
||||
if duration_cap is not None and samples_written >= duration_cap:
|
||||
audio_done = True
|
||||
return cap
|
||||
|
||||
try:
|
||||
for packet in container.demux(*streams):
|
||||
if video_done and audio_done:
|
||||
break
|
||||
|
||||
if packet.stream == video_stream and not video_done:
|
||||
try:
|
||||
frames = packet.decode()
|
||||
except av.error.InvalidDataError:
|
||||
logging.info("pyav decode error")
|
||||
continue
|
||||
for frame in frames:
|
||||
if frame.pts is not None and frame.pts < start_pts:
|
||||
continue
|
||||
if end_pts is not None and frame.pts is not None and frame.pts >= end_pts:
|
||||
video_done = True
|
||||
if last_video_pts is not None:
|
||||
# the source continues past the window: hold the last kept frame to the window end
|
||||
end_offset = video_pts_offset if video_pts_offset is not None else start_pts
|
||||
last_video_end = max(last_video_end, end_pts - end_offset)
|
||||
break
|
||||
# the source's true display duration of this frame; average_rate is not a
|
||||
# frame duration (sparse/VFR sources), so it is only the fallback
|
||||
frame_duration = frame.duration if frame.duration else pts_step
|
||||
if end_pts is not None and frame.pts is not None:
|
||||
frame_duration = min(frame_duration, end_pts - frame.pts)
|
||||
if output is None:
|
||||
rotation_k = int(round(frame.rotation // 90)) % 4 if frame.rotation else 0
|
||||
if rotation_k % 2:
|
||||
out_width, out_height = frame.height, frame.width
|
||||
else:
|
||||
out_width, out_height = frame.width, frame.height
|
||||
if out_width % 2 or out_height % 2:
|
||||
raise ValueError(f"H.264 output requires even dimensions, got {out_width}x{out_height}")
|
||||
source_size = (frame.width, frame.height)
|
||||
output = av.open(path, **open_kwargs)
|
||||
# Add metadata before writing any streams
|
||||
write_output_metadata(container, output, metadata)
|
||||
out_video = output.add_stream("h264", rate=rate)
|
||||
# no B-frames: reordering makes mp4 sample durations follow decode order,
|
||||
# so irregular-VFR spans and trim windows land wrong
|
||||
out_video.codec_context.max_b_frames = 0
|
||||
out_video.width = out_width
|
||||
out_video.height = out_height
|
||||
out_video.pix_fmt = pix_fmt
|
||||
# source pts pass through (rebased to 0), so variable frame rate survives
|
||||
out_video.codec_context.time_base = video_stream.time_base
|
||||
if audio_stream is not None:
|
||||
out_audio = output.add_stream("aac", rate=sample_rate, layout=layout)
|
||||
if (frame.width, frame.height) != source_size:
|
||||
# encoding would silently rescale the new geometry into the old one
|
||||
raise ValueError(
|
||||
f"Video resolution changes mid-stream "
|
||||
f"({source_size[0]}x{source_size[1]} -> {frame.width}x{frame.height}); cannot transcode"
|
||||
)
|
||||
if rotation_k:
|
||||
if rotation_filter is None:
|
||||
g = av.filter.Graph()
|
||||
g_src = g.add_buffer(width=frame.width, height=frame.height,
|
||||
format=frame.format.name, time_base=video_stream.time_base)
|
||||
tail = g_src
|
||||
for filter_name, filter_args in {1: [("transpose", "cclock")],
|
||||
2: [("hflip", None), ("vflip", None)],
|
||||
3: [("transpose", "clock")]}[rotation_k]:
|
||||
step = g.add(filter_name, filter_args)
|
||||
tail.link_to(step)
|
||||
tail = step
|
||||
g_sink = g.add("buffersink")
|
||||
tail.link_to(g_sink)
|
||||
g.configure()
|
||||
rotation_filter = (g_src, g_sink)
|
||||
rotation_filter[0].push(frame)
|
||||
frame = rotation_filter[1].pull()
|
||||
if frame.color_range == ColorRange.JPEG:
|
||||
# compress full-range sources (yuvj/MJPEG) to limited range
|
||||
frame = frame.reformat(format=pix_fmt, src_color_range="JPEG", dst_color_range="MPEG")
|
||||
else:
|
||||
frame = frame.reformat(format=pix_fmt)
|
||||
frame_output_end = None
|
||||
if frame.pts is not None:
|
||||
if video_pts_offset is None:
|
||||
video_pts_offset = frame.pts
|
||||
frame.pts -= video_pts_offset
|
||||
if output_end_pts is not None:
|
||||
frame_output_end = output_end_pts - video_pts_offset
|
||||
if frame.pts + frame_duration > frame_output_end:
|
||||
clamped_pts = frame_output_end - frame_duration
|
||||
if clamped_pts >= 0 and (last_video_pts is None or clamped_pts > last_video_pts):
|
||||
frame.pts = min(frame.pts, clamped_pts)
|
||||
elif frame.pts < frame_output_end:
|
||||
frame_duration = frame_output_end - frame.pts
|
||||
else:
|
||||
continue
|
||||
if frame.pts is None or (last_video_pts is not None and frame.pts <= last_video_pts):
|
||||
# broken sources emit missing/backward timestamps mid-stream, which the
|
||||
# muxer rejects; nudge them forward by one nominal frame interval
|
||||
frame.pts = 0 if last_video_pts is None else last_video_pts + pts_step
|
||||
if frame_output_end is not None and frame.pts + frame_duration > frame_output_end:
|
||||
if frame.pts >= frame_output_end:
|
||||
continue
|
||||
frame_duration = frame_output_end - frame.pts
|
||||
last_video_pts = frame.pts
|
||||
last_video_end = frame.pts + frame_duration
|
||||
video_frame_durations[frame.pts] = frame_duration
|
||||
# the decoded pict_type would force x264's frame types (intra-only
|
||||
# sources like MJPEG/ProRes would come out all-keyframe)
|
||||
frame.pict_type = 0
|
||||
for out_packet in out_video.encode(frame):
|
||||
out_packet.duration = video_frame_durations.pop(out_packet.pts, 0)
|
||||
output.mux(out_packet)
|
||||
drain_audio()
|
||||
|
||||
elif packet.stream == audio_stream and not audio_done:
|
||||
for resampled in itertools.chain.from_iterable(map(resampler.resample, packet.decode())):
|
||||
frame_start = None
|
||||
if resampled.pts is not None:
|
||||
# passthrough frames keep the source stream's time base
|
||||
tb = resampled.time_base if resampled.time_base else audio_time_base
|
||||
frame_start = float(resampled.pts * tb)
|
||||
if duration and not audio_started and frame_start >= start_time + duration:
|
||||
audio_done = True
|
||||
break
|
||||
if not audio_started:
|
||||
if frame_start is None:
|
||||
frame_start = 0.0
|
||||
to_skip = max(0, int((start_time - frame_start) * sample_rate))
|
||||
if to_skip >= resampled.samples:
|
||||
continue
|
||||
audio_started = True
|
||||
if duration and frame_start > start_time:
|
||||
duration_cap = min(duration_cap, math.ceil((start_time + duration - frame_start) * sample_rate))
|
||||
if to_skip:
|
||||
pending_audio.append(audio_frame_from_ndarray(resampled.to_ndarray()[..., to_skip:]))
|
||||
continue
|
||||
pending_audio.append(resampled)
|
||||
if video_done:
|
||||
# the video window is complete so the cap is final, but containers
|
||||
# that interleave audio behind video (fragmented mp4) still owe most
|
||||
# of it: stop only once the demuxed audio covers the cap
|
||||
cap = drain_audio()
|
||||
if pending_audio or samples_written >= cap:
|
||||
drain_audio(final=True)
|
||||
audio_done = True
|
||||
break
|
||||
|
||||
if output is None:
|
||||
raise ValueError(f"No decodable video frames found in file '{self.__file}'")
|
||||
if out_audio is not None and not audio_done:
|
||||
drain_audio(final=True)
|
||||
window_fill = last_video_end - last_video_pts if video_done and last_video_pts is not None else 0
|
||||
for out_packet in out_video.encode(None):
|
||||
duration = video_frame_durations.pop(out_packet.pts, 0)
|
||||
if out_packet.pts == last_video_pts:
|
||||
duration = max(duration, window_fill)
|
||||
out_packet.duration = duration
|
||||
output.mux(out_packet)
|
||||
if out_audio is not None:
|
||||
output.mux(out_audio.encode(None))
|
||||
except BaseException:
|
||||
if output is not None:
|
||||
output.close()
|
||||
if isinstance(path, (str, os.PathLike)) and os.path.exists(path):
|
||||
os.remove(path)
|
||||
raise
|
||||
else:
|
||||
if output is not None:
|
||||
output.close()
|
||||
|
||||
def _get_first_video_stream(self, container: InputContainer):
|
||||
if len(container.streams.video):
|
||||
return container.streams.video[0]
|
||||
@@ -829,12 +527,22 @@ class VideoFromComponents(VideoInput):
|
||||
bit_depth: int | None = None,
|
||||
):
|
||||
"""Save the video to a file path or BytesIO buffer."""
|
||||
open_kwargs = mp4_output_open_kwargs(path, format, codec)
|
||||
if format != VideoContainer.AUTO and format != VideoContainer.MP4:
|
||||
raise ValueError("Only MP4 format is supported for now")
|
||||
if codec != VideoCodec.AUTO and codec != VideoCodec.H264:
|
||||
raise ValueError("Only H264 codec is supported for now")
|
||||
# None means "use the depth this video was created with" (CreateVideo's choice).
|
||||
if bit_depth is None:
|
||||
bit_depth = self.__bit_depth
|
||||
is_10bit = bit_depth >= 10
|
||||
with av.open(path, **open_kwargs) as output:
|
||||
extra_kwargs = {}
|
||||
if isinstance(format, VideoContainer) and format != VideoContainer.AUTO:
|
||||
extra_kwargs["format"] = format.value
|
||||
elif isinstance(path, io.BytesIO):
|
||||
# BytesIO has no file extension, so av.open can't infer the format.
|
||||
# Default to mp4 since that's the only supported format anyway.
|
||||
extra_kwargs["format"] = "mp4"
|
||||
with av.open(path, mode='w', options={'movflags': 'use_metadata_tags'}, **extra_kwargs) as output:
|
||||
# Add metadata before writing any streams
|
||||
if metadata is not None:
|
||||
for key, value in metadata.items():
|
||||
|
||||
@@ -1261,6 +1261,155 @@ class DynamicSlot(ComfyTypeI):
|
||||
out_dict[input_type][finalized_id] = value
|
||||
out_dict["dynamic_paths"][finalized_id] = finalize_prefix(curr_prefix, curr_prefix[-1])
|
||||
|
||||
@comfytype(io_type="COMFY_DYNAMICGROUP_V3")
|
||||
class DynamicGroup(ComfyTypeI):
|
||||
"""A repeatable group of widget inputs (e.g. lora_name + strength stacked into N rows).
|
||||
|
||||
At execution time the node receives a ``list[dict]`` where each element is a row.
|
||||
|
||||
Example::
|
||||
|
||||
io.DynamicGroup.Input(
|
||||
"loras",
|
||||
template=[
|
||||
io.Combo.Input("lora_name", options=folder_paths.get_filename_list("loras")),
|
||||
io.Float.Input("strength", default=1.0, min=-100, max=100, step=0.01),
|
||||
],
|
||||
min=0,
|
||||
max=50,
|
||||
)
|
||||
# execute receives: loras: list[dict] = [{"lora_name": "x.safetensors", "strength": 1.0}, ...]
|
||||
"""
|
||||
|
||||
Type = list[dict[str, Any]]
|
||||
_MaxRows = 100
|
||||
|
||||
class Input(DynamicInput):
|
||||
def __init__(
|
||||
self,
|
||||
id: str,
|
||||
template: list["Input"],
|
||||
min: int = 0,
|
||||
max: int = 50,
|
||||
display_name: str = None,
|
||||
optional: bool = False,
|
||||
tooltip: str = None,
|
||||
lazy: bool = None,
|
||||
extra_dict=None,
|
||||
group_name: str = "Group",
|
||||
):
|
||||
super().__init__(id, display_name, optional, tooltip, lazy, extra_dict)
|
||||
assert len(template) > 0, "DynamicGroup template must have at least one field."
|
||||
for t in template:
|
||||
assert isinstance(t, WidgetInput), (
|
||||
f"DynamicGroup template field '{t.id}' must be a WidgetInput subclass "
|
||||
f"(Combo, Float, Int, String, Boolean, Color). Got {type(t).__name__}."
|
||||
)
|
||||
assert not isinstance(t, DynamicInput), (
|
||||
f"DynamicGroup template field '{t.id}' must not be a DynamicInput. "
|
||||
"Nesting dynamic inputs inside DynamicGroup is not supported."
|
||||
)
|
||||
field_ids = [t.id for t in template]
|
||||
assert len(field_ids) == len(set(field_ids)), (
|
||||
f"DynamicGroup template field ids must be unique within a row. Got: {field_ids}"
|
||||
)
|
||||
# Reject "." in group id and template field ids: slot_id encoding uses "." as a
|
||||
# delimiter (<group_id>.<row>.<field_id>), so any "." in these names would cause
|
||||
# path.split(".") to produce the wrong number of segments during decoding.
|
||||
assert "." not in id, (
|
||||
f"DynamicGroup id must not contain '.'. Got: '{id}'"
|
||||
)
|
||||
for t in template:
|
||||
assert "." not in t.id, (
|
||||
f"DynamicGroup template field id must not contain '.'. Got: '{t.id}'"
|
||||
)
|
||||
assert min >= 0, "DynamicGroup min must be >= 0."
|
||||
assert max >= 1, "DynamicGroup max must be >= 1."
|
||||
assert max <= DynamicGroup._MaxRows, f"DynamicGroup max must be <= {DynamicGroup._MaxRows}."
|
||||
assert min <= max, "DynamicGroup min must be <= max."
|
||||
self.template = template
|
||||
self.min = min
|
||||
self.max = max
|
||||
self.group_name = group_name
|
||||
|
||||
def get_all(self) -> list["Input"]:
|
||||
return [self] + list(self.template)
|
||||
|
||||
def as_dict(self):
|
||||
return super().as_dict() | prune_dict({
|
||||
"template": create_input_dict_v1(self.template),
|
||||
"min": self.min,
|
||||
"max": self.max,
|
||||
"group_name": self.group_name,
|
||||
})
|
||||
|
||||
def validate(self):
|
||||
for t in self.template:
|
||||
t.validate()
|
||||
|
||||
@staticmethod
|
||||
def _expand_schema_for_dynamic(
|
||||
out_dict: dict[str, Any],
|
||||
live_inputs: dict[str, Any],
|
||||
value: tuple[str, dict[str, Any]],
|
||||
input_type: str,
|
||||
curr_prefix: list[str] | None,
|
||||
):
|
||||
info = value[1]
|
||||
min_rows: int = info.get("min", 0)
|
||||
max_rows: int = info.get("max", DynamicGroup._MaxRows)
|
||||
template: dict[str, Any] = info.get("template", {})
|
||||
|
||||
# Collect all template field specs across required/optional sections
|
||||
field_specs: list[tuple[str, tuple[str, dict[str, Any]], bool]] = []
|
||||
for field_required_key in ("required", "optional"):
|
||||
section = template.get(field_required_key, {})
|
||||
is_required_field = field_required_key == "required"
|
||||
for field_id, field_value in section.items():
|
||||
field_specs.append((field_id, field_value, is_required_field))
|
||||
|
||||
# Determine how many rows are currently present by scanning live_inputs
|
||||
finalized_prefix = finalize_prefix(curr_prefix)
|
||||
present_rows = 0
|
||||
for live_key in live_inputs:
|
||||
# Keys look like "<prefix>.<row>.<field_id>"
|
||||
if live_key.startswith(finalized_prefix + "."):
|
||||
remainder = live_key[len(finalized_prefix) + 1:]
|
||||
parts = remainder.split(".", 1)
|
||||
if len(parts) >= 1:
|
||||
try:
|
||||
row_idx = int(parts[0])
|
||||
present_rows = max(present_rows, row_idx + 1)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if present_rows > max_rows:
|
||||
raise ValueError(
|
||||
f"DynamicGroup input '{finalized_prefix}' received {present_rows} rows but max is {max_rows}."
|
||||
)
|
||||
row_count = max(min_rows, present_rows)
|
||||
|
||||
for row in range(row_count):
|
||||
for field_id, field_value, is_required_field in field_specs:
|
||||
slot_id = f"{finalized_prefix}.{row}.{field_id}"
|
||||
if row < min_rows and is_required_field:
|
||||
out_dict["required"][slot_id] = field_value
|
||||
else:
|
||||
out_dict["optional"][slot_id] = field_value
|
||||
# Register into dynamic_paths so build_nested_inputs places value at the right path
|
||||
out_dict["dynamic_paths"][slot_id] = slot_id
|
||||
|
||||
# Track the list root path so build_nested_inputs can convert the index dict to a list
|
||||
out_dict.setdefault("list_paths", set()).add(finalized_prefix)
|
||||
|
||||
# Handle the empty case (0 rows) – emit an empty-list default for the parent.
|
||||
# This must only fire when there are genuinely no rows; otherwise the parent
|
||||
# path would clobber the per-row dict built from the slot ids above.
|
||||
if row_count == 0:
|
||||
out_dict["dynamic_paths"][finalized_prefix] = finalized_prefix
|
||||
out_dict["dynamic_paths_default_value"][finalized_prefix] = DynamicPathsDefaultValue.EMPTY_LIST
|
||||
|
||||
|
||||
@comfytype(io_type="IMAGECOMPARE")
|
||||
class ImageCompare(ComfyTypeI):
|
||||
Type = dict
|
||||
@@ -1418,6 +1567,8 @@ def setup_dynamic_input_funcs():
|
||||
register_dynamic_input_func(DynamicCombo.io_type, DynamicCombo._expand_schema_for_dynamic)
|
||||
# DynamicSlot.Input
|
||||
register_dynamic_input_func(DynamicSlot.io_type, DynamicSlot._expand_schema_for_dynamic)
|
||||
# DynamicGroup.Input
|
||||
register_dynamic_input_func(DynamicGroup.io_type, DynamicGroup._expand_schema_for_dynamic)
|
||||
|
||||
if len(DYNAMIC_INPUT_LOOKUP) == 0:
|
||||
setup_dynamic_input_funcs()
|
||||
@@ -1429,6 +1580,8 @@ class V3Data(TypedDict):
|
||||
'Dictionary where the keys are the input ids and the values dictate how to turn the inputs into a nested dictionary.'
|
||||
dynamic_paths_default_value: dict[str, Any]
|
||||
'Dictionary where the keys are the input ids and the values are a string from DynamicPathsDefaultValue for the inputs if value is None.'
|
||||
list_paths: set[str]
|
||||
'Set of top-level keys whose index-keyed dict values should be converted to a sorted list[dict] after build_nested_inputs runs.'
|
||||
create_dynamic_tuple: bool
|
||||
'When True, the value of the dynamic input will be in the format (value, path_key).'
|
||||
|
||||
@@ -1770,6 +1923,7 @@ def get_finalized_class_inputs(d: dict[str, Any], live_inputs: dict[str, Any], i
|
||||
"optional": {},
|
||||
"dynamic_paths": {},
|
||||
"dynamic_paths_default_value": {},
|
||||
"list_paths": set(),
|
||||
}
|
||||
d = d.copy()
|
||||
# ignore hidden for parsing
|
||||
@@ -1785,6 +1939,10 @@ def get_finalized_class_inputs(d: dict[str, Any], live_inputs: dict[str, Any], i
|
||||
dynamic_paths_default_value = out_dict.pop("dynamic_paths_default_value", None)
|
||||
if dynamic_paths_default_value is not None and len(dynamic_paths_default_value) > 0:
|
||||
v3_data["dynamic_paths_default_value"] = dynamic_paths_default_value
|
||||
# list_paths: keys whose nested dict should be post-converted to a sorted list[dict]
|
||||
list_paths = out_dict.pop("list_paths", None)
|
||||
if list_paths:
|
||||
v3_data["list_paths"] = list_paths
|
||||
return out_dict, hidden, v3_data
|
||||
|
||||
def parse_class_inputs(out_dict: dict[str, Any], live_inputs: dict[str, Any], curr_dict: dict[str, Any], curr_prefix: list[str] | None=None) -> None:
|
||||
@@ -1820,10 +1978,12 @@ def add_to_dict_v1(i: Input, d: dict):
|
||||
|
||||
class DynamicPathsDefaultValue:
|
||||
EMPTY_DICT = "empty_dict"
|
||||
EMPTY_LIST = "empty_list"
|
||||
|
||||
def build_nested_inputs(values: dict[str, Any], v3_data: V3Data):
|
||||
paths = v3_data.get("dynamic_paths", None)
|
||||
default_value_dict = v3_data.get("dynamic_paths_default_value", {})
|
||||
list_paths: set[str] = v3_data.get("list_paths", set()) or set()
|
||||
if paths is None:
|
||||
return values
|
||||
values = values.copy()
|
||||
@@ -1846,6 +2006,8 @@ def build_nested_inputs(values: dict[str, Any], v3_data: V3Data):
|
||||
default_option = default_value_dict.get(key, None)
|
||||
if default_option == DynamicPathsDefaultValue.EMPTY_DICT:
|
||||
value = {}
|
||||
elif default_option == DynamicPathsDefaultValue.EMPTY_LIST:
|
||||
value = []
|
||||
if create_tuple:
|
||||
value = (value, key)
|
||||
current[p] = value
|
||||
@@ -1853,6 +2015,34 @@ def build_nested_inputs(values: dict[str, Any], v3_data: V3Data):
|
||||
current = current.setdefault(p, {})
|
||||
|
||||
values.update(result)
|
||||
|
||||
# Post-pass: convert index-keyed dicts to sorted lists for io.DynamicGroup fields
|
||||
for list_path in list_paths:
|
||||
parts = list_path.split(".")
|
||||
# Navigate to the parent container, then convert the leaf
|
||||
container = values
|
||||
for part in parts[:-1]:
|
||||
if not isinstance(container, dict) or part not in container:
|
||||
container = None
|
||||
break
|
||||
container = container[part]
|
||||
if container is None:
|
||||
continue
|
||||
leaf_key = parts[-1]
|
||||
leaf = container.get(leaf_key, None)
|
||||
if isinstance(leaf, dict):
|
||||
try:
|
||||
sorted_rows = [leaf[k] for k in sorted(leaf.keys(), key=int)]
|
||||
container[leaf_key] = sorted_rows
|
||||
except (ValueError, TypeError):
|
||||
# Keys are not all integers; leave as-is
|
||||
pass
|
||||
elif isinstance(leaf, list):
|
||||
# Already a list (e.g. the EMPTY_LIST default was applied above)
|
||||
pass
|
||||
elif leaf is None:
|
||||
container[leaf_key] = []
|
||||
|
||||
return values
|
||||
|
||||
|
||||
@@ -2417,7 +2607,9 @@ __all__ = [
|
||||
# Dynamic Types
|
||||
"MatchType",
|
||||
"DynamicCombo",
|
||||
"DynamicSlot",
|
||||
"Autogrow",
|
||||
"DynamicGroup",
|
||||
# Other classes
|
||||
"HiddenHolder",
|
||||
"Hidden",
|
||||
|
||||
@@ -17,10 +17,6 @@ class Seedream4Options(BaseModel):
|
||||
max_images: int = Field(15)
|
||||
|
||||
|
||||
class Seedream5OptimizePromptOptions(BaseModel):
|
||||
thinking: Literal["auto", "enabled", "disabled"] = Field(...)
|
||||
|
||||
|
||||
class Seedream4TaskCreationRequest(BaseModel):
|
||||
model: str = Field(...)
|
||||
prompt: str = Field(...)
|
||||
@@ -32,7 +28,6 @@ class Seedream4TaskCreationRequest(BaseModel):
|
||||
sequential_image_generation_options: Seedream4Options | None = Field(Seedream4Options(max_images=15))
|
||||
watermark: bool = Field(False)
|
||||
output_format: str | None = None
|
||||
optimize_prompt_options: Seedream5OptimizePromptOptions | None = None
|
||||
|
||||
|
||||
class ImageTaskCreationResponse(BaseModel):
|
||||
|
||||
@@ -77,7 +77,6 @@ class To3DUVTaskRequest(BaseModel):
|
||||
|
||||
class To3DPartTaskRequest(BaseModel):
|
||||
File: TaskFile3DInput = Field(...)
|
||||
EnableStagedGeneration: bool | None = Field(None)
|
||||
|
||||
|
||||
class TextureEditImageInfo(BaseModel):
|
||||
|
||||
@@ -128,7 +128,7 @@ class OpenAIResponse(ModelResponseProperties, ResponseProperties):
|
||||
parallel_tool_calls: bool | None = Field(True)
|
||||
status: str | None = Field(
|
||||
None,
|
||||
description="One of `completed`, `failed`, `in_progress`, `incomplete`, `queued`, or `cancelled`.",
|
||||
description="One of `completed`, `failed`, `in_progress`, or `incomplete`.",
|
||||
)
|
||||
usage: ResponseUsage | None = Field(None)
|
||||
|
||||
|
||||
@@ -1,49 +0,0 @@
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class SyncInputItem(BaseModel):
|
||||
type: str = Field(..., description="Input kind: 'video', 'image' or 'audio'.")
|
||||
url: str = Field(...)
|
||||
|
||||
|
||||
class SyncActiveSpeakerDetection(BaseModel):
|
||||
auto_detect: bool | None = Field(
|
||||
None, description="Detect the active speaker automatically. Video input only; rejected for images."
|
||||
)
|
||||
frame_number: int | None = Field(
|
||||
None, description="Frame used for manual speaker selection. Must be 0 for image inputs."
|
||||
)
|
||||
coordinates: list[int] | None = Field(
|
||||
None, description="Pixel [x, y] of the speaker's face in the frame selected by frame_number."
|
||||
)
|
||||
|
||||
|
||||
class SyncGenerationOptions(BaseModel):
|
||||
sync_mode: str | None = Field(
|
||||
None,
|
||||
description="How to resolve an audio/video duration mismatch: "
|
||||
"cut_off, bounce, loop, silence or remap. Ignored for image inputs.",
|
||||
)
|
||||
i2v_prompt: str | None = Field(
|
||||
None, description="Motion prompt for image-to-video generation. Image input only."
|
||||
)
|
||||
active_speaker_detection: SyncActiveSpeakerDetection | None = Field(None)
|
||||
|
||||
|
||||
class SyncGenerationRequest(BaseModel):
|
||||
model: str = Field(..., description="Generation model, e.g. 'sync-3'.")
|
||||
input: list[SyncInputItem] = Field(
|
||||
..., description="Exactly one visual input (video or image) plus one audio input."
|
||||
)
|
||||
options: SyncGenerationOptions | None = Field(None)
|
||||
|
||||
|
||||
class SyncGeneration(BaseModel):
|
||||
"""Subset of the Generation object returned by POST /v2/generate and GET /v2/generate/{id}."""
|
||||
|
||||
id: str = Field(...)
|
||||
status: str = Field(..., description="PENDING | PROCESSING | COMPLETED | FAILED | REJECTED")
|
||||
outputUrl: str | None = Field(None)
|
||||
outputDuration: float | None = Field(None)
|
||||
error: str | None = Field(None, description="Human-readable failure message.")
|
||||
errorCode: str | None = Field(None, description="Stable machine-readable code from the GET /v2/errors catalog.")
|
||||
@@ -34,7 +34,6 @@ from comfy_api_nodes.apis.bytedance import (
|
||||
SeedanceVirtualLibraryCreateAssetRequest,
|
||||
Seedream4Options,
|
||||
Seedream4TaskCreationRequest,
|
||||
Seedream5OptimizePromptOptions,
|
||||
TaskAudioContent,
|
||||
TaskAudioContentUrl,
|
||||
TaskCreationResponse,
|
||||
@@ -876,17 +875,6 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
tooltip='Whether to add an "AI generated" watermark to the image.',
|
||||
advanced=True,
|
||||
),
|
||||
IO.Boolean.Input(
|
||||
"thinking",
|
||||
default=True,
|
||||
tooltip=(
|
||||
"Enable the model's prompt-optimization reasoning ('thinking') for better adherence. "
|
||||
"Can substantially increase generation time — notably on Seedream 5.0 Pro. "
|
||||
"Can only be disabled for text-to-image (not when reference images are provided)."
|
||||
),
|
||||
optional=True,
|
||||
advanced=True,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
IO.Image.Output(),
|
||||
@@ -932,7 +920,6 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
model: dict,
|
||||
seed: int = 0,
|
||||
watermark: bool = False,
|
||||
thinking: bool = True,
|
||||
) -> IO.NodeOutput:
|
||||
validate_string(prompt, strip_whitespace=True, min_length=1)
|
||||
model_id = SEEDREAM_MODELS[model["model"]]
|
||||
@@ -992,10 +979,6 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
raise ValueError(
|
||||
"The maximum number of generated images plus the number of reference images cannot exceed 15."
|
||||
)
|
||||
if not thinking and n_input_images > 0:
|
||||
raise ValueError(
|
||||
"'thinking' can only be disabled for text-to-image; enable it when using reference images."
|
||||
)
|
||||
|
||||
reference_images_urls: list[str] = []
|
||||
if image_tensors:
|
||||
@@ -1009,9 +992,6 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
wait_label="Uploading reference images",
|
||||
)
|
||||
|
||||
optimize_prompt_options = None
|
||||
if n_input_images == 0:
|
||||
optimize_prompt_options = Seedream5OptimizePromptOptions(thinking="enabled" if thinking else "disabled")
|
||||
response = await sync_op(
|
||||
cls,
|
||||
ApiEndpoint(path=BYTEPLUS_IMAGE_ENDPOINT, method="POST"),
|
||||
@@ -1025,7 +1005,6 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
sequential_image_generation=None if is_pro else sequential_image_generation,
|
||||
sequential_image_generation_options=None if is_pro else Seedream4Options(max_images=max_images),
|
||||
watermark=watermark,
|
||||
optimize_prompt_options=optimize_prompt_options,
|
||||
),
|
||||
)
|
||||
if len(response.data) == 1:
|
||||
|
||||
@@ -1133,9 +1133,7 @@ class GeminiImage2(IO.ComfyNode):
|
||||
) -> IO.NodeOutput:
|
||||
validate_string(prompt, strip_whitespace=True, min_length=1)
|
||||
if model == "Nano Banana 2 (Gemini 3.1 Flash Image)":
|
||||
model = "gemini-3.1-flash-image"
|
||||
elif model == "gemini-3-pro-image-preview":
|
||||
model = "gemini-3-pro-image"
|
||||
model = "gemini-3.1-flash-image-preview"
|
||||
|
||||
parts: list[GeminiPart] = [GeminiPart(text=prompt)]
|
||||
if images is not None:
|
||||
@@ -1509,7 +1507,7 @@ class GeminiNanoBanana2V2(IO.ComfyNode):
|
||||
validate_string(prompt, strip_whitespace=True, min_length=1)
|
||||
model_choice = model["model"]
|
||||
if model_choice == "Nano Banana 2 (Gemini 3.1 Flash Image)":
|
||||
model_id = "gemini-3.1-flash-image"
|
||||
model_id = "gemini-3.1-flash-image-preview"
|
||||
elif model_choice == "Nano Banana 2 Lite":
|
||||
model_id = "gemini-3.1-flash-lite-image"
|
||||
else:
|
||||
|
||||
@@ -642,7 +642,6 @@ class Tencent3DPartNode(IO.ComfyNode):
|
||||
response_model=To3DProTaskCreateResponse,
|
||||
data=To3DPartTaskRequest(
|
||||
File=TaskFile3DInput(Type=file_format.upper(), Url=model_url),
|
||||
EnableStagedGeneration=True,
|
||||
),
|
||||
is_rate_limited=_is_tencent_rate_limited,
|
||||
)
|
||||
|
||||
@@ -41,9 +41,6 @@ STARTING_POINT_ID_PATTERN = r"<starting_point_id:(.*)>"
|
||||
|
||||
|
||||
class SupportedOpenAIModel(str, Enum):
|
||||
gpt_5_6_sol = "gpt-5.6-sol"
|
||||
gpt_5_6_terra = "gpt-5.6-terra"
|
||||
gpt_5_6_luna = "gpt-5.6-luna"
|
||||
gpt_5_5_pro = "gpt-5.5-pro"
|
||||
gpt_5_5 = "gpt-5.5"
|
||||
gpt_5 = "gpt-5"
|
||||
@@ -1066,21 +1063,6 @@ class OpenAIChatNode(IO.ComfyNode):
|
||||
"usd": [0.002, 0.008],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "gpt-5.6-terra") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.0025, 0.015],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "gpt-5.6-luna") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.001, 0.006],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "gpt-5.6") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.005, 0.03],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "gpt-5.5-pro") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.03, 0.18],
|
||||
|
||||
@@ -1,391 +0,0 @@
|
||||
from typing_extensions import override
|
||||
|
||||
from comfy_api.latest import IO, ComfyExtension, Input
|
||||
from comfy_api_nodes.apis.sync_so import (
|
||||
SyncActiveSpeakerDetection,
|
||||
SyncGeneration,
|
||||
SyncGenerationOptions,
|
||||
SyncGenerationRequest,
|
||||
SyncInputItem,
|
||||
)
|
||||
from comfy_api_nodes.util import (
|
||||
ApiEndpoint,
|
||||
download_url_to_video_output,
|
||||
downscale_image_tensor,
|
||||
downscale_image_tensor_by_max_side,
|
||||
get_image_dimensions,
|
||||
get_number_of_images,
|
||||
poll_op,
|
||||
sync_op,
|
||||
upload_audio_to_comfyapi,
|
||||
upload_image_to_comfyapi,
|
||||
upload_video_to_comfyapi,
|
||||
validate_audio_duration,
|
||||
)
|
||||
|
||||
|
||||
class SyncLipSyncNode(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="SyncLipSyncNode",
|
||||
display_name="sync.so Lip Sync",
|
||||
category="partner/video/sync.so",
|
||||
description=(
|
||||
"Re-sync mouth movement in a video to new speech audio using sync.so. "
|
||||
"Handles close-ups, profiles and obstructions automatically while preserving "
|
||||
"the speaker's expression. Cost scales with output duration."
|
||||
),
|
||||
inputs=[
|
||||
IO.Video.Input(
|
||||
"video",
|
||||
tooltip="Footage of the speaker to re-sync. Up to 4K (4096x2160); "
|
||||
"a constant frame rate of 24/25/30 fps works best.",
|
||||
),
|
||||
IO.Audio.Input(
|
||||
"audio",
|
||||
tooltip="Speech audio to sync the mouth to.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
tooltip="Seed controls whether the node should re-run; "
|
||||
"results are non-deterministic regardless of seed.",
|
||||
),
|
||||
IO.DynamicCombo.Input(
|
||||
"model",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"sync-3",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"sync_mode",
|
||||
options=["bounce", "cut_off", "loop", "silence", "remap"],
|
||||
default="bounce",
|
||||
tooltip=(
|
||||
"How to handle a duration mismatch between video and audio; "
|
||||
"this also sets the output length. "
|
||||
"bounce: video plays forward then backward until the audio ends "
|
||||
"(output = audio length). "
|
||||
"loop: video restarts until the audio ends (output = audio length). "
|
||||
"remap: video is time-stretched to match the audio (output = audio length). "
|
||||
"cut_off: the longer track is trimmed (output = shorter length). "
|
||||
"silence: nothing is trimmed; the shorter track is padded "
|
||||
"(output = longer length)."
|
||||
),
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"speaker_selection",
|
||||
options=["default", "auto-detect", "coordinates"],
|
||||
default="default",
|
||||
tooltip=(
|
||||
"Which face to lipsync when several people are visible. "
|
||||
"default: let the model decide. "
|
||||
"auto-detect: detect and follow the active speaker. "
|
||||
"coordinates: target the face at pixel (speaker_x, speaker_y) "
|
||||
"in the frame chosen by speaker_frame."
|
||||
),
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_frame",
|
||||
default=0,
|
||||
min=0,
|
||||
max=1_000_000,
|
||||
advanced=True,
|
||||
tooltip="Video frame used to locate the speaker. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_x",
|
||||
default=0,
|
||||
min=0,
|
||||
max=4096,
|
||||
advanced=True,
|
||||
tooltip="X pixel coordinate of the speaker's face. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_y",
|
||||
default=0,
|
||||
min=0,
|
||||
max=4096,
|
||||
advanced=True,
|
||||
tooltip="Y pixel coordinate of the speaker's face. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
],
|
||||
)
|
||||
],
|
||||
tooltip="sync.so generation model.",
|
||||
),
|
||||
],
|
||||
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.19019,"format":{"approximate":true,"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
video: Input.Video,
|
||||
audio: Input.Audio,
|
||||
seed: int,
|
||||
model: dict,
|
||||
) -> IO.NodeOutput:
|
||||
try:
|
||||
width, height = video.get_dimensions()
|
||||
except Exception:
|
||||
width = height = None
|
||||
if width and height and (max(width, height) > 4096 or width * height > 4096 * 2160):
|
||||
raise ValueError(
|
||||
f"sync.so rejects videos above 4K (4096x2160); got {width}x{height}. Downscale the video first."
|
||||
)
|
||||
validate_audio_duration(audio, max_duration=600)
|
||||
|
||||
if model["speaker_selection"] == "auto-detect":
|
||||
speaker_detection = SyncActiveSpeakerDetection(auto_detect=True)
|
||||
elif model["speaker_selection"] == "coordinates":
|
||||
speaker_detection = SyncActiveSpeakerDetection(
|
||||
frame_number=model["speaker_frame"],
|
||||
coordinates=[model["speaker_x"], model["speaker_y"]],
|
||||
)
|
||||
else:
|
||||
speaker_detection = None
|
||||
|
||||
video_url = await upload_video_to_comfyapi(cls, video, max_duration=600)
|
||||
audio_url = await upload_audio_to_comfyapi(cls, audio)
|
||||
|
||||
generation = await sync_op(
|
||||
cls,
|
||||
ApiEndpoint(path="/proxy/synclabs/v2/generate", method="POST"),
|
||||
response_model=SyncGeneration,
|
||||
data=SyncGenerationRequest(
|
||||
model=model["model"],
|
||||
input=[
|
||||
SyncInputItem(type="video", url=video_url),
|
||||
SyncInputItem(type="audio", url=audio_url),
|
||||
],
|
||||
options=SyncGenerationOptions(
|
||||
sync_mode=model["sync_mode"],
|
||||
active_speaker_detection=speaker_detection,
|
||||
),
|
||||
),
|
||||
)
|
||||
generation = await poll_op(
|
||||
cls,
|
||||
ApiEndpoint(path=f"/proxy/synclabs/v2/generate/{generation.id}"),
|
||||
response_model=SyncGeneration,
|
||||
status_extractor=lambda g: g.status,
|
||||
completed_statuses=["COMPLETED", "FAILED", "REJECTED"],
|
||||
failed_statuses=[],
|
||||
queued_statuses=["PENDING"],
|
||||
poll_interval=10.0,
|
||||
)
|
||||
if generation.status != "COMPLETED":
|
||||
code = f" [{generation.errorCode}]" if generation.errorCode else ""
|
||||
raise ValueError(
|
||||
f"sync.so generation {generation.status.lower()}{code}: "
|
||||
f"{generation.error or 'no error details provided'}"
|
||||
)
|
||||
if not generation.outputUrl:
|
||||
raise ValueError("sync.so generation completed but no output URL was returned.")
|
||||
return IO.NodeOutput(await download_url_to_video_output(generation.outputUrl))
|
||||
|
||||
|
||||
class SyncTalkingImageNode(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="SyncTalkingImageNode",
|
||||
display_name="sync.so Talking Image",
|
||||
category="partner/video/sync.so",
|
||||
description=(
|
||||
"Animate a still portrait into a talking video driven by speech audio, "
|
||||
"using sync.so's sync-3 model. The output duration matches the audio. "
|
||||
"Cost scales with output duration."
|
||||
),
|
||||
inputs=[
|
||||
IO.Image.Input(
|
||||
"image",
|
||||
tooltip="A single image with a clearly visible face, up to 4K (4096x2160).",
|
||||
),
|
||||
IO.Audio.Input(
|
||||
"audio",
|
||||
tooltip="Speech audio driving the talking video; the output duration matches it. "
|
||||
"Chain any TTS node here to drive the animation from text.",
|
||||
),
|
||||
IO.String.Input(
|
||||
"prompt",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Optional guidance for how the portrait comes to life, e.g. "
|
||||
"'make the subject smile and look at the camera'. "
|
||||
"Leave empty for natural talking motion.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
tooltip="Seed controls whether the node should re-run; "
|
||||
"results are non-deterministic regardless of seed.",
|
||||
),
|
||||
IO.DynamicCombo.Input(
|
||||
"model",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"sync-3",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"speaker_selection",
|
||||
options=["default", "coordinates"],
|
||||
default="default",
|
||||
tooltip=(
|
||||
"Which face to animate when several people are visible. "
|
||||
"default: let the model decide. "
|
||||
"coordinates: target the face at pixel (speaker_x, speaker_y) "
|
||||
"in the image. Auto-detection is not supported for images."
|
||||
),
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_x",
|
||||
default=0,
|
||||
min=0,
|
||||
max=4096,
|
||||
advanced=True,
|
||||
tooltip="X pixel coordinate of the speaker's face. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_y",
|
||||
default=0,
|
||||
min=0,
|
||||
max=4096,
|
||||
advanced=True,
|
||||
tooltip="Y pixel coordinate of the speaker's face. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
IO.Boolean.Input(
|
||||
"auto_downscale",
|
||||
default=True,
|
||||
advanced=True,
|
||||
tooltip="Automatically downscale the image if it exceeds the 4K "
|
||||
"(4096x2160) input limit; speaker coordinates are scaled to match. "
|
||||
"When disabled, an oversized image raises an error instead.",
|
||||
),
|
||||
],
|
||||
)
|
||||
],
|
||||
tooltip="sync.so generation model. Image input is exclusive to sync-3.",
|
||||
),
|
||||
],
|
||||
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.19019,"format":{"approximate":true,"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
image: Input.Image,
|
||||
audio: Input.Audio,
|
||||
prompt: str,
|
||||
seed: int,
|
||||
model: dict,
|
||||
) -> IO.NodeOutput:
|
||||
if get_number_of_images(image) != 1:
|
||||
raise ValueError("Exactly one image is required; got a batch. Pick one frame first.")
|
||||
validate_audio_duration(audio, max_duration=600)
|
||||
|
||||
height, width = get_image_dimensions(image)
|
||||
speaker_x, speaker_y = model["speaker_x"], model["speaker_y"]
|
||||
if max(width, height) > 4096 or width * height > 4096 * 2160:
|
||||
if not model["auto_downscale"]:
|
||||
raise ValueError(
|
||||
f"sync.so rejects images above 4K (4096x2160); got {width}x{height}. "
|
||||
"Downscale the image first or enable auto_downscale."
|
||||
)
|
||||
image = downscale_image_tensor(image, total_pixels=4096 * 2160)
|
||||
image = downscale_image_tensor_by_max_side(image, max_side=4096)
|
||||
new_height, new_width = get_image_dimensions(image)
|
||||
# speaker coordinates are given in the original image's pixel space
|
||||
speaker_x = min(new_width - 1, round(speaker_x * new_width / width))
|
||||
speaker_y = min(new_height - 1, round(speaker_y * new_height / height))
|
||||
|
||||
if model["speaker_selection"] == "coordinates":
|
||||
speaker_detection = SyncActiveSpeakerDetection(
|
||||
frame_number=0, # images have a single frame; auto_detect is rejected by the API
|
||||
coordinates=[speaker_x, speaker_y],
|
||||
)
|
||||
else:
|
||||
speaker_detection = None
|
||||
|
||||
image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None)
|
||||
audio_url = await upload_audio_to_comfyapi(cls, audio)
|
||||
|
||||
generation = await sync_op(
|
||||
cls,
|
||||
ApiEndpoint(path="/proxy/synclabs/v2/generate", method="POST"),
|
||||
response_model=SyncGeneration,
|
||||
data=SyncGenerationRequest(
|
||||
model=model["model"],
|
||||
input=[
|
||||
SyncInputItem(type="image", url=image_url),
|
||||
SyncInputItem(type="audio", url=audio_url),
|
||||
],
|
||||
options=SyncGenerationOptions(
|
||||
i2v_prompt=prompt.strip() or None,
|
||||
active_speaker_detection=speaker_detection,
|
||||
),
|
||||
),
|
||||
)
|
||||
generation = await poll_op(
|
||||
cls,
|
||||
ApiEndpoint(path=f"/proxy/synclabs/v2/generate/{generation.id}"),
|
||||
response_model=SyncGeneration,
|
||||
status_extractor=lambda g: g.status,
|
||||
completed_statuses=["COMPLETED", "FAILED", "REJECTED"],
|
||||
failed_statuses=[],
|
||||
queued_statuses=["PENDING"],
|
||||
poll_interval=10.0,
|
||||
)
|
||||
if generation.status != "COMPLETED":
|
||||
code = f" [{generation.errorCode}]" if generation.errorCode else ""
|
||||
raise ValueError(
|
||||
f"sync.so generation {generation.status.lower()}{code}: "
|
||||
f"{generation.error or 'no error details provided'}"
|
||||
)
|
||||
if not generation.outputUrl:
|
||||
raise ValueError("sync.so generation completed but no output URL was returned.")
|
||||
return IO.NodeOutput(await download_url_to_video_output(generation.outputUrl))
|
||||
|
||||
|
||||
class SyncExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[IO.ComfyNode]]:
|
||||
return [
|
||||
SyncLipSyncNode,
|
||||
SyncTalkingImageNode,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> SyncExtension:
|
||||
return SyncExtension()
|
||||
@@ -15,8 +15,6 @@ from comfy.comfy_api_env import normalize_comfy_api_base
|
||||
from comfy.deploy_environment import get_deploy_environment
|
||||
from comfy.model_management import processing_interrupted
|
||||
from comfy_api.latest import IO
|
||||
from comfy_execution.utils import get_executing_context
|
||||
from comfyui_version import __version__ as comfyui_version
|
||||
|
||||
from .common_exceptions import ProcessingInterrupted
|
||||
|
||||
@@ -58,16 +56,11 @@ def get_comfy_api_headers(node_cls: type[IO.ComfyNode]) -> dict[str, str]:
|
||||
relative/cloud URLs resolved against ``default_base_url()``; because the result
|
||||
includes auth, callers must not attach it to arbitrary absolute/presigned URLs.
|
||||
"""
|
||||
headers = {
|
||||
return {
|
||||
**get_auth_header(node_cls),
|
||||
"Comfy-Env": get_deploy_environment(),
|
||||
"Comfy-Usage-Source": get_usage_source(node_cls),
|
||||
"Comfy-Core-Version": comfyui_version,
|
||||
}
|
||||
ctx = get_executing_context()
|
||||
if ctx is not None:
|
||||
headers["Comfy-Job-Id"] = ctx.prompt_id
|
||||
return headers
|
||||
|
||||
|
||||
def default_base_url() -> str:
|
||||
|
||||
@@ -503,8 +503,6 @@ RAM_CACHE_DEFAULT_RAM_USAGE = 0.05
|
||||
|
||||
RAM_CACHE_OLD_WORKFLOW_OOM_MULTIPLIER = 1.3
|
||||
|
||||
RAM_CACHE_LARGE_INTERMEDIATE = 512 * 1024 ** 2
|
||||
|
||||
|
||||
def all_outputs_dynamic(outputs):
|
||||
if outputs is None:
|
||||
@@ -519,6 +517,7 @@ def all_outputs_dynamic(outputs):
|
||||
|
||||
return True
|
||||
|
||||
|
||||
class RAMPressureCache(LRUCache):
|
||||
|
||||
def __init__(self, key_class, enable_providers=False):
|
||||
@@ -540,9 +539,9 @@ class RAMPressureCache(LRUCache):
|
||||
self.timestamps[self.cache_key_set.get_data_key(node_id)] = time.time()
|
||||
super().set_local(node_id, value)
|
||||
|
||||
def ram_release(self, target, free_active=False, min_entry_size=0):
|
||||
def ram_release(self, target, free_active=False):
|
||||
if psutil.virtual_memory().available >= target:
|
||||
return 0
|
||||
return
|
||||
|
||||
clean_list = []
|
||||
|
||||
@@ -556,9 +555,8 @@ class RAMPressureCache(LRUCache):
|
||||
oom_score = RAM_CACHE_OLD_WORKFLOW_OOM_MULTIPLIER ** (self.generation - self.used_generation[key])
|
||||
|
||||
ram_usage = RAM_CACHE_DEFAULT_RAM_USAGE
|
||||
oom_ram_usage = ram_usage
|
||||
def scan_list_for_ram_usage(outputs):
|
||||
nonlocal ram_usage, oom_ram_usage
|
||||
nonlocal ram_usage
|
||||
if outputs is None:
|
||||
return
|
||||
for output in outputs:
|
||||
@@ -566,26 +564,19 @@ class RAMPressureCache(LRUCache):
|
||||
scan_list_for_ram_usage(output)
|
||||
elif isinstance(output, torch.Tensor) and output.device.type == 'cpu':
|
||||
ram_usage += output.numel() * output.element_size()
|
||||
oom_ram_usage += output.numel() * output.element_size()
|
||||
elif isinstance(output, ModelPatcher) and self.used_generation[key] != self.generation:
|
||||
#old ModelPatchers are the first to go
|
||||
oom_ram_usage = 1e30
|
||||
ram_usage = 1e30
|
||||
scan_list_for_ram_usage(cache_entry.outputs)
|
||||
|
||||
if ram_usage < min_entry_size:
|
||||
continue
|
||||
|
||||
oom_score *= oom_ram_usage
|
||||
oom_score *= ram_usage
|
||||
#In the case where we have no information on the node ram usage at all,
|
||||
#break OOM score ties on the last touch timestamp (pure LRU)
|
||||
bisect.insort(clean_list, (oom_score, self.timestamps[key], key, ram_usage))
|
||||
bisect.insort(clean_list, (oom_score, self.timestamps[key], key))
|
||||
|
||||
freed = 0
|
||||
while psutil.virtual_memory().available < target and clean_list:
|
||||
_, _, key, ram_usage = clean_list.pop()
|
||||
_, _, key = clean_list.pop()
|
||||
del self.cache[key]
|
||||
self.used_generation.pop(key, None)
|
||||
self.timestamps.pop(key, None)
|
||||
self.children.pop(key, None)
|
||||
freed += ram_usage
|
||||
return freed
|
||||
|
||||
+3
-12
@@ -56,9 +56,6 @@ PREVIEWABLE_MEDIA_TYPES = frozenset({'images', 'video', 'audio', '3d', 'text'})
|
||||
# 3D file extensions for preview fallback (no dedicated media_type exists)
|
||||
THREE_D_EXTENSIONS = frozenset({'.obj', '.fbx', '.gltf', '.glb', '.usdz'})
|
||||
|
||||
# Text file extensions for preview fallback (the formats SaveText can produce)
|
||||
TEXT_EXTENSIONS = frozenset({'.txt', '.md', '.json'})
|
||||
|
||||
|
||||
def has_3d_extension(filename: str) -> bool:
|
||||
lower = filename.lower()
|
||||
@@ -146,10 +143,9 @@ def is_previewable(media_type: str, item: dict) -> bool:
|
||||
Maintains backwards compatibility with existing logic.
|
||||
|
||||
Priority:
|
||||
1. media_type is 'images', 'video', 'audio', '3d', or 'text'
|
||||
1. media_type is 'images', 'video', 'audio', or '3d'
|
||||
2. format field starts with 'video/' or 'audio/'
|
||||
3. filename has a 3D extension (.obj, .fbx, .gltf, .glb, .usdz)
|
||||
4. filename has a text extension (.txt, .md, .json, ...)
|
||||
"""
|
||||
if media_type in PREVIEWABLE_MEDIA_TYPES:
|
||||
return True
|
||||
@@ -160,12 +156,10 @@ def is_previewable(media_type: str, item: dict) -> bool:
|
||||
if fmt and (fmt.startswith('video/') or fmt.startswith('audio/')):
|
||||
return True
|
||||
|
||||
# Check for 3D and text files by extension
|
||||
# Check for 3D files by extension
|
||||
filename = item.get('filename', '').lower()
|
||||
if any(filename.endswith(ext) for ext in THREE_D_EXTENSIONS):
|
||||
return True
|
||||
if any(filename.endswith(ext) for ext in TEXT_EXTENSIONS):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
@@ -261,10 +255,6 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]:
|
||||
Preview priority (matching frontend):
|
||||
1. type="output" with previewable media
|
||||
2. Any previewable media
|
||||
|
||||
Text content entries (strings under 'text') are preview-only metadata,
|
||||
matching the frontend's METADATA_KEYS: they can serve as the fallback
|
||||
preview but are not counted as outputs.
|
||||
"""
|
||||
count = 0
|
||||
preview_output = None
|
||||
@@ -285,6 +275,7 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]:
|
||||
if normalized is None:
|
||||
# Not a 3D file string — check for text preview
|
||||
if media_type == 'text':
|
||||
count += 1
|
||||
if preview_output is None:
|
||||
if isinstance(item, tuple):
|
||||
text_value = item[0] if item else ''
|
||||
|
||||
@@ -298,7 +298,6 @@ class PreviewAudio(IO.ComfyNode):
|
||||
search_aliases=["play audio"],
|
||||
display_name="Preview Audio",
|
||||
category="audio",
|
||||
description="Preview the audio without saving it to the ComfyUI output directory.",
|
||||
inputs=[
|
||||
IO.Audio.Input("audio"),
|
||||
],
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw, ImageEnhance, ImageFont
|
||||
@@ -168,111 +166,6 @@ def boxes_to_regions(boxes, width: int, height: int) -> list:
|
||||
return regions
|
||||
|
||||
|
||||
def normalize_incoming_boxes(bboxes) -> list:
|
||||
if isinstance(bboxes, dict):
|
||||
frame = [bboxes]
|
||||
elif not isinstance(bboxes, list) or not bboxes:
|
||||
frame = []
|
||||
elif isinstance(bboxes[0], dict):
|
||||
frame = bboxes
|
||||
else:
|
||||
frame = bboxes[0] if isinstance(bboxes[0], list) else []
|
||||
boxes = []
|
||||
for box in frame:
|
||||
if not isinstance(box, dict):
|
||||
continue
|
||||
norm = {
|
||||
"x": box.get("x", 0),
|
||||
"y": box.get("y", 0),
|
||||
"width": box.get("width", 0),
|
||||
"height": box.get("height", 0),
|
||||
}
|
||||
meta = box.get("metadata")
|
||||
if isinstance(meta, dict):
|
||||
norm["metadata"] = meta
|
||||
boxes.append(norm)
|
||||
return boxes
|
||||
|
||||
|
||||
def _looks_like_element(box: dict) -> bool:
|
||||
bbox = box.get("bbox")
|
||||
return isinstance(bbox, (list, tuple)) and len(bbox) == 4
|
||||
|
||||
|
||||
def _looks_like_bbox(box: dict) -> bool:
|
||||
return all(key in box for key in ("x", "y", "width", "height"))
|
||||
|
||||
|
||||
def elements_to_boxes(elements: list, width: int, height: int) -> list:
|
||||
boxes = []
|
||||
for element in elements:
|
||||
if not isinstance(element, dict):
|
||||
continue
|
||||
bbox = element.get("bbox")
|
||||
if not (isinstance(bbox, (list, tuple)) and len(bbox) == 4):
|
||||
raise ValueError("bboxes element is missing a valid 'bbox' [ymin, xmin, ymax, xmax]")
|
||||
try:
|
||||
ymin, xmin, ymax, xmax = (float(v) / 1000.0 for v in bbox)
|
||||
except (TypeError, ValueError):
|
||||
raise ValueError("bboxes element 'bbox' must contain four numbers")
|
||||
etype = "text" if element.get("type") == "text" else "obj"
|
||||
boxes.append({
|
||||
"x": round(min(xmin, xmax) * width),
|
||||
"y": round(min(ymin, ymax) * height),
|
||||
"width": round(abs(xmax - xmin) * width),
|
||||
"height": round(abs(ymax - ymin) * height),
|
||||
"metadata": {
|
||||
"type": etype,
|
||||
"text": element.get("text", "") if etype == "text" else "",
|
||||
"desc": element.get("desc", ""),
|
||||
"palette": element.get("color_palette", []) or [],
|
||||
},
|
||||
})
|
||||
return boxes
|
||||
|
||||
|
||||
def boxes_from_input(data, width: int, height: int) -> list:
|
||||
if data is None:
|
||||
return []
|
||||
if isinstance(data, str):
|
||||
text = data.strip()
|
||||
if not text:
|
||||
return []
|
||||
try:
|
||||
data = json.loads(text)
|
||||
except (ValueError, TypeError) as exc:
|
||||
raise ValueError(f"bboxes string input is not valid JSON: {exc}") from exc
|
||||
if isinstance(data, dict):
|
||||
if _looks_like_element(data):
|
||||
return elements_to_boxes([data], width, height)
|
||||
if _looks_like_bbox(data):
|
||||
return normalize_incoming_boxes(data)
|
||||
raise ValueError(
|
||||
"bboxes dict must be a bounding box (x, y, width, height) or an element (with a 'bbox')"
|
||||
)
|
||||
if not isinstance(data, list):
|
||||
raise ValueError(
|
||||
"bboxes input must be bounding boxes, elements, or a JSON string, "
|
||||
f"got {type(data).__name__}"
|
||||
)
|
||||
if not data:
|
||||
return []
|
||||
first = data[0]
|
||||
if isinstance(first, list):
|
||||
return normalize_incoming_boxes(data)
|
||||
if isinstance(first, dict):
|
||||
if _looks_like_element(first):
|
||||
return elements_to_boxes(data, width, height)
|
||||
if _looks_like_bbox(first):
|
||||
return normalize_incoming_boxes(data)
|
||||
raise ValueError(
|
||||
"bboxes items must be bounding boxes (x, y, width, height) or elements (with a 'bbox')"
|
||||
)
|
||||
raise ValueError(
|
||||
f"bboxes list must contain bounding boxes or elements, got {type(first).__name__}"
|
||||
)
|
||||
|
||||
|
||||
def _norm_bbox(region: dict) -> list[int]:
|
||||
def grid(value: float) -> int:
|
||||
return max(0, min(1000, round(value * 1000)))
|
||||
@@ -324,48 +217,29 @@ class CreateBoundingBoxes(io.ComfyNode):
|
||||
optional=True,
|
||||
tooltip="Optional image used as background in the canvas and preview.",
|
||||
),
|
||||
io.MultiType.Input(
|
||||
"bboxes",
|
||||
[io.BoundingBox, io.Array, io.String],
|
||||
optional=True,
|
||||
tooltip="Bounding boxes, elements, or a JSON string to initialize the canvas. A new upstream value initializes the canvas; edits made on the canvas take priority and are kept until the upstream value changes again.",
|
||||
),
|
||||
io.Int.Input("width", default=1024, min=64, max=16384, step=16,
|
||||
tooltip="Width of the canvas and the pixel grid for the bounding boxes."),
|
||||
io.Int.Input("height", default=1024, min=64, max=16384, step=16,
|
||||
tooltip="Height of the canvas and the pixel grid for the bounding boxes."),
|
||||
editor_state,
|
||||
io.BoundingBoxes.Input(
|
||||
"last_incoming",
|
||||
optional=True,
|
||||
tooltip="Internal state managed by the canvas: the upstream bboxes value that last initialized it. Leave empty to re-initialize the canvas from the bboxes input on the next run.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Image.Output(display_name="preview"),
|
||||
io.BoundingBox.Output(display_name="bboxes"),
|
||||
io.Array.Output(display_name="elements"),
|
||||
],
|
||||
is_output_node=True,
|
||||
is_experimental=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, width, height, editor_state=None, last_incoming=None, background=None, bboxes=None) -> io.NodeOutput:
|
||||
incoming = boxes_from_input(bboxes, width, height)
|
||||
applied = last_incoming if isinstance(last_incoming, list) else []
|
||||
upstream_changed = bool(incoming) and incoming != applied
|
||||
source = incoming if upstream_changed else (editor_state or [])
|
||||
regions = boxes_to_regions(source, width, height)
|
||||
def execute(cls, width, height, editor_state=None, background=None) -> io.NodeOutput:
|
||||
regions = boxes_to_regions(editor_state, width, height)
|
||||
preview = render_preview(regions, width, height, _bg_from_image(background))
|
||||
ui = {"dims": [width, height]}
|
||||
if incoming:
|
||||
ui["input_bboxes"] = incoming
|
||||
return io.NodeOutput(
|
||||
preview,
|
||||
fractions_to_bbox_frame(regions, width, height),
|
||||
build_elements(regions),
|
||||
ui=ui,
|
||||
ui={"dims": [width, height]},
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -844,18 +844,15 @@ class ImageMergeTileList(IO.ComfyNode):
|
||||
# Format specifications
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Maps (file_format, bit_depth, num_channels) -> (quantization scale, numpy dtype,
|
||||
# av frame pix_fmt, stream pix_fmt). Keeps the encode path declarative instead of branchy.
|
||||
# Maps (file_format, bit_depth, has_alpha) -> (numpy dtype scale, av pixel format,
|
||||
# stream pix_fmt). Keeps the encode path declarative instead of branchy.
|
||||
_FORMAT_SPECS = {
|
||||
("png", "8-bit", 1): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "gray", "stream_fmt": "gray"},
|
||||
("png", "8-bit", 3): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgb24", "stream_fmt": "rgb24"},
|
||||
("png", "8-bit", 4): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgba", "stream_fmt": "rgba"},
|
||||
("png", "16-bit", 1): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "gray16le", "stream_fmt": "gray16be"},
|
||||
("png", "16-bit", 3): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgb48le", "stream_fmt": "rgb48be"},
|
||||
("png", "16-bit", 4): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgba64le", "stream_fmt": "rgba64be"},
|
||||
("exr", "32-bit float", 1): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "grayf32le", "stream_fmt": "grayf32le"},
|
||||
("exr", "32-bit float", 3): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrpf32le", "stream_fmt": "gbrpf32le"},
|
||||
("exr", "32-bit float", 4): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrapf32le", "stream_fmt": "gbrapf32le"},
|
||||
("png", "8-bit", False): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgb24", "stream_fmt": "rgb24"},
|
||||
("png", "8-bit", True): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgba", "stream_fmt": "rgba"},
|
||||
("png", "16-bit", False): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgb48le", "stream_fmt": "rgb48be"},
|
||||
("png", "16-bit", True): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgba64le", "stream_fmt": "rgba64be"},
|
||||
("exr", "32-bit float", False): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrpf32le", "stream_fmt": "gbrpf32le"},
|
||||
("exr", "32-bit float", True): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrapf32le", "stream_fmt": "gbrapf32le"},
|
||||
}
|
||||
|
||||
|
||||
@@ -894,11 +891,10 @@ def hlg_to_linear(t: torch.Tensor) -> torch.Tensor:
|
||||
return torch.cat([hlg_to_linear(rgb), alpha], dim=-1)
|
||||
|
||||
# Piecewise: sqrt branch below 0.5, log branch above.
|
||||
# Clamp the log branch at the 0.5 branch point (not above it) so the
|
||||
# unselected lane stays finite in exp() without altering selected values;
|
||||
# Clamp inside the log branch so negative / out-of-range values don't blow up;
|
||||
# values above 1.0 are allowed and extrapolate naturally.
|
||||
low = (t ** 2) / 3.0
|
||||
high = (torch.exp((t.clamp(min=0.5) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0
|
||||
high = (torch.exp((t.clamp(min=_HLG_C) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0
|
||||
return torch.where(t <= 0.5, low, high)
|
||||
|
||||
|
||||
@@ -1091,8 +1087,7 @@ def _encode_image(
|
||||
bit_depth: str,
|
||||
colorspace: str,
|
||||
) -> bytes:
|
||||
"""Encode a single HxWxC (or channel-less HxW grayscale) tensor to PNG or
|
||||
EXR bytes in memory. Grayscale is written as single-channel PNG / Y-only EXR.
|
||||
"""Encode a single HxWxC tensor to PNG or EXR bytes in memory.
|
||||
|
||||
For EXR the input is interpreted according to `colorspace` and converted
|
||||
to scene-linear (EXR's convention) before writing:
|
||||
@@ -1106,16 +1101,10 @@ def _encode_image(
|
||||
For PNG, colorspace selection does not modify pixels — PNG is delivered
|
||||
sRGB-encoded and there is no PNG path for wide-gamut HDR in this node.
|
||||
"""
|
||||
if img_tensor.ndim == 2:
|
||||
img_tensor = img_tensor.unsqueeze(-1) # Some nodes emit grayscale as (H, W) with no channel dim, mask-style.
|
||||
height, width, num_channels = img_tensor.shape
|
||||
has_alpha = num_channels == 4
|
||||
|
||||
spec = _FORMAT_SPECS.get((file_format, bit_depth, num_channels))
|
||||
if spec is None:
|
||||
raise ValueError(
|
||||
f"No {file_format}/{bit_depth} encoder for {num_channels}-channel images: "
|
||||
"supported channel counts are 1 (grayscale), 3 (RGB) and 4 (RGBA)."
|
||||
)
|
||||
spec = _FORMAT_SPECS[(file_format, bit_depth, has_alpha)]
|
||||
|
||||
if spec["dtype"] == np.float32:
|
||||
# EXR path: preserve full range, no clamp.
|
||||
|
||||
@@ -1,102 +0,0 @@
|
||||
from typing_extensions import override
|
||||
|
||||
import comfy.utils
|
||||
import node_helpers
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
|
||||
|
||||
# fmt: off
|
||||
BUCKETS_1024 = [
|
||||
(512, 1792), (512, 1856), (512, 1920), (512, 1984), (512, 2048),
|
||||
(576, 1600), (576, 1664), (576, 1728), (576, 1792),
|
||||
(640, 1472), (640, 1536), (640, 1600),
|
||||
(704, 1344), (704, 1408), (704, 1472),
|
||||
(768, 1216), (768, 1280), (768, 1344),
|
||||
(832, 1152), (832, 1216),
|
||||
(896, 1088), (896, 1152),
|
||||
(960, 1024), (960, 1088),
|
||||
(1024, 960), (1024, 1024),
|
||||
(1088, 896), (1088, 960),
|
||||
(1152, 832), (1152, 896),
|
||||
(1216, 768), (1216, 832),
|
||||
(1280, 768),
|
||||
(1344, 704), (1344, 768),
|
||||
(1408, 704),
|
||||
(1472, 640), (1472, 704),
|
||||
(1536, 640),
|
||||
(1600, 576), (1600, 640),
|
||||
(1664, 576),
|
||||
(1728, 576),
|
||||
(1792, 512), (1792, 576),
|
||||
(1856, 512),
|
||||
(1920, 512),
|
||||
(1984, 512),
|
||||
(2048, 512),
|
||||
]
|
||||
# fmt: on
|
||||
|
||||
|
||||
def _find_best_bucket(height: int, width: int) -> tuple[int, int]:
|
||||
target_ratio = height / width
|
||||
return min(BUCKETS_1024, key=lambda hw: abs(hw[0] / hw[1] - target_ratio))
|
||||
|
||||
|
||||
def _resize_reference(image):
|
||||
if image.shape[0] != 1:
|
||||
raise ValueError("JoyImage reference inputs must contain one image each")
|
||||
samples = image.movedim(-1, 1)
|
||||
bucket_h, bucket_w = _find_best_bucket(samples.shape[2], samples.shape[3])
|
||||
resized = comfy.utils.common_upscale(samples, bucket_w, bucket_h, "bilinear", "center")
|
||||
return resized.movedim(1, -1)[:, :, :, :3]
|
||||
|
||||
|
||||
def _encode(clip, prompt, vae, images):
|
||||
resized_images = [_resize_reference(image) for image in images]
|
||||
conditioning = clip.encode_from_tokens_scheduled(clip.tokenize(prompt, images=resized_images))
|
||||
if vae is not None and resized_images:
|
||||
ref_latents = [vae.encode(image) for image in resized_images]
|
||||
conditioning = node_helpers.conditioning_set_values(
|
||||
conditioning, {"reference_latents": ref_latents}, append=True,
|
||||
)
|
||||
return conditioning
|
||||
|
||||
|
||||
class TextEncodeJoyImageEdit(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
image_template = io.Autogrow.TemplatePrefix(
|
||||
io.Image.Input("image"),
|
||||
prefix="image",
|
||||
min=0,
|
||||
max=6,
|
||||
)
|
||||
return io.Schema(
|
||||
node_id="TextEncodeJoyImageEdit",
|
||||
category="model/conditioning/joyimage",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.String.Input("prompt", multiline=True, dynamic_prompts=True),
|
||||
io.Vae.Input("vae", optional=True),
|
||||
io.Autogrow.Input("images", template=image_template, optional=True),
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output(),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, clip, prompt, vae=None, images: io.Autogrow.Type = None) -> io.NodeOutput:
|
||||
images = images or {}
|
||||
return io.NodeOutput(_encode(clip, prompt, vae, list(images.values())))
|
||||
|
||||
|
||||
class JoyImageExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
TextEncodeJoyImageEdit,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> JoyImageExtension:
|
||||
return JoyImageExtension()
|
||||
@@ -61,10 +61,14 @@ class Load3D(IO.ComfyNode):
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model_file, image, **kwargs) -> IO.NodeOutput:
|
||||
image_path = folder_paths.get_annotated_filepath(image['image'])
|
||||
mask_path = folder_paths.get_annotated_filepath(image['mask'])
|
||||
normal_path = folder_paths.get_annotated_filepath(image['normal'])
|
||||
|
||||
load_image_node = nodes.LoadImage()
|
||||
output_image, ignore_mask = load_image_node.load_image(image=image['image'])
|
||||
ignore_image, output_mask = load_image_node.load_image(image=image['mask'])
|
||||
normal_image, ignore_mask2 = load_image_node.load_image(image=image['normal'])
|
||||
output_image, ignore_mask = load_image_node.load_image(image=image_path)
|
||||
ignore_image, output_mask = load_image_node.load_image(image=mask_path)
|
||||
normal_image, ignore_mask2 = load_image_node.load_image(image=normal_path)
|
||||
|
||||
video = None
|
||||
|
||||
@@ -92,7 +96,6 @@ class Preview3D(IO.ComfyNode):
|
||||
search_aliases=["view mesh", "3d viewer"],
|
||||
display_name="Preview 3D & Animation",
|
||||
category="3d",
|
||||
description="Preview a 3D model file without saving it to the ComfyUI output directory.",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
@@ -137,7 +140,6 @@ class Preview3DAdvanced(IO.ComfyNode):
|
||||
display_name="Preview 3D (Advanced)",
|
||||
search_aliases=["preview 3d", "3d viewer", "view mesh", "frame 3d", "3d camera output"],
|
||||
category="3d",
|
||||
description="Preview a 3D model file without saving it to the ComfyUI output directory.",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
@@ -174,9 +176,8 @@ class Preview3DAdvanced(IO.ComfyNode):
|
||||
filename = f"preview3d_advanced_{uuid.uuid4().hex}.{model_3d.format}"
|
||||
model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename))
|
||||
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
camera_info_input = kwargs.get("camera_info", None)
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
|
||||
model_3d_info_input = kwargs.get("model_3d_info", None)
|
||||
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(
|
||||
@@ -196,7 +197,6 @@ class PreviewGaussianSplat(IO.ComfyNode):
|
||||
node_id="PreviewGaussianSplat",
|
||||
display_name="Preview Splat",
|
||||
category="3d",
|
||||
description="Preview a gaussian splat 3D file without saving it to the ComfyUI output directory.",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
search_aliases=[
|
||||
@@ -244,9 +244,8 @@ class PreviewGaussianSplat(IO.ComfyNode):
|
||||
filename = f"preview_splat_{uuid.uuid4().hex}.{model_3d.format}"
|
||||
model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename))
|
||||
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
camera_info_input = kwargs.get("camera_info", None)
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
|
||||
model_3d_info_input = kwargs.get("model_3d_info", None)
|
||||
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(
|
||||
@@ -266,7 +265,6 @@ class PreviewPointCloud(IO.ComfyNode):
|
||||
node_id="PreviewPointCloud",
|
||||
display_name="Preview Point Cloud",
|
||||
category="3d",
|
||||
description="Preview a point cloud 3D file without saving it to the ComfyUI output directory.",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
search_aliases=[
|
||||
@@ -305,9 +303,8 @@ class PreviewPointCloud(IO.ComfyNode):
|
||||
filename = f"preview_pointcloud_{uuid.uuid4().hex}.{model_3d.format}"
|
||||
model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename))
|
||||
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
camera_info_input = kwargs.get("camera_info", None)
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
|
||||
model_3d_info_input = kwargs.get("model_3d_info", None)
|
||||
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(
|
||||
@@ -378,9 +375,8 @@ class Load3DAdvanced(IO.ComfyNode):
|
||||
file_3d = None
|
||||
if model_file and model_file != "none":
|
||||
file_3d = Types.File3D(folder_paths.get_annotated_filepath(model_file))
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
model_3d_info = viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(file_3d, model_3d_info, viewport_state.get('camera_info'), width, height)
|
||||
return IO.NodeOutput(file_3d, model_3d_info, viewport_state['camera_info'], width, height)
|
||||
|
||||
|
||||
class Load3DExtension(ComfyExtension):
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing_extensions import override
|
||||
|
||||
import comfy.sd
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
|
||||
|
||||
def _load_lora_file(lora_name: str):
|
||||
lora_path = folder_paths.get_full_path_or_raise("loras", lora_name)
|
||||
return comfy.utils.load_torch_file(lora_path, safe_load=True, return_metadata=True)
|
||||
|
||||
|
||||
def _lora_template() -> list[io.Input]:
|
||||
return [
|
||||
io.Combo.Input("lora_name", options=folder_paths.get_filename_list("loras"),
|
||||
tooltip="The name of the LoRA file to apply."),
|
||||
io.Float.Input("strength", default=1.0, min=-100.0, max=100.0, step=0.01,
|
||||
tooltip="How strongly to apply this LoRA. 0 = off, negative inverts the effect."),
|
||||
]
|
||||
|
||||
|
||||
class LoadLoraModel(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LoadLoraModel",
|
||||
display_name="Load LoRA (Model)",
|
||||
search_aliases=["lora", "load lora", "apply lora", "lora model", "lora stack"],
|
||||
category="model/loaders",
|
||||
description="Apply a stack of LoRAs to a diffusion model. Add one row per LoRA; "
|
||||
"each row picks a LoRA file and its strength.",
|
||||
inputs=[
|
||||
io.Model.Input("model", tooltip="The diffusion model the LoRAs will be applied to."),
|
||||
io.DynamicGroup.Input(
|
||||
"loras",
|
||||
template=_lora_template(),
|
||||
min=1,
|
||||
max=50,
|
||||
tooltip="Each row applies one LoRA to the model.",
|
||||
group_name="LoRA",
|
||||
),
|
||||
],
|
||||
outputs=[io.Model.Output(tooltip="The modified diffusion model.")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model, loras: list[dict]) -> io.NodeOutput:
|
||||
for row in loras:
|
||||
lora_name = row.get("lora_name")
|
||||
strength = row.get("strength", 1.0)
|
||||
if not lora_name or lora_name == "none" or strength == 0:
|
||||
continue
|
||||
lora, metadata = _load_lora_file(lora_name)
|
||||
model, _ = comfy.sd.load_lora_for_models(model, None, lora, strength, 0, lora_metadata=metadata)
|
||||
return io.NodeOutput(model)
|
||||
|
||||
|
||||
class LoadLoraTextEncoder(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LoadLoraTextEncoder",
|
||||
display_name="Load LoRA (Text Encoder)",
|
||||
search_aliases=["lora", "load lora", "apply lora", "clip lora", "lora stack"],
|
||||
category="model/loaders",
|
||||
description="Apply a stack of LoRAs to a CLIP text encoder. Add one row per LoRA; "
|
||||
"each row picks a LoRA file and its strength.",
|
||||
inputs=[
|
||||
io.Clip.Input("clip", tooltip="The CLIP text encoder the LoRAs will be applied to."),
|
||||
io.DynamicGroup.Input(
|
||||
"loras",
|
||||
template=_lora_template(),
|
||||
min=1,
|
||||
max=50,
|
||||
tooltip="Each row applies one LoRA to the text encoder.",
|
||||
group_name="LoRA",
|
||||
),
|
||||
],
|
||||
outputs=[io.Clip.Output(tooltip="The modified CLIP text encoder.")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, clip, loras: list[dict]) -> io.NodeOutput:
|
||||
for row in loras:
|
||||
lora_name = row.get("lora_name")
|
||||
strength = row.get("strength", 1.0)
|
||||
if not lora_name or lora_name == "none" or strength == 0:
|
||||
continue
|
||||
lora, metadata = _load_lora_file(lora_name)
|
||||
_, clip = comfy.sd.load_lora_for_models(None, clip, lora, 0, strength, lora_metadata=metadata)
|
||||
return io.NodeOutput(clip)
|
||||
|
||||
|
||||
class LoraStackExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
LoadLoraModel,
|
||||
LoadLoraTextEncoder,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> LoraStackExtension:
|
||||
return LoraStackExtension()
|
||||
@@ -419,18 +419,17 @@ class MaskPreview(IO.ComfyNode):
|
||||
search_aliases=["show mask", "view mask", "inspect mask", "debug mask"],
|
||||
display_name="Preview Mask",
|
||||
category="image/mask",
|
||||
description="Preview the masks without saving them to the ComfyUI output directory.",
|
||||
description="Saves the input images to your ComfyUI output directory.",
|
||||
inputs=[
|
||||
IO.Mask.Input("mask"),
|
||||
],
|
||||
hidden=[IO.Hidden.prompt, IO.Hidden.extra_pnginfo],
|
||||
is_output_node=True,
|
||||
outputs=[IO.Mask.Output(display_name="mask")]
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, mask, filename_prefix="ComfyUI") -> IO.NodeOutput:
|
||||
return IO.NodeOutput(mask, ui=UI.PreviewMask(mask))
|
||||
return IO.NodeOutput(ui=UI.PreviewMask(mask))
|
||||
|
||||
|
||||
class MaskExtension(ComfyExtension):
|
||||
|
||||
@@ -18,7 +18,6 @@ class PreviewAny():
|
||||
|
||||
CATEGORY = "utilities"
|
||||
SEARCH_ALIASES = ["show output", "inspect", "debug", "print value", "show text"]
|
||||
DESCRIPTION = "Preview any input value as text."
|
||||
|
||||
def main(self, source=None):
|
||||
torch.set_printoptions(edgeitems=6)
|
||||
|
||||
@@ -10,10 +10,11 @@ class String(io.ComfyNode):
|
||||
return io.Schema(
|
||||
node_id="PrimitiveString",
|
||||
search_aliases=["text", "string", "text box", "prompt"],
|
||||
display_name="Text",
|
||||
display_name="Text String (DEPRECATED)",
|
||||
category="utilities/primitive",
|
||||
inputs=[io.String.Input("value")],
|
||||
outputs=[io.String.Output()]
|
||||
outputs=[io.String.Output()],
|
||||
is_deprecated=True
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -27,7 +28,7 @@ class StringMultiline(io.ComfyNode):
|
||||
return io.Schema(
|
||||
node_id="PrimitiveStringMultiline",
|
||||
search_aliases=["text", "string", "text multiline", "string multiline", "text box", "prompt"],
|
||||
display_name="Text (Multiline)",
|
||||
display_name="Input Text",
|
||||
category="utilities/primitive",
|
||||
essentials_category="Basics",
|
||||
inputs=[io.String.Input("value", multiline=True)],
|
||||
|
||||
@@ -13,7 +13,7 @@ from typing_extensions import override
|
||||
|
||||
import folder_paths
|
||||
from comfy.cli_args import args
|
||||
from comfy_api.latest import ComfyExtension, IO, Types, UI
|
||||
from comfy_api.latest import ComfyExtension, IO, Types
|
||||
|
||||
|
||||
def pack_variable_mesh_batch(vertices, faces, colors=None, uvs=None, texture=None, unlit=False):
|
||||
@@ -406,165 +406,10 @@ class SaveGLB(IO.ComfyNode):
|
||||
return IO.NodeOutput(ui={"3d": results})
|
||||
|
||||
|
||||
def _save_file3d_to_output(model_3d: Types.File3D, filename_prefix: str) -> str:
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
|
||||
filename_prefix, folder_paths.get_output_directory()
|
||||
)
|
||||
ext = model_3d.format or "glb"
|
||||
saved_filename = f"{filename}_{counter:05}.{ext}"
|
||||
model_3d.save_to(os.path.join(full_output_folder, saved_filename))
|
||||
return f"{subfolder}/{saved_filename}" if subfolder else saved_filename
|
||||
|
||||
|
||||
def execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs) -> IO.NodeOutput:
|
||||
model_file = _save_file3d_to_output(model_3d, filename_prefix)
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
camera_info_input = kwargs.get("camera_info", None)
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
|
||||
model_3d_info_input = kwargs.get("model_3d_info", None)
|
||||
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(
|
||||
model_3d,
|
||||
model_3d_info,
|
||||
camera_info,
|
||||
width,
|
||||
height,
|
||||
ui=UI.PreviewUI3DAdvanced(model_file, camera_info, model_3d_info),
|
||||
)
|
||||
|
||||
|
||||
class Save3DAdvanced(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return IO.Schema(
|
||||
node_id="Save3DAdvanced",
|
||||
display_name="Save 3D (Advanced)",
|
||||
search_aliases=["save 3d", "export 3d model", "save mesh advanced"],
|
||||
category="3d",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
IO.MultiType.Input(
|
||||
"model_3d",
|
||||
types=[
|
||||
IO.File3DGLB,
|
||||
IO.File3DGLTF,
|
||||
IO.File3DFBX,
|
||||
IO.File3DOBJ,
|
||||
IO.File3DSTL,
|
||||
IO.File3DUSDZ,
|
||||
IO.File3DAny,
|
||||
],
|
||||
tooltip="3D model file from an upstream 3D node.",
|
||||
),
|
||||
IO.String.Input("filename_prefix", default="3d/ComfyUI"),
|
||||
IO.Load3D.Input("viewport_state"),
|
||||
IO.Load3DModelInfo.Input("model_3d_info", optional=True, advanced=True),
|
||||
IO.Load3DCamera.Input("camera_info", optional=True, advanced=True),
|
||||
IO.Int.Input("width", default=1024, min=1, max=4096, step=1),
|
||||
IO.Int.Input("height", default=1024, min=1, max=4096, step=1),
|
||||
],
|
||||
outputs=[
|
||||
IO.File3DAny.Output(display_name="model_3d"),
|
||||
IO.Load3DModelInfo.Output(display_name="model_3d_info"),
|
||||
IO.Load3DCamera.Output(display_name="camera_info"),
|
||||
IO.Int.Output(display_name="width"),
|
||||
IO.Int.Output(display_name="height"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model_3d: Types.File3D, viewport_state, width: int, height: int, filename_prefix: str, **kwargs) -> IO.NodeOutput:
|
||||
return execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs)
|
||||
|
||||
|
||||
class SaveGaussianSplat(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return IO.Schema(
|
||||
node_id="SaveGaussianSplat",
|
||||
display_name="Save Splat",
|
||||
search_aliases=["save splat", "save gaussian splat", "export gaussian", "export splat"],
|
||||
category="3d",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
IO.MultiType.Input(
|
||||
"model_3d",
|
||||
types=[
|
||||
IO.File3DSplatAny,
|
||||
IO.File3DPLY,
|
||||
IO.File3DSPLAT,
|
||||
IO.File3DSPZ,
|
||||
IO.File3DKSPLAT,
|
||||
],
|
||||
tooltip="A gaussian splat 3D file.",
|
||||
),
|
||||
IO.String.Input("filename_prefix", default="3d/ComfyUI"),
|
||||
IO.Load3D.Input("viewport_state"),
|
||||
IO.Load3DModelInfo.Input("model_3d_info", optional=True, advanced=True),
|
||||
IO.Load3DCamera.Input("camera_info", optional=True, advanced=True),
|
||||
IO.Int.Input("width", default=1024, min=1, max=4096, step=1),
|
||||
IO.Int.Input("height", default=1024, min=1, max=4096, step=1),
|
||||
],
|
||||
outputs=[
|
||||
IO.File3DSplatAny.Output(display_name="model_3d"),
|
||||
IO.Load3DModelInfo.Output(display_name="model_3d_info"),
|
||||
IO.Load3DCamera.Output(display_name="camera_info"),
|
||||
IO.Int.Output(display_name="width"),
|
||||
IO.Int.Output(display_name="height"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model_3d: Types.File3D, viewport_state, width: int, height: int, filename_prefix: str, **kwargs) -> IO.NodeOutput:
|
||||
return execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs)
|
||||
|
||||
|
||||
class SavePointCloud(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return IO.Schema(
|
||||
node_id="SavePointCloud",
|
||||
display_name="Save Point Cloud",
|
||||
search_aliases=["save point cloud", "save pointcloud", "export point cloud"],
|
||||
category="3d",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
IO.MultiType.Input(
|
||||
"model_3d",
|
||||
types=[
|
||||
IO.File3DPointCloudAny,
|
||||
IO.File3DPLY,
|
||||
],
|
||||
tooltip="Point cloud file (.ply)",
|
||||
),
|
||||
IO.String.Input("filename_prefix", default="3d/ComfyUI"),
|
||||
IO.Load3D.Input("viewport_state"),
|
||||
IO.Load3DModelInfo.Input("model_3d_info", optional=True, advanced=True),
|
||||
IO.Load3DCamera.Input("camera_info", optional=True, advanced=True),
|
||||
IO.Int.Input("width", default=1024, min=1, max=4096, step=1),
|
||||
IO.Int.Input("height", default=1024, min=1, max=4096, step=1),
|
||||
],
|
||||
outputs=[
|
||||
IO.File3DPointCloudAny.Output(display_name="model_3d"),
|
||||
IO.Load3DModelInfo.Output(display_name="model_3d_info"),
|
||||
IO.Load3DCamera.Output(display_name="camera_info"),
|
||||
IO.Int.Output(display_name="width"),
|
||||
IO.Int.Output(display_name="height"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model_3d: Types.File3D, viewport_state, width: int, height: int, filename_prefix: str, **kwargs) -> IO.NodeOutput:
|
||||
return execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs)
|
||||
|
||||
|
||||
class Save3DExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[IO.ComfyNode]]:
|
||||
return [SaveGLB, Save3DAdvanced, SaveGaussianSplat, SavePointCloud]
|
||||
return [SaveGLB]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> Save3DExtension:
|
||||
|
||||
@@ -1,614 +0,0 @@
|
||||
import logging
|
||||
|
||||
from typing_extensions import override
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
import torch
|
||||
|
||||
import comfy.model_management
|
||||
from comfy.ldm.seedvr.color_fix import (
|
||||
adain_color_transfer,
|
||||
lab_color_transfer,
|
||||
wavelet_color_transfer,
|
||||
)
|
||||
from comfy.ldm.seedvr.constants import (
|
||||
BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE,
|
||||
SEEDVR2_ADAIN_SCALE_MULTIPLIER,
|
||||
SEEDVR2_CHUNK_GIB_PER_MPX_FRAME,
|
||||
SEEDVR2_CHUNK_RESERVED_GIB,
|
||||
SEEDVR2_CHUNK_SIGMA_GIB,
|
||||
SEEDVR2_CHUNK_SIGMA_K,
|
||||
SEEDVR2_COLOR_MEM_HEADROOM,
|
||||
SEEDVR2_DTYPE_BYTES_FLOOR,
|
||||
SEEDVR2_LAB_SCALE_MULTIPLIER,
|
||||
SEEDVR2_LATENT_CHANNELS,
|
||||
SEEDVR2_OOM_BACKOFF_DIVISOR,
|
||||
SEEDVR2_WAVELET_SCALE_MULTIPLIER,
|
||||
)
|
||||
|
||||
from torchvision.transforms import functional as TVF
|
||||
from torchvision.transforms.functional import InterpolationMode
|
||||
|
||||
|
||||
_SEEDVR2_INVALID_MODEL_MSG_PREFIX = "SeedVR2Conditioning: model object does not match expected SeedVR2 structure"
|
||||
_ATTR_MISSING = object()
|
||||
|
||||
|
||||
def _resolve_seedvr2_diffusion_model(model):
|
||||
inner = getattr(model, "model", _ATTR_MISSING)
|
||||
if inner is _ATTR_MISSING:
|
||||
raise RuntimeError(
|
||||
f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: input has no 'model' attribute "
|
||||
f"(got type {type(model).__name__})."
|
||||
)
|
||||
if inner is None:
|
||||
raise RuntimeError(
|
||||
f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: input.model is None "
|
||||
f"(input type {type(model).__name__})."
|
||||
)
|
||||
diffusion_model = getattr(inner, "diffusion_model", _ATTR_MISSING)
|
||||
if diffusion_model is _ATTR_MISSING:
|
||||
raise RuntimeError(
|
||||
f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: 'model.model' has no "
|
||||
f"'diffusion_model' attribute (got type {type(inner).__name__})."
|
||||
)
|
||||
if diffusion_model is None:
|
||||
raise RuntimeError(
|
||||
f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: 'model.model.diffusion_model' "
|
||||
f"is None (model.model type {type(inner).__name__})."
|
||||
)
|
||||
return diffusion_model
|
||||
|
||||
|
||||
def div_pad(image, factor):
|
||||
height_factor, width_factor = factor
|
||||
height, width = image.shape[-2:]
|
||||
|
||||
pad_height = (height_factor - (height % height_factor)) % height_factor
|
||||
pad_width = (width_factor - (width % width_factor)) % width_factor
|
||||
|
||||
if pad_height == 0 and pad_width == 0:
|
||||
return image
|
||||
|
||||
padding = (0, pad_width, 0, pad_height)
|
||||
return torch.nn.functional.pad(image, padding, mode='constant', value=0.0)
|
||||
|
||||
def cut_videos(videos):
|
||||
t = videos.size(1)
|
||||
if t < 1:
|
||||
raise ValueError("SeedVR2Preprocess expected at least one frame.")
|
||||
if t == 1:
|
||||
return videos
|
||||
if t <= 4:
|
||||
padding = videos[:, -1:].repeat(1, 4 - t + 1, 1, 1, 1)
|
||||
return torch.cat([videos, padding], dim=1)
|
||||
if (t - 1) % 4 == 0:
|
||||
return videos
|
||||
padding = videos[:, -1:].repeat(1, 4 - ((t - 1) % 4), 1, 1, 1)
|
||||
videos = torch.cat([videos, padding], dim=1)
|
||||
if (videos.size(1) - 1) % 4 != 0:
|
||||
raise ValueError(f"SeedVR2Preprocess failed to pad video length to 4n+1; got {videos.size(1)} frames.")
|
||||
return videos
|
||||
|
||||
def _seedvr2_input_shorter_edge(images, node_name):
|
||||
if images.dim() == 4:
|
||||
return min(images.shape[1], images.shape[2])
|
||||
if images.dim() == 5:
|
||||
return min(images.shape[2], images.shape[3])
|
||||
raise ValueError(
|
||||
f"{node_name}: expected 4-D or 5-D IMAGE tensor, "
|
||||
f"got shape {tuple(images.shape)}"
|
||||
)
|
||||
|
||||
|
||||
def _seedvr2_pad(images, upscaled_shorter_edge, node_name):
|
||||
if upscaled_shorter_edge < 2:
|
||||
raise ValueError(
|
||||
f"{node_name}: input shorter edge must be at least 2 pixels; "
|
||||
f"got {upscaled_shorter_edge}."
|
||||
)
|
||||
if images.shape[-1] > 3:
|
||||
images = images[..., :3]
|
||||
if images.dim() == 4:
|
||||
# Comfy video components arrive as a 4-D IMAGE frame sequence:
|
||||
# (frames, H, W, C). SeedVR2 consumes that as one video.
|
||||
images = images.unsqueeze(0)
|
||||
elif images.dim() != 5:
|
||||
raise ValueError(
|
||||
f"{node_name}: expected 4-D or 5-D IMAGE tensor, "
|
||||
f"got shape {tuple(images.shape)}"
|
||||
)
|
||||
images = images.permute(0, 1, 4, 2, 3)
|
||||
|
||||
b, t, c, h, w = images.shape
|
||||
images = images.reshape(b * t, c, h, w)
|
||||
|
||||
images = torch.clamp(images, 0.0, 1.0)
|
||||
images = div_pad(images, (16, 16))
|
||||
_, _, new_h, new_w = images.shape
|
||||
|
||||
images = images.reshape(b, t, c, new_h, new_w)
|
||||
images = cut_videos(images)
|
||||
images_bthwc = images.permute(0, 1, 3, 4, 2).contiguous()
|
||||
|
||||
return io.NodeOutput(images_bthwc)
|
||||
|
||||
|
||||
class SeedVR2Preprocess(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2Preprocess",
|
||||
display_name="Pre-Process SeedVR2 Input",
|
||||
category="image/pre-processors",
|
||||
description="Pad a resized image for SeedVR2 model. Alpha channel is dropped. The node Post-Process SeedVR2 Output re-applies it from the original resized image.",
|
||||
search_aliases=["seedvr2", "upscale", "video upscale", "pad", "preprocess"],
|
||||
inputs=[
|
||||
io.Image.Input("resized_images", tooltip="The resized image to process."),
|
||||
],
|
||||
outputs=[
|
||||
io.Image.Output("images", tooltip="The padded image for VAE encoding."),
|
||||
]
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, resized_images):
|
||||
upscaled_shorter_edge = _seedvr2_input_shorter_edge(resized_images, "SeedVR2Preprocess")
|
||||
return _seedvr2_pad(
|
||||
resized_images, upscaled_shorter_edge, "SeedVR2Preprocess",
|
||||
)
|
||||
|
||||
|
||||
class SeedVR2PostProcessing(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2PostProcessing",
|
||||
display_name="Post-Process SeedVR2 Output",
|
||||
category="image/post-processors",
|
||||
description="Align the generated image with the original resized image and apply color correction.",
|
||||
search_aliases=["seedvr2", "upscale", "color correction", "color match", "postprocess"],
|
||||
inputs=[
|
||||
io.Image.Input("images", tooltip="The generated image to process."),
|
||||
io.Image.Input("original_resized_images", tooltip="The original resized image before pre-processing, used as reference."),
|
||||
io.Combo.Input("color_correction_method", options=["lab", "wavelet", "adain", "none"], default="lab", tooltip="Method to match the generated image colors to the original image. lab: transfer color in CIELAB space, preserving detail (most faithful). wavelet: transfer low-frequency color, keeping upscaled high-frequency detail. adain: match per-channel mean/std (fastest, global tint). none: skip color transfer (geometry alignment only)."),
|
||||
],
|
||||
outputs=[io.Image.Output(display_name="images", tooltip="The aligned, color-corrected image.")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, images, original_resized_images, color_correction_method):
|
||||
alpha_input = None
|
||||
if original_resized_images.shape[-1] == 4:
|
||||
alpha_input = original_resized_images[..., 3:4]
|
||||
original_resized_images = original_resized_images[..., :3]
|
||||
decoded_5d, decoded_was_4d = cls._as_bthwc(images)
|
||||
reference_full, _ = cls._as_bthwc(original_resized_images)
|
||||
decoded_5d = cls._restore_reference_batch_time(decoded_5d, reference_full)
|
||||
|
||||
b = min(decoded_5d.shape[0], reference_full.shape[0])
|
||||
t = min(decoded_5d.shape[1], reference_full.shape[1])
|
||||
reference_h = reference_full.shape[2]
|
||||
reference_w = reference_full.shape[3]
|
||||
|
||||
decoded_5d = decoded_5d[:b, :t, :, :, :]
|
||||
target_h = min(decoded_5d.shape[2], reference_h)
|
||||
target_w = min(decoded_5d.shape[3], reference_w)
|
||||
decoded_5d = decoded_5d[:, :, :target_h, :target_w, :]
|
||||
if color_correction_method in ("lab", "wavelet", "adain"):
|
||||
reference_5d = reference_full[:b, :t, :, :, :]
|
||||
reference_5d = cls._resize_reference(reference_5d, target_h, target_w)
|
||||
output_device = decoded_5d.device
|
||||
decoded_raw = cls._to_seedvr2_raw(decoded_5d)
|
||||
reference_raw = cls._to_seedvr2_raw(reference_5d)
|
||||
decoded_flat = decoded_raw.permute(0, 1, 4, 2, 3).reshape(b * t, decoded_raw.shape[4], target_h, target_w)
|
||||
reference_flat = reference_raw.permute(0, 1, 4, 2, 3).reshape(b * t, reference_raw.shape[4], target_h, target_w)
|
||||
output = cls._color_transfer_chunked(
|
||||
decoded_flat, reference_flat, output_device, color_correction_method,
|
||||
)
|
||||
output = output.reshape(b, t, output.shape[1], output.shape[2], output.shape[3]).permute(0, 1, 3, 4, 2)
|
||||
output = output.add(1.0).div(2.0).clamp(0.0, 1.0)
|
||||
elif color_correction_method == "none":
|
||||
output = decoded_5d
|
||||
else:
|
||||
raise ValueError(f"SeedVR2PostProcessing: unknown color_correction_method {color_correction_method!r}")
|
||||
|
||||
if alpha_input is not None:
|
||||
alpha_5d, _ = cls._as_bthwc(alpha_input)
|
||||
alpha_5d = alpha_5d[:output.shape[0], :output.shape[1], :output.shape[2], :output.shape[3], :]
|
||||
output = torch.cat([output, alpha_5d.to(dtype=output.dtype, device=output.device)], dim=-1)
|
||||
h2 = output.shape[-3] - (output.shape[-3] % 2)
|
||||
w2 = output.shape[-2] - (output.shape[-2] % 2)
|
||||
output = output[:, :, :h2, :w2, :]
|
||||
if decoded_was_4d:
|
||||
output = output.reshape(-1, output.shape[-3], output.shape[-2], output.shape[-1])
|
||||
return io.NodeOutput(output)
|
||||
|
||||
@staticmethod
|
||||
def _as_bthwc(images):
|
||||
if images.ndim == 4:
|
||||
return images.unsqueeze(0), True
|
||||
if images.ndim == 5:
|
||||
return images, False
|
||||
raise ValueError(
|
||||
f"SeedVR2PostProcessing: expected 4-D or 5-D IMAGE tensor, got shape {tuple(images.shape)}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _restore_reference_batch_time(decoded, reference):
|
||||
if decoded.shape[0] != 1:
|
||||
return decoded
|
||||
ref_b, ref_t = reference.shape[:2]
|
||||
if ref_b < 1 or decoded.shape[1] % ref_b != 0:
|
||||
return decoded
|
||||
decoded_t = decoded.shape[1] // ref_b
|
||||
if decoded_t < ref_t:
|
||||
return decoded
|
||||
return decoded.reshape(ref_b, decoded_t, decoded.shape[2], decoded.shape[3], decoded.shape[4])
|
||||
|
||||
@staticmethod
|
||||
def _to_seedvr2_raw(images):
|
||||
return images.mul(2.0).sub(1.0)
|
||||
|
||||
@staticmethod
|
||||
def _color_transfer_on_vae_device(decoded_flat, reference_flat, output_device, transfer_fn):
|
||||
color_device = comfy.model_management.vae_device()
|
||||
decoded_flat = decoded_flat.to(device=color_device)
|
||||
reference_flat = reference_flat.to(device=color_device)
|
||||
output = transfer_fn(decoded_flat, reference_flat)
|
||||
return output.to(device=output_device)
|
||||
|
||||
@staticmethod
|
||||
def _lab_color_transfer_on_vae_device(decoded_flat, reference_flat, output_device):
|
||||
color_device = comfy.model_management.vae_device()
|
||||
result = None
|
||||
for start in range(decoded_flat.shape[0]):
|
||||
decoded_frame = decoded_flat[start:start + 1].to(device=color_device).clone()
|
||||
reference_frame = reference_flat[start:start + 1].to(device=color_device).clone()
|
||||
output = lab_color_transfer(decoded_frame, reference_frame).to(device=output_device)
|
||||
if result is None:
|
||||
result = torch.empty(
|
||||
(decoded_flat.shape[0],) + tuple(output.shape[1:]),
|
||||
device=output_device,
|
||||
dtype=output.dtype,
|
||||
)
|
||||
result[start:start + 1].copy_(output)
|
||||
if result is None:
|
||||
raise ValueError("SeedVR2PostProcessing: LAB color correction requires at least one frame.")
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def _color_transfer_chunked(cls, decoded_flat, reference_flat, output_device, color_correction_method):
|
||||
chunk_size = cls._estimate_color_correction_chunk_size(decoded_flat, color_correction_method)
|
||||
while True:
|
||||
try:
|
||||
return cls._run_color_transfer_chunks(
|
||||
decoded_flat, reference_flat, output_device, color_correction_method, chunk_size,
|
||||
)
|
||||
except Exception as e:
|
||||
comfy.model_management.raise_non_oom(e)
|
||||
if chunk_size <= 1:
|
||||
raise RuntimeError(
|
||||
"SeedVR2PostProcessing: color correction OOM at one frame; "
|
||||
f"color_correction_method={color_correction_method}, shape={tuple(decoded_flat.shape)}."
|
||||
) from e
|
||||
chunk_size = max(1, chunk_size // SEEDVR2_OOM_BACKOFF_DIVISOR)
|
||||
|
||||
@classmethod
|
||||
def _run_color_transfer_chunks(cls, decoded_flat, reference_flat, output_device, color_correction_method, chunk_size):
|
||||
result = None
|
||||
for start in range(0, decoded_flat.shape[0], chunk_size):
|
||||
end = min(start + chunk_size, decoded_flat.shape[0])
|
||||
decoded_chunk = decoded_flat[start:end]
|
||||
reference_chunk = reference_flat[start:end]
|
||||
if color_correction_method == "lab":
|
||||
output = cls._lab_color_transfer_on_vae_device(decoded_chunk, reference_chunk, output_device)
|
||||
elif color_correction_method == "wavelet":
|
||||
output = cls._color_transfer_on_vae_device(
|
||||
decoded_chunk, reference_chunk, output_device, wavelet_color_transfer,
|
||||
)
|
||||
else:
|
||||
output = cls._color_transfer_on_vae_device(
|
||||
decoded_chunk, reference_chunk, output_device, adain_color_transfer,
|
||||
)
|
||||
if result is None:
|
||||
result = torch.empty(
|
||||
(decoded_flat.shape[0],) + tuple(output.shape[1:]),
|
||||
device=output_device,
|
||||
dtype=output.dtype,
|
||||
)
|
||||
result[start:end].copy_(output)
|
||||
if result is None:
|
||||
raise ValueError("SeedVR2PostProcessing: color correction requires at least one frame.")
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def _estimate_color_correction_chunk_size(cls, decoded_flat, color_correction_method):
|
||||
multiplier = cls._color_correction_memory_multiplier(color_correction_method)
|
||||
frames = decoded_flat.shape[0]
|
||||
_, channels, height, width = decoded_flat.shape
|
||||
dtype_bytes = max(decoded_flat.element_size(), SEEDVR2_DTYPE_BYTES_FLOOR)
|
||||
bytes_per_frame = height * width * channels * dtype_bytes * multiplier
|
||||
if bytes_per_frame <= 0:
|
||||
return frames
|
||||
color_device = comfy.model_management.vae_device()
|
||||
free_memory = comfy.model_management.get_free_memory(color_device)
|
||||
chunk_size = int((free_memory * SEEDVR2_COLOR_MEM_HEADROOM) // bytes_per_frame)
|
||||
return max(1, min(frames, chunk_size))
|
||||
|
||||
@staticmethod
|
||||
def _color_correction_memory_multiplier(color_correction_method):
|
||||
if color_correction_method == "lab":
|
||||
return SEEDVR2_LAB_SCALE_MULTIPLIER
|
||||
if color_correction_method == "wavelet":
|
||||
return SEEDVR2_WAVELET_SCALE_MULTIPLIER
|
||||
if color_correction_method == "adain":
|
||||
return SEEDVR2_ADAIN_SCALE_MULTIPLIER
|
||||
raise ValueError(f"SeedVR2PostProcessing: unknown color_correction_method {color_correction_method!r}")
|
||||
|
||||
@staticmethod
|
||||
def _resize_reference(reference, height, width):
|
||||
if reference.shape[2] == height and reference.shape[3] == width:
|
||||
return reference
|
||||
b, t = reference.shape[:2]
|
||||
reference_flat = reference.permute(0, 1, 4, 2, 3).reshape(b * t, reference.shape[4], reference.shape[2], reference.shape[3])
|
||||
resized = TVF.resize(
|
||||
reference_flat,
|
||||
size=(height, width),
|
||||
interpolation=InterpolationMode.BICUBIC,
|
||||
antialias=not (isinstance(reference_flat, torch.Tensor) and reference_flat.device.type == "mps"),
|
||||
)
|
||||
return resized.reshape(b, t, resized.shape[1], height, width).permute(0, 1, 3, 4, 2)
|
||||
|
||||
|
||||
class SeedVR2Conditioning(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2Conditioning",
|
||||
display_name="Apply SeedVR2 Conditioning",
|
||||
category="model/conditioning",
|
||||
description="Build SeedVR2 positive/negative conditioning from a VAE latent.",
|
||||
search_aliases=["seedvr2", "upscale", "conditioning"],
|
||||
inputs=[
|
||||
io.Model.Input("model", tooltip="The SeedVR2 model."),
|
||||
io.Latent.Input("vae_conditioning", display_name="latent"),
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output(display_name="positive", tooltip="The positive conditioning for sampling."),
|
||||
io.Conditioning.Output(display_name="negative", tooltip="The negative conditioning for sampling."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model, vae_conditioning) -> io.NodeOutput:
|
||||
|
||||
vae_conditioning = vae_conditioning["samples"]
|
||||
if vae_conditioning.ndim != 5:
|
||||
raise ValueError(
|
||||
"SeedVR2Conditioning expects a 5-D VAE latent in Comfy "
|
||||
f"channel-first layout; got shape {tuple(vae_conditioning.shape)}."
|
||||
)
|
||||
if vae_conditioning.shape[1] != SEEDVR2_LATENT_CHANNELS:
|
||||
if vae_conditioning.shape[-1] == SEEDVR2_LATENT_CHANNELS:
|
||||
raise ValueError(
|
||||
"SeedVR2Conditioning expects SeedVR2 VAE latents in Comfy "
|
||||
f"channel-first layout (B, {SEEDVR2_LATENT_CHANNELS}, T, H, W); "
|
||||
f"got channel-last shape {tuple(vae_conditioning.shape)}."
|
||||
)
|
||||
raise ValueError(
|
||||
"SeedVR2Conditioning expects SeedVR2 VAE latents with "
|
||||
f"{SEEDVR2_LATENT_CHANNELS} channels; got shape {tuple(vae_conditioning.shape)}."
|
||||
)
|
||||
vae_conditioning = vae_conditioning.movedim(1, -1).contiguous()
|
||||
model = _resolve_seedvr2_diffusion_model(model)
|
||||
pos_cond = model.positive_conditioning
|
||||
neg_cond = model.negative_conditioning
|
||||
|
||||
mask = vae_conditioning.new_ones(vae_conditioning.shape[:-1] + (1,))
|
||||
condition = torch.cat((vae_conditioning, mask), dim=-1)
|
||||
condition = condition.movedim(-1, 1)
|
||||
|
||||
negative = [[neg_cond.unsqueeze(0), {"condition": condition}]]
|
||||
positive = [[pos_cond.unsqueeze(0), {"condition": condition}]]
|
||||
|
||||
return io.NodeOutput(positive, negative)
|
||||
|
||||
def _seedvr2_chunk_crossfade_weights(overlap, device, dtype):
|
||||
"""Descending previous-chunk weights across the overlap (next chunk gets ``1 - w``): a Hann fade over the middle third, flat shoulders on the outer thirds."""
|
||||
ramp = torch.linspace(0.0, 1.0, steps=overlap, device=device, dtype=dtype)
|
||||
ramp = ((ramp - 1.0 / 3.0) / (1.0 / 3.0)).clamp(0.0, 1.0)
|
||||
return 0.5 + 0.5 * torch.cos(torch.pi * ramp)
|
||||
|
||||
|
||||
class SeedVR2TemporalChunk(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2TemporalChunk",
|
||||
display_name="Split SeedVR2 Latent",
|
||||
category="model/latent/batch",
|
||||
description="Split a SeedVR2 video latent into overlapping temporal chunks small enough to sample one at a time within VRAM, wiring latents outputs to both Apply SeedVR2 Conditioning and the sampler latent input before recombining with Merge SeedVR2 Latents.",
|
||||
search_aliases=["seedvr2", "split", "chunk", "temporal", "video upscale", "rebatch"],
|
||||
inputs=[
|
||||
io.Latent.Input("latent", tooltip="The VAE-encoded SeedVR2 latent to split."),
|
||||
io.Int.Input("temporal_overlap", default=0, min=0, max=16384,
|
||||
tooltip="Latent frames shared between adjacent chunks and crossfaded at merge; 0 = no overlap."),
|
||||
io.DynamicCombo.Input("chunking_mode",
|
||||
tooltip="manual = use frames_per_chunk exactly; auto = predict the largest chunk that fits free VRAM.",
|
||||
options=[
|
||||
io.DynamicCombo.Option("auto", []),
|
||||
io.DynamicCombo.Option("manual", [
|
||||
io.Int.Input("frames_per_chunk", default=21, min=1, max=16384, step=4,
|
||||
tooltip="Pixel frames per temporal chunk (4n+1: 1, 5, 9, 13, ...)."),
|
||||
]),
|
||||
]),
|
||||
],
|
||||
outputs=[
|
||||
io.Latent.Output(display_name="latents", is_output_list=True,
|
||||
tooltip="The temporal chunks in sequence order."),
|
||||
io.Int.Output(display_name="temporal_overlap",
|
||||
tooltip="The effective latent-frame overlap between adjacent chunks, for Merge SeedVR2 Latents."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, latent, temporal_overlap, chunking_mode) -> io.NodeOutput:
|
||||
samples = latent["samples"]
|
||||
if samples.ndim != 5:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: expected a 5-D video latent (B, C, T, H, W); "
|
||||
f"got shape {tuple(samples.shape)}."
|
||||
)
|
||||
if samples.shape[1] != SEEDVR2_LATENT_CHANNELS:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: expected {SEEDVR2_LATENT_CHANNELS} latent channels; "
|
||||
f"got shape {tuple(samples.shape)}."
|
||||
)
|
||||
if temporal_overlap < 0:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: temporal_overlap must be >= 0; got {temporal_overlap}."
|
||||
)
|
||||
mode = chunking_mode["chunking_mode"]
|
||||
if mode not in ("auto", "manual"):
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: chunking_mode must be 'auto' or 'manual'; "
|
||||
f"got {mode!r}."
|
||||
)
|
||||
t_latent = samples.shape[2]
|
||||
t_pixel = 4 * (t_latent - 1) + 1
|
||||
|
||||
if mode == "auto":
|
||||
free_gb = comfy.model_management.get_free_memory(
|
||||
comfy.model_management.get_torch_device()) / (1024 ** 3)
|
||||
mpx_per_frame = (samples.shape[0] * samples.shape[3] * samples.shape[4]) * (BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE ** 2) / 1e6
|
||||
budget_gb = free_gb - SEEDVR2_CHUNK_RESERVED_GIB - SEEDVR2_CHUNK_SIGMA_K * SEEDVR2_CHUNK_SIGMA_GIB
|
||||
chunk_latent_max = max(1, int(budget_gb / (SEEDVR2_CHUNK_GIB_PER_MPX_FRAME * mpx_per_frame)))
|
||||
frames_per_chunk = min(4 * (chunk_latent_max - 1) + 1, t_pixel)
|
||||
logging.info(
|
||||
"SeedVR2TemporalChunk auto: free=%.2fGiB, %.2fMpx -> frames_per_chunk=%d (t_pixel=%d).",
|
||||
free_gb, mpx_per_frame, frames_per_chunk, t_pixel,
|
||||
)
|
||||
else:
|
||||
frames_per_chunk = chunking_mode["frames_per_chunk"]
|
||||
if frames_per_chunk < 1 or (frames_per_chunk - 1) % 4 != 0:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: frames_per_chunk must be a 4n+1 pixel-frame count "
|
||||
f"(1, 5, 9, 13, 17, 21, ...); got {frames_per_chunk}."
|
||||
)
|
||||
|
||||
if t_pixel <= frames_per_chunk:
|
||||
return io.NodeOutput([latent], 0)
|
||||
|
||||
chunk_latent = (frames_per_chunk - 1) // 4 + 1
|
||||
temporal_overlap = min(temporal_overlap, chunk_latent - 1)
|
||||
step = chunk_latent - temporal_overlap
|
||||
|
||||
chunks = []
|
||||
for start in range(0, t_latent, step):
|
||||
end = min(start + chunk_latent, t_latent)
|
||||
chunk = latent.copy()
|
||||
chunk["samples"] = samples[:, :, start:end].contiguous()
|
||||
chunks.append(chunk)
|
||||
if end >= t_latent:
|
||||
break
|
||||
return io.NodeOutput(chunks, temporal_overlap)
|
||||
|
||||
|
||||
class SeedVR2TemporalMerge(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2TemporalMerge",
|
||||
display_name="Merge SeedVR2 Latents",
|
||||
category="model/latent/batch",
|
||||
is_input_list=True,
|
||||
description="Recombine sampled SeedVR2 latent temporal chunks into one latent, crossfading each overlap with a Hann window sized by the temporal_overlap wired from Split SeedVR2 Latent.",
|
||||
search_aliases=["seedvr2", "merge", "temporal", "hann", "crossfade"],
|
||||
inputs=[
|
||||
io.Latent.Input("latents", tooltip="The sampled temporal chunks in sequence order."),
|
||||
io.Int.Input("temporal_overlap", default=0, min=0, max=16384, force_input=True,
|
||||
tooltip="The temporal_overlap output of Split SeedVR2 Latent. 0 = plain concatenation."),
|
||||
],
|
||||
outputs=[
|
||||
io.Latent.Output(display_name="latent", tooltip="The recombined full-length latent."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, latents, temporal_overlap) -> io.NodeOutput:
|
||||
temporal_overlap = temporal_overlap[0]
|
||||
if temporal_overlap < 0:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalMerge: temporal_overlap must be >= 0; got {temporal_overlap}."
|
||||
)
|
||||
chunks = [entry["samples"] for entry in latents]
|
||||
first = chunks[0]
|
||||
if first.ndim != 5:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalMerge: expected 5-D video latents (B, C, T, H, W); "
|
||||
f"chunk 0 has shape {tuple(first.shape)}."
|
||||
)
|
||||
for i, chunk in enumerate(chunks[1:], start=1):
|
||||
if chunk.shape[:2] != first.shape[:2] or chunk.shape[3:] != first.shape[3:]:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalMerge: chunk {i} shape {tuple(chunk.shape)} does not "
|
||||
f"match chunk 0 shape {tuple(first.shape)} outside the temporal axis."
|
||||
)
|
||||
if i < len(chunks) - 1 and chunk.shape[2] != first.shape[2]:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalMerge: chunk {i} has {chunk.shape[2]} latent frames but "
|
||||
f"chunk 0 has {first.shape[2]}; only the final chunk may be shorter."
|
||||
)
|
||||
|
||||
out = latents[0].copy()
|
||||
out.pop("noise_mask", None)
|
||||
|
||||
if len(chunks) == 1:
|
||||
out["samples"] = first
|
||||
return io.NodeOutput(out)
|
||||
if temporal_overlap == 0:
|
||||
out["samples"] = torch.cat(chunks, dim=2)
|
||||
return io.NodeOutput(out)
|
||||
|
||||
chunk_latent = first.shape[2]
|
||||
step = chunk_latent - min(temporal_overlap, chunk_latent - 1)
|
||||
t_total = step * (len(chunks) - 1) + chunks[-1].shape[2]
|
||||
b, c, _, h, w = first.shape
|
||||
merged = torch.empty((b, c, t_total, h, w), device=first.device, dtype=first.dtype)
|
||||
|
||||
merged[:, :, :chunk_latent] = first
|
||||
filled = chunk_latent
|
||||
for i, chunk in enumerate(chunks[1:], start=1):
|
||||
start = i * step
|
||||
end = start + chunk.shape[2]
|
||||
# Crossfade width is bounded by the previous fill frontier and by a runt
|
||||
# final chunk shorter than the configured overlap.
|
||||
fade = min(filled - start, chunk.shape[2])
|
||||
if fade > 0:
|
||||
w_prev = _seedvr2_chunk_crossfade_weights(
|
||||
fade, chunk.device, chunk.dtype).view(1, 1, fade, 1, 1)
|
||||
merged[:, :, start:start + fade] = (
|
||||
merged[:, :, start:start + fade] * w_prev + chunk[:, :, :fade] * (1.0 - w_prev)
|
||||
)
|
||||
merged[:, :, start + fade:end] = chunk[:, :, fade:]
|
||||
else:
|
||||
merged[:, :, start:end] = chunk
|
||||
filled = end
|
||||
|
||||
out["samples"] = merged
|
||||
return io.NodeOutput(out)
|
||||
|
||||
|
||||
class SeedVRExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
SeedVR2Conditioning,
|
||||
SeedVR2Preprocess,
|
||||
SeedVR2PostProcessing,
|
||||
SeedVR2TemporalChunk,
|
||||
SeedVR2TemporalMerge,
|
||||
]
|
||||
|
||||
async def comfy_entrypoint() -> SeedVRExtension:
|
||||
return SeedVRExtension()
|
||||
@@ -1,71 +0,0 @@
|
||||
import os
|
||||
import json
|
||||
from typing_extensions import override
|
||||
from comfy_api.latest import io, ComfyExtension, ui
|
||||
import folder_paths
|
||||
|
||||
|
||||
class SaveTextNode(io.ComfyNode):
|
||||
"""Save text content to .txt, .md, or .json."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SaveText",
|
||||
search_aliases=["save text", "write text", "export text"],
|
||||
display_name="Save Text",
|
||||
category="text",
|
||||
description="Save text content to a file in the output directory.",
|
||||
inputs=[
|
||||
io.String.Input("text", force_input=True),
|
||||
io.String.Input("filename_prefix", default="ComfyUI"),
|
||||
io.Combo.Input("format", options=["txt", "md", "json"], default="txt"),
|
||||
],
|
||||
outputs=[io.String.Output(display_name="text")],
|
||||
is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, text, filename_prefix, format):
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
|
||||
filename_prefix,
|
||||
folder_paths.get_output_directory(),
|
||||
1,
|
||||
1,
|
||||
)
|
||||
|
||||
file = f"{filename}_{counter:05}.{format}"
|
||||
filepath = os.path.join(full_output_folder, file)
|
||||
|
||||
if format == "json":
|
||||
# tries to pretty print otherwise saves normally
|
||||
try:
|
||||
data = json.loads(text)
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, indent=2, ensure_ascii=False)
|
||||
except json.JSONDecodeError:
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
f.write(text)
|
||||
else:
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
f.write(text)
|
||||
|
||||
return io.NodeOutput(
|
||||
text,
|
||||
ui={
|
||||
"text": (text,),
|
||||
"files": [
|
||||
ui.SavedResult(file, subfolder, io.FolderType.output)
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
class TextExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
SaveTextNode
|
||||
]
|
||||
|
||||
async def comfy_entrypoint() -> TextExtension:
|
||||
return TextExtension()
|
||||
@@ -81,7 +81,7 @@ class SaveVideo(io.ComfyNode):
|
||||
display_name="Save Video",
|
||||
category="video",
|
||||
essentials_category="Basics",
|
||||
description="Saves the input videos to your ComfyUI output directory.",
|
||||
description="Saves the input images to your ComfyUI output directory.",
|
||||
inputs=[
|
||||
io.Video.Input("video", tooltip="The video to save."),
|
||||
io.String.Input("filename_prefix", default="video/ComfyUI", tooltip="The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes."),
|
||||
|
||||
+1
-1
@@ -1,3 +1,3 @@
|
||||
# This file is automatically generated by the build process when version is
|
||||
# updated in pyproject.toml.
|
||||
__version__ = "0.28.0"
|
||||
__version__ = "0.27.0"
|
||||
|
||||
+8
-13
@@ -29,7 +29,6 @@ from comfy_execution.caching import (
|
||||
HierarchicalCache,
|
||||
LRUCache,
|
||||
RAMPressureCache,
|
||||
RAM_CACHE_LARGE_INTERMEDIATE,
|
||||
)
|
||||
from comfy_execution.graph import (
|
||||
DynamicPrompt,
|
||||
@@ -426,12 +425,12 @@ def _is_intermediate_output(dynprompt, node_id):
|
||||
|
||||
|
||||
def _send_cached_ui(server, node_id, display_node_id, cached, prompt_id, ui_outputs):
|
||||
if cached.ui is not None:
|
||||
ui_outputs[node_id] = cached.ui
|
||||
if server.client_id is None:
|
||||
return
|
||||
cached_ui = cached.ui or {}
|
||||
server.send_sync("executed", { "node": node_id, "display_node": display_node_id, "output": cached_ui.get("output", None), "prompt_id": prompt_id }, server.client_id)
|
||||
if cached.ui is not None:
|
||||
ui_outputs[node_id] = cached.ui
|
||||
|
||||
async def execute(server, dynprompt, caches, current_item, extra_data, executed, prompt_id, execution_list, pending_subgraph_results, pending_async_nodes, ui_outputs):
|
||||
unique_id = current_item
|
||||
@@ -795,16 +794,12 @@ class PromptExecutor:
|
||||
if self.cache_type == CacheType.RAM_PRESSURE:
|
||||
ram_release_callback(ram_inactive_headroom)
|
||||
ram_shortfall = ram_headroom - psutil.virtual_memory().available
|
||||
if ram_shortfall > 0:
|
||||
freed = ram_release_callback(ram_headroom, free_active=True, min_entry_size=RAM_CACHE_LARGE_INTERMEDIATE)
|
||||
ram_shortfall -= freed
|
||||
if comfy.model_management.should_free_pins_for_ram_pressure(ram_shortfall):
|
||||
freed = comfy.model_management.free_pins(ram_shortfall + 512 * (1024 ** 2))
|
||||
if freed < ram_shortfall:
|
||||
if freed > 64 * (1024 ** 2):
|
||||
# AIMDO MEM_DECOMMIT can outrun psutil.available catching up.
|
||||
time.sleep(0.05)
|
||||
ram_release_callback(ram_headroom, free_active=True)
|
||||
freed = comfy.model_management.free_pins(ram_shortfall + 512 * (1024 ** 2))
|
||||
if freed < ram_shortfall:
|
||||
if freed > 64 * (1024 ** 2):
|
||||
# AIMDO MEM_DECOMMIT can outrun psutil.available catching up.
|
||||
time.sleep(0.05)
|
||||
ram_release_callback(ram_headroom, free_active=True)
|
||||
else:
|
||||
# Only execute when the while-loop ends without break
|
||||
# Send cached UI for intermediate output nodes that weren't executed
|
||||
|
||||
@@ -992,7 +992,7 @@ class CLIPLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "clip_name": (folder_paths.get_filename_list("text_encoders"), ),
|
||||
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage"], ),
|
||||
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2"], ),
|
||||
},
|
||||
"optional": {
|
||||
"device": (["default", "cpu"], {"advanced": True}),
|
||||
@@ -1002,7 +1002,7 @@ class CLIPLoader:
|
||||
|
||||
CATEGORY = "model/loaders"
|
||||
|
||||
DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\njoyimage: qwen3-vl 8B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm"
|
||||
DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm"
|
||||
|
||||
def load_clip(self, clip_name, type="stable_diffusion", device="default"):
|
||||
clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION)
|
||||
@@ -1709,7 +1709,6 @@ class PreviewImage(SaveImage):
|
||||
self.compress_level = 1
|
||||
|
||||
SEARCH_ALIASES = ["preview", "preview image", "show image", "view image", "display image", "image viewer"]
|
||||
DESCRIPTION = "Preview the images without saving them to the ComfyUI output directory."
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -2459,10 +2458,8 @@ async def init_builtin_extra_nodes():
|
||||
"nodes_camera_trajectory.py",
|
||||
"nodes_edit_model.py",
|
||||
"nodes_tcfg.py",
|
||||
"nodes_seedvr.py",
|
||||
"nodes_context_windows.py",
|
||||
"nodes_qwen.py",
|
||||
"nodes_joyimage.py",
|
||||
"nodes_boogu.py",
|
||||
"nodes_chroma_radiance.py",
|
||||
"nodes_pid.py",
|
||||
@@ -2506,7 +2503,7 @@ async def init_builtin_extra_nodes():
|
||||
"nodes_triposplat.py",
|
||||
"nodes_depth_anything_3.py",
|
||||
"nodes_seed.py",
|
||||
"nodes_text.py",
|
||||
"nodes_lora_stack.py",
|
||||
]
|
||||
|
||||
import_failed = []
|
||||
|
||||
+16
-88
@@ -7,18 +7,18 @@ components:
|
||||
description: Timestamp when the asset was created
|
||||
format: date-time
|
||||
type: string
|
||||
display_name:
|
||||
description: Display name of the asset. Mirrors name for backwards compatibility.
|
||||
nullable: true
|
||||
type: string
|
||||
file_path:
|
||||
description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors")
|
||||
nullable: true
|
||||
type: string
|
||||
hash:
|
||||
description: Blake3 hash of the asset content.
|
||||
pattern: ^blake3:[a-f0-9]{64}$
|
||||
type: string
|
||||
loader_path:
|
||||
description: The value a loader consumes to load this asset. Null when no loader can resolve the file.
|
||||
nullable: true
|
||||
type: string
|
||||
display_name:
|
||||
description: Human-facing label for the asset. Not unique.
|
||||
nullable: true
|
||||
type: string
|
||||
id:
|
||||
description: Unique identifier for the asset
|
||||
format: uuid
|
||||
@@ -144,14 +144,6 @@ components:
|
||||
AssetUpdated:
|
||||
description: Response returned when an existing asset is successfully updated.
|
||||
properties:
|
||||
display_name:
|
||||
description: Display name of the asset. Mirrors name for backwards compatibility.
|
||||
nullable: true
|
||||
type: string
|
||||
file_path:
|
||||
description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors")
|
||||
nullable: true
|
||||
type: string
|
||||
hash:
|
||||
description: Blake3 hash of the asset content.
|
||||
pattern: ^blake3:[a-f0-9]{64}$
|
||||
@@ -1644,7 +1636,7 @@ paths:
|
||||
format: uuid
|
||||
type: string
|
||||
tags:
|
||||
description: JSON-encoded array of freeform tag strings, e.g. '["models","checkpoint"]'. Common types include "models", "input", "output", and "temp", but any tag can be used in any order.
|
||||
description: JSON-encoded array of tag strings. For new byte uploads, include exactly one destination role (`input`, `output`, or `models`); `models` uploads also require exactly one `model_type:<folder_name>` tag. Extra tags are stored as labels and do not create path components.
|
||||
type: string
|
||||
user_metadata:
|
||||
description: Custom JSON metadata as a string
|
||||
@@ -1829,7 +1821,7 @@ paths:
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/AssetUpdated'
|
||||
$ref: '#/components/schemas/Asset'
|
||||
description: Asset updated successfully
|
||||
"400":
|
||||
content:
|
||||
@@ -2470,6 +2462,9 @@ paths:
|
||||
supports_preview_metadata:
|
||||
description: Whether the server supports preview metadata
|
||||
type: boolean
|
||||
supports_model_type_tags:
|
||||
description: Whether the server supports namespaced model type asset tags
|
||||
type: boolean
|
||||
type: object
|
||||
description: Success
|
||||
headers:
|
||||
@@ -3297,12 +3292,6 @@ paths:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Invalid request parameters
|
||||
"401":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Unauthorized - Authentication required
|
||||
"500":
|
||||
content:
|
||||
application/json:
|
||||
@@ -3808,24 +3797,7 @@ paths:
|
||||
/api/upload/mask:
|
||||
post:
|
||||
description: |
|
||||
Upload a mask image to be composited onto an existing image server-side.
|
||||
|
||||
The uploaded mask's alpha channel replaces the original image's alpha.
|
||||
If the mask resolution differs from the original (e.g. the mask was
|
||||
painted on a downscaled preview obtained via /api/view's preview and
|
||||
max_size parameters), the mask alpha is upscaled to the original's
|
||||
resolution before compositing, so the result keeps the original's full
|
||||
resolution.
|
||||
|
||||
An optional paint layer can be supplied to additionally produce painted
|
||||
composites: 'painted' (original + paint) and 'painted_masked'
|
||||
(original + paint + mask alpha), saved next to the main output under
|
||||
the provided filenames.
|
||||
|
||||
Clients that already have the fully composited result client-side
|
||||
should upload it as a plain image via /api/upload/image instead;
|
||||
this endpoint exists for server-side compositing at original
|
||||
resolution.
|
||||
Upload a mask image to be applied to an existing image.
|
||||
|
||||
Image limits apply to both the uploaded mask and the referenced
|
||||
original image:
|
||||
@@ -3841,25 +3813,12 @@ paths:
|
||||
schema:
|
||||
properties:
|
||||
image:
|
||||
description: The mask image file to upload; its alpha channel is composited onto the original
|
||||
description: The mask image file to upload
|
||||
format: binary
|
||||
type: string
|
||||
original_ref:
|
||||
description: JSON string containing reference to the original image
|
||||
type: string
|
||||
paint:
|
||||
description: Optional RGBA paint-stroke layer to composite over the original
|
||||
format: binary
|
||||
type: string
|
||||
paint_filename:
|
||||
description: Optional filename to save the (upscaled) paint layer as, next to the main output
|
||||
type: string
|
||||
painted_filename:
|
||||
description: Optional filename to save the original+paint composite as
|
||||
type: string
|
||||
painted_masked_filename:
|
||||
description: Optional filename to save the original+paint+mask composite as
|
||||
type: string
|
||||
required:
|
||||
- image
|
||||
- original_ref
|
||||
@@ -4480,32 +4439,6 @@ paths:
|
||||
maximum: 1024
|
||||
minimum: 64
|
||||
type: integer
|
||||
- description: |
|
||||
Compressed preview request in the form "<format>[;<quality>]", e.g. "webp;90".
|
||||
Format may be webp or jpeg; requests that need the alpha channel (channel
|
||||
containing 'a') are forced to an alpha-capable format. Quality is an integer
|
||||
from 1 to 100 (defaults to 90); values outside that range are ignored.
|
||||
Takes precedence over 'res' and 'channel' processing.
|
||||
in: query
|
||||
name: preview
|
||||
schema:
|
||||
example: webp;90
|
||||
type: string
|
||||
- description: |
|
||||
Requests a downscaled compressed preview, giving the maximum dimension
|
||||
(width or height) while preserving aspect ratio. Larger images are
|
||||
downscaled to fit; smaller images are never upscaled. Any integer is
|
||||
accepted and clamped to 512-8192; omit the parameter for no downscale.
|
||||
Only used together with 'preview'. Requires the assets system
|
||||
(--enable-assets): generated previews are keyed by the source's content
|
||||
hash and registered as assets linked to the source via preview_id. When
|
||||
assets are disabled, max_size is ignored and the image is only
|
||||
recompressed at full resolution.
|
||||
in: query
|
||||
name: max_size
|
||||
schema:
|
||||
example: 4096
|
||||
type: integer
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
@@ -4519,12 +4452,7 @@ paths:
|
||||
description: Processed PNG image with extracted channel
|
||||
format: binary
|
||||
type: string
|
||||
image/webp:
|
||||
schema:
|
||||
description: Compressed preview (returned when the preview parameter is used)
|
||||
format: binary
|
||||
type: string
|
||||
description: Success - File content returned (used when channel, res, or preview parameter is present)
|
||||
description: Success - File content returned (used when channel or res parameter is present)
|
||||
"302":
|
||||
description: Redirect to GCS signed URL
|
||||
headers:
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "ComfyUI"
|
||||
version = "0.28.0"
|
||||
version = "0.27.0"
|
||||
readme = "README.md"
|
||||
license = { file = "LICENSE" }
|
||||
requires-python = ">=3.10"
|
||||
|
||||
+4
-4
@@ -1,6 +1,6 @@
|
||||
comfyui-frontend-package==1.45.21
|
||||
comfyui-workflow-templates==0.11.9
|
||||
comfyui-embedded-docs==0.5.8
|
||||
comfyui-frontend-package==1.45.20
|
||||
comfyui-workflow-templates==0.11.6
|
||||
comfyui-embedded-docs==0.5.7
|
||||
torch
|
||||
torchsde
|
||||
torchvision
|
||||
@@ -22,7 +22,7 @@ alembic
|
||||
SQLAlchemy>=2.0.0
|
||||
filelock
|
||||
av>=16.0.0
|
||||
comfy-kitchen==0.2.20
|
||||
comfy-kitchen==0.2.16
|
||||
comfy-aimdo==0.4.10
|
||||
requests
|
||||
simpleeval>=1.0.0
|
||||
|
||||
@@ -49,7 +49,6 @@ from app.assets.api.routes import register_assets_routes
|
||||
from app.assets.services.ingest import register_file_in_place
|
||||
from app.assets.services.path_utils import get_known_subfolder_tags
|
||||
from app.assets.services.asset_management import resolve_hash_to_path
|
||||
from app.assets.services.preview import get_or_create_preview_file
|
||||
|
||||
from app.user_manager import UserManager
|
||||
from app.model_manager import ModelFileManager
|
||||
@@ -474,104 +473,45 @@ class PromptServer():
|
||||
|
||||
def image_save_function(image, post, filepath):
|
||||
original_ref = json.loads(post.get("original_ref"))
|
||||
ref_filename = original_ref['filename']
|
||||
filename, output_dir = folder_paths.annotated_filepath(original_ref['filename'])
|
||||
|
||||
if ref_filename.startswith("blake3:"):
|
||||
owner_id = self.user_manager.get_request_user_id(request)
|
||||
result = resolve_hash_to_path(ref_filename, owner_id=owner_id)
|
||||
if result is None:
|
||||
raise web.HTTPBadRequest()
|
||||
file = result.abs_path
|
||||
else:
|
||||
filename, output_dir = folder_paths.annotated_filepath(ref_filename)
|
||||
if not filename:
|
||||
return web.Response(status=400)
|
||||
|
||||
if not filename:
|
||||
raise web.HTTPBadRequest()
|
||||
# validation for security: prevent accessing arbitrary path
|
||||
if filename[0] == '/' or '..' in filename:
|
||||
return web.Response(status=400)
|
||||
|
||||
# validation for security: prevent accessing arbitrary path
|
||||
if filename[0] == '/' or '..' in filename:
|
||||
raise web.HTTPBadRequest()
|
||||
if output_dir is None:
|
||||
type = original_ref.get("type", "output")
|
||||
output_dir = folder_paths.get_directory_by_type(type)
|
||||
|
||||
if output_dir is None:
|
||||
type = original_ref.get("type", "output")
|
||||
output_dir = folder_paths.get_directory_by_type(type)
|
||||
if output_dir is None:
|
||||
return web.Response(status=400)
|
||||
|
||||
if output_dir is None:
|
||||
raise web.HTTPBadRequest()
|
||||
if original_ref.get("subfolder", "") != "":
|
||||
full_output_dir = os.path.join(output_dir, original_ref["subfolder"])
|
||||
if os.path.commonpath((os.path.abspath(full_output_dir), output_dir)) != output_dir:
|
||||
return web.Response(status=403)
|
||||
output_dir = full_output_dir
|
||||
|
||||
if original_ref.get("subfolder", "") != "":
|
||||
full_output_dir = os.path.join(output_dir, original_ref["subfolder"])
|
||||
if os.path.commonpath((os.path.abspath(full_output_dir), output_dir)) != output_dir:
|
||||
raise web.HTTPForbidden()
|
||||
output_dir = full_output_dir
|
||||
file = os.path.join(output_dir, filename)
|
||||
|
||||
file = os.path.join(output_dir, filename)
|
||||
if os.path.isfile(file):
|
||||
with Image.open(file) as original_pil:
|
||||
metadata = PngInfo()
|
||||
if hasattr(original_pil,'text'):
|
||||
for key in original_pil.text:
|
||||
metadata.add_text(key, original_pil.text[key])
|
||||
original_pil = original_pil.convert('RGBA')
|
||||
mask_pil = Image.open(image.file).convert('RGBA')
|
||||
|
||||
if not os.path.isfile(file):
|
||||
raise web.HTTPBadRequest()
|
||||
# alpha copy
|
||||
new_alpha = mask_pil.getchannel('A')
|
||||
original_pil.putalpha(new_alpha)
|
||||
original_pil.save(filepath, compress_level=4, pnginfo=metadata)
|
||||
|
||||
with Image.open(file) as original_pil:
|
||||
metadata = PngInfo()
|
||||
if hasattr(original_pil,'text'):
|
||||
for key in original_pil.text:
|
||||
metadata.add_text(key, original_pil.text[key])
|
||||
|
||||
original_pil = ImageOps.exif_transpose(original_pil.convert('RGBA'))
|
||||
original_pil.putalpha(Image.new('L', original_pil.size, 255))
|
||||
mask_pil = Image.open(image.file).convert('RGBA')
|
||||
|
||||
# alpha copy; the mask may come from a downscaled
|
||||
# preview edit, so upscale it to the original size
|
||||
new_alpha = mask_pil.getchannel('A')
|
||||
if new_alpha.size != original_pil.size:
|
||||
new_alpha = new_alpha.resize(original_pil.size, Image.Resampling.LANCZOS)
|
||||
|
||||
masked_pil = original_pil.copy()
|
||||
masked_pil.putalpha(new_alpha)
|
||||
outputs = [(filepath, masked_pil, {"pnginfo": metadata})]
|
||||
|
||||
paint = post.get("paint")
|
||||
if paint is not None and paint.file:
|
||||
save_dir = os.path.dirname(filepath)
|
||||
|
||||
paint_pil = Image.open(paint.file).convert('RGBA')
|
||||
if paint_pil.size != original_pil.size:
|
||||
paint_pil = paint_pil.resize(original_pil.size, Image.Resampling.LANCZOS)
|
||||
painted_pil = Image.alpha_composite(original_pil, paint_pil)
|
||||
painted_masked_pil = painted_pil.copy()
|
||||
painted_masked_pil.putalpha(new_alpha)
|
||||
|
||||
sibling_outputs = [
|
||||
("paint_filename", paint_pil, {}),
|
||||
("painted_filename", painted_pil, {"pnginfo": metadata}),
|
||||
("painted_masked_filename", painted_masked_pil, {"pnginfo": metadata}),
|
||||
]
|
||||
for field, pil_image, save_kwargs in sibling_outputs:
|
||||
name = post.get(field)
|
||||
if not name:
|
||||
continue
|
||||
outputs.append((os.path.join(save_dir, os.path.basename(name)), pil_image, save_kwargs))
|
||||
|
||||
staged = []
|
||||
try:
|
||||
for dest, pil_image, save_kwargs in outputs:
|
||||
tmp = f"{dest}.{uuid.uuid4().hex}.tmp"
|
||||
pil_image.save(tmp, format="PNG", compress_level=4, **save_kwargs)
|
||||
staged.append((tmp, dest))
|
||||
for tmp, dest in staged:
|
||||
os.replace(tmp, dest)
|
||||
except Exception:
|
||||
for tmp, _ in staged:
|
||||
if os.path.exists(tmp):
|
||||
try:
|
||||
os.remove(tmp)
|
||||
except OSError:
|
||||
pass
|
||||
raise
|
||||
|
||||
# Compositing large originals can take seconds; keep it off the event loop
|
||||
return await asyncio.get_running_loop().run_in_executor(
|
||||
None, image_upload, post, image_save_function)
|
||||
return image_upload(post, image_save_function)
|
||||
|
||||
@routes.get("/view")
|
||||
async def view_image(request):
|
||||
@@ -617,68 +557,24 @@ class PromptServer():
|
||||
|
||||
if os.path.isfile(file):
|
||||
if 'preview' in request.rel_url.query:
|
||||
preview_info = request.rel_url.query['preview'].split(';')
|
||||
channel = request.rel_url.query.get('channel', '')
|
||||
image_format = preview_info[0]
|
||||
if image_format not in ['webp', 'jpeg'] or 'a' in channel:
|
||||
image_format = 'webp'
|
||||
with Image.open(file) as img:
|
||||
preview_info = request.rel_url.query['preview'].split(';')
|
||||
image_format = preview_info[0]
|
||||
if image_format not in ['webp', 'jpeg'] or 'a' in request.rel_url.query.get('channel', ''):
|
||||
image_format = 'webp'
|
||||
|
||||
quality = 90
|
||||
if preview_info[-1].isdigit():
|
||||
parsed_quality = int(preview_info[-1])
|
||||
if 1 <= parsed_quality <= 100:
|
||||
quality = parsed_quality
|
||||
quality = 90
|
||||
if preview_info[-1].isdigit():
|
||||
quality = int(preview_info[-1])
|
||||
|
||||
# A max_size (any integer) requests downscaling and is
|
||||
# clamped into the permitted range; absent or non-integer
|
||||
# means no downscale.
|
||||
max_size = None
|
||||
max_size_param = request.rel_url.query.get('max_size', '')
|
||||
if max_size_param:
|
||||
try:
|
||||
max_size = min(8192, max(512, int(max_size_param)))
|
||||
except ValueError:
|
||||
max_size = None
|
||||
buffer = BytesIO()
|
||||
if image_format in ['jpeg'] or request.rel_url.query.get('channel', '') == 'rgb':
|
||||
img = img.convert("RGB")
|
||||
img.save(buffer, format=image_format, quality=quality)
|
||||
buffer.seek(0)
|
||||
|
||||
safe_filename = filename.replace("\\", "\\\\").replace('"', '\\"')
|
||||
preview_headers = {"Content-Disposition": f'filename="{safe_filename}"'}
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
preview_file = None
|
||||
preview_failed = False
|
||||
if max_size is not None and args.enable_assets:
|
||||
try:
|
||||
preview_file = await loop.run_in_executor(
|
||||
None, get_or_create_preview_file, file, max_size, quality)
|
||||
except Exception:
|
||||
logging.warning("Failed to generate preview asset, downscaling in memory without caching", exc_info=True)
|
||||
preview_failed = True
|
||||
|
||||
needs_convert = image_format == 'jpeg' or channel == 'rgb'
|
||||
if preview_file is not None and not needs_convert:
|
||||
return web.FileResponse(preview_file, headers={
|
||||
**preview_headers,
|
||||
"Content-Type": "image/webp",
|
||||
})
|
||||
|
||||
render_source = preview_file or file
|
||||
|
||||
def render_preview():
|
||||
with Image.open(render_source) as img:
|
||||
preview_img = img
|
||||
|
||||
if preview_failed and max_size is not None and max(img.size) > max_size:
|
||||
preview_img = ImageOps.contain(
|
||||
img, (max_size, max_size), Image.Resampling.LANCZOS)
|
||||
if needs_convert:
|
||||
preview_img = preview_img.convert("RGB")
|
||||
buffer = BytesIO()
|
||||
preview_img.save(buffer, format=image_format, quality=quality)
|
||||
return buffer.getvalue()
|
||||
|
||||
body = await loop.run_in_executor(None, render_preview)
|
||||
return web.Response(body=body, content_type=f'image/{image_format}',
|
||||
headers=preview_headers)
|
||||
return web.Response(body=buffer.read(), content_type=f'image/{image_format}',
|
||||
headers={"Content-Disposition": f"filename=\"{filename}\""})
|
||||
|
||||
if 'channel' not in request.rel_url.query:
|
||||
channel = 'rgba'
|
||||
|
||||
@@ -24,28 +24,6 @@ def app(model_manager):
|
||||
app.add_routes(routes)
|
||||
return app
|
||||
|
||||
async def test_get_model_folders_includes_registered_extensions(aiohttp_client, app, tmp_path):
|
||||
"""Folders expose their registered extension set verbatim; an empty list
|
||||
means match-all (filter_files_extensions semantics)."""
|
||||
with patch('folder_paths.folder_names_and_paths', {
|
||||
'test_checkpoints': ([str(tmp_path)], {'.safetensors', '.ckpt'}),
|
||||
'test_configs': ([str(tmp_path)], ['.yaml']),
|
||||
'test_match_all': ([str(tmp_path)], set()),
|
||||
'configs': ([str(tmp_path)], ['.yaml']),
|
||||
}):
|
||||
client = await aiohttp_client(app)
|
||||
response = await client.get('/experiment/models')
|
||||
|
||||
assert response.status == 200
|
||||
folders = {f['name']: f for f in await response.json()}
|
||||
|
||||
assert 'configs' not in folders # blocklisted
|
||||
assert folders['test_checkpoints']['folders'] == [str(tmp_path)]
|
||||
assert folders['test_checkpoints']['extensions'] == ['.ckpt', '.safetensors']
|
||||
assert folders['test_configs']['extensions'] == ['.yaml']
|
||||
# Match-all registrations are exposed honestly, not substituted.
|
||||
assert folders['test_match_all']['extensions'] == []
|
||||
|
||||
async def test_get_model_preview_safetensors(aiohttp_client, app, tmp_path):
|
||||
img = Image.new('RGB', (100, 100), 'white')
|
||||
img_byte_arr = BytesIO()
|
||||
|
||||
@@ -0,0 +1,204 @@
|
||||
"""Unit tests for io.DynamicGroup: expansion/reconstruction (0-row and N-row cases)."""
|
||||
import sys
|
||||
import types
|
||||
import pytest
|
||||
|
||||
# Stub torch (type-hint only in _io.py; real torch not available in unit-test env)
|
||||
if "torch" not in sys.modules:
|
||||
_torch_stub = types.ModuleType("torch")
|
||||
_torch_stub.Tensor = object # type: ignore[attr-defined]
|
||||
sys.modules["torch"] = _torch_stub
|
||||
|
||||
from comfy_api.latest._io import ( # noqa: E402
|
||||
DynamicGroup,
|
||||
Float,
|
||||
Int,
|
||||
String,
|
||||
Boolean,
|
||||
get_finalized_class_inputs,
|
||||
build_nested_inputs,
|
||||
create_input_dict_v1,
|
||||
setup_dynamic_input_funcs,
|
||||
)
|
||||
|
||||
# Make sure dynamic input funcs are registered (may already be done at import time)
|
||||
setup_dynamic_input_funcs()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_class_inputs(group_input: DynamicGroup.Input) -> dict:
|
||||
"""Wrap a DynamicGroup.Input into the required/optional dict structure."""
|
||||
return create_input_dict_v1([group_input])
|
||||
|
||||
|
||||
def _run(group_input: DynamicGroup.Input, live_values: dict) -> dict:
|
||||
"""End-to-end helper: expand schema + reconstruct values.
|
||||
|
||||
Mirrors the production split in execution.py:
|
||||
1. get_finalized_class_inputs (schema expansion, line 162)
|
||||
2. build_nested_inputs (value reconstruction, line 281)
|
||||
|
||||
The two steps are separate in production because the engine resolves
|
||||
linked node outputs between them, but in tests we supply values directly.
|
||||
"""
|
||||
class_inputs = _make_class_inputs(group_input)
|
||||
_, _, v3_data = get_finalized_class_inputs(class_inputs, live_values)
|
||||
return build_nested_inputs(dict(live_values), v3_data)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Schema construction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestDynamicGroupInputConstruction:
|
||||
def test_basic_construction(self):
|
||||
inp = DynamicGroup.Input(
|
||||
"loras",
|
||||
template=[
|
||||
Float.Input("strength", default=1.0),
|
||||
String.Input("name"),
|
||||
],
|
||||
min=0,
|
||||
max=10,
|
||||
)
|
||||
assert inp.id == "loras"
|
||||
assert inp.min == 0
|
||||
assert inp.max == 10
|
||||
assert len(inp.template) == 2
|
||||
|
||||
def test_get_all_includes_self_and_template(self):
|
||||
inp = DynamicGroup.Input(
|
||||
"items",
|
||||
template=[Float.Input("value")],
|
||||
)
|
||||
all_inputs = inp.get_all()
|
||||
assert all_inputs[0] is inp
|
||||
assert all_inputs[1].id == "value"
|
||||
|
||||
def test_as_dict_has_template_min_max(self):
|
||||
inp = DynamicGroup.Input(
|
||||
"items",
|
||||
template=[Float.Input("val", default=0.5)],
|
||||
min=1,
|
||||
max=5,
|
||||
)
|
||||
d = inp.as_dict()
|
||||
assert "template" in d
|
||||
assert d["min"] == 1
|
||||
assert d["max"] == 5
|
||||
|
||||
def test_duplicate_field_ids_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input(
|
||||
"bad",
|
||||
template=[Float.Input("x"), Float.Input("x")],
|
||||
)
|
||||
|
||||
def test_empty_template_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input("bad", template=[])
|
||||
|
||||
def test_min_gt_max_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input("bad", template=[Float.Input("x")], min=5, max=3)
|
||||
|
||||
def test_max_exceeds_limit_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input("bad", template=[Float.Input("x")], max=101)
|
||||
|
||||
def test_dynamic_input_in_template_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input(
|
||||
"bad",
|
||||
template=[DynamicGroup.Input("nested", template=[Float.Input("x")])],
|
||||
)
|
||||
|
||||
def test_validate_calls_through(self):
|
||||
inp = DynamicGroup.Input("items", template=[Float.Input("val", min=-1.0, max=1.0)])
|
||||
inp.validate() # should not raise
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 0-row case
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestZeroRows:
|
||||
def test_empty_live_inputs_produces_empty_list(self):
|
||||
"""With min=0 and no live values, the result should be an empty list."""
|
||||
inp = DynamicGroup.Input("loras", template=[Float.Input("strength", default=1.0)], min=0, max=10)
|
||||
assert _run(inp, {}).get("loras") == []
|
||||
|
||||
def test_min_zero_with_values(self):
|
||||
"""min=0 but 2 rows of live data."""
|
||||
inp = DynamicGroup.Input("loras", template=[Float.Input("strength", default=1.0)], min=0, max=10)
|
||||
result = _run(inp, {"loras.0.strength": 0.8, "loras.1.strength": 0.5})
|
||||
assert result["loras"] == [{"strength": 0.8}, {"strength": 0.5}]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# N-row case
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestNRows:
|
||||
def test_two_rows_two_fields(self):
|
||||
"""Two rows with two fields each produce a list[dict]."""
|
||||
inp = DynamicGroup.Input(
|
||||
"loras",
|
||||
template=[String.Input("lora_name"), Float.Input("strength", default=1.0)],
|
||||
min=0, max=50,
|
||||
)
|
||||
result = _run(inp, {
|
||||
"loras.0.lora_name": "model_a.safetensors", "loras.0.strength": 0.9,
|
||||
"loras.1.lora_name": "model_b.safetensors", "loras.1.strength": 0.4,
|
||||
})
|
||||
assert result["loras"] == [
|
||||
{"lora_name": "model_a.safetensors", "strength": 0.9},
|
||||
{"lora_name": "model_b.safetensors", "strength": 0.4},
|
||||
]
|
||||
|
||||
def test_rows_are_sorted_by_index(self):
|
||||
"""Rows must be in ascending index order even if dict iteration is unordered."""
|
||||
inp = DynamicGroup.Input("items", template=[Int.Input("v", default=0)], min=0, max=10)
|
||||
result = _run(inp, {"items.0.v": 10, "items.2.v": 30, "items.1.v": 20})
|
||||
assert [row["v"] for row in result["items"]] == [10, 20, 30]
|
||||
|
||||
def test_min_rows_schema_slots(self):
|
||||
"""With min=2 and no live data, 2 slots must appear in the expanded schema."""
|
||||
inp = DynamicGroup.Input("items", template=[Float.Input("val", default=0.0)], min=2, max=5)
|
||||
out, _, _ = get_finalized_class_inputs(_make_class_inputs(inp), {})
|
||||
all_slots = {**out.get("required", {}), **out.get("optional", {})}
|
||||
assert "items.0.val" in all_slots
|
||||
assert "items.1.val" in all_slots
|
||||
|
||||
def test_min_rows_reconstructs_when_no_values(self):
|
||||
"""min=2 with NO live values must still yield a 2-element list,
|
||||
not collapse to [] (regression: parent-path clobber)."""
|
||||
inp = DynamicGroup.Input("items", template=[Float.Input("val", default=0.0)], min=2, max=5)
|
||||
result = _run(inp, {})
|
||||
assert len(result["items"]) == 2
|
||||
assert all("val" in row for row in result["items"])
|
||||
|
||||
def test_min_rows_reconstructs_with_partial_values(self):
|
||||
"""min=2 with only the first row's value present still yields 2 rows."""
|
||||
inp = DynamicGroup.Input("items", template=[Float.Input("val", default=0.0)], min=2, max=5)
|
||||
result = _run(inp, {"items.0.val": 0.7})
|
||||
assert len(result["items"]) == 2
|
||||
assert result["items"][0]["val"] == 0.7
|
||||
assert result["items"][1]["val"] is None
|
||||
|
||||
def test_list_paths_in_v3_data(self):
|
||||
"""list_paths must contain the group id so build_nested_inputs knows to convert."""
|
||||
inp = DynamicGroup.Input("things", template=[Boolean.Input("flag")], min=0, max=5)
|
||||
_, _, v3_data = get_finalized_class_inputs(_make_class_inputs(inp), {})
|
||||
assert "things" in v3_data.get("list_paths", set())
|
||||
|
||||
def test_no_leftover_flat_keys(self):
|
||||
"""Flat keys must be consumed; only the reconstructed list remains."""
|
||||
inp = DynamicGroup.Input("rows", template=[Float.Input("x", default=0.0)], min=0, max=5)
|
||||
result = _run(inp, {"rows.0.x": 1.0, "rows.1.x": 2.0})
|
||||
assert "rows.0.x" not in result
|
||||
assert "rows.1.x" not in result
|
||||
assert isinstance(result["rows"], list)
|
||||
@@ -2,12 +2,11 @@ import pytest
|
||||
import torch
|
||||
import tempfile
|
||||
import os
|
||||
import sys
|
||||
import av
|
||||
import io
|
||||
from fractions import Fraction
|
||||
from comfy_api.input_impl.video_types import VideoFromFile, VideoFromComponents
|
||||
from comfy_api.util.video_types import VideoComponents, VideoContainer, VideoCodec
|
||||
from comfy_api.util.video_types import VideoComponents
|
||||
from comfy_api.input.basic_types import AudioInput
|
||||
from av.error import InvalidDataError
|
||||
|
||||
@@ -238,526 +237,3 @@ def test_duration_consistency(video_components):
|
||||
manual_duration = float(components.images.shape[0] / components.frame_rate)
|
||||
|
||||
assert duration == pytest.approx(manual_duration)
|
||||
|
||||
|
||||
def create_transcode_source(
|
||||
width=64, height=64, frames=30, fps=30, audio_streams=1, undecodable_audio=0, rotation=False,
|
||||
container_format="mov", audio_codec="pcm_s16le",
|
||||
):
|
||||
"""Create a temp video that save_to must transcode (mpeg4 video, so codec != h264).
|
||||
|
||||
``undecodable_audio`` trailing PCM streams get their fourcc corrupted so no decoder exists
|
||||
(``codec_context is None``), like the APAC track in iPhone spatial-audio recordings.
|
||||
``rotation`` patches a 90-degree display matrix into the video track header.
|
||||
"""
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format=container_format) as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=fps)
|
||||
video_stream.width = width
|
||||
video_stream.height = height
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio = []
|
||||
for _ in range(audio_streams + undecodable_audio):
|
||||
stream = container.add_stream(audio_codec, rate=44100)
|
||||
stream.sample_rate = 44100
|
||||
audio.append(stream)
|
||||
|
||||
for i in range(frames):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((height, width, 3), (i * 7) % 256, dtype=torch.uint8).numpy(),
|
||||
format="rgb24",
|
||||
)
|
||||
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
|
||||
# write audio in 1024-sample frames, like real decoders produce, so the
|
||||
# per-frame skip/cap logic in the transcode path actually runs
|
||||
for stream in audio:
|
||||
for offset in range(0, 44100 * frames // fps, 1024):
|
||||
n = min(1024, 44100 * frames // fps - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(1, n, dtype=torch.int16).numpy(), format="s16", layout="mono"
|
||||
)
|
||||
audio_frame.sample_rate = 44100
|
||||
audio_frame.pts = offset
|
||||
container.mux(stream.encode(audio_frame))
|
||||
for stream in [video_stream, *audio]:
|
||||
container.mux(stream.encode(None))
|
||||
|
||||
data = bytearray(buffer.getvalue())
|
||||
end = len(data)
|
||||
for _ in range(undecodable_audio):
|
||||
end = data.rindex(b"sowt", 0, end)
|
||||
data[end:end + 4] = b"Xpac"
|
||||
if rotation:
|
||||
# the 3x3 display matrix sits 40 bytes into the version-0 tkhd payload; first tkhd
|
||||
# inside moov = video track (search from moov so mdat bytes can't false-match)
|
||||
matrix_offset = data.index(b"tkhd", data.rindex(b"moov")) + 4 + 40
|
||||
values = [0, 1 << 16, 0, -(1 << 16), 0, 0, 0, 0, 1 << 30]
|
||||
data[matrix_offset:matrix_offset + 36] = b"".join(v.to_bytes(4, "big", signed=True) for v in values)
|
||||
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=f".{container_format}", delete=False)
|
||||
tmp.write(bytes(data))
|
||||
tmp.close()
|
||||
return tmp.name
|
||||
|
||||
|
||||
def transcode_and_probe(video):
|
||||
buffer = io.BytesIO()
|
||||
video.save_to(buffer, format=VideoContainer.MP4, codec=VideoCodec.H264)
|
||||
buffer.seek(0)
|
||||
with av.open(buffer) as container:
|
||||
video_stream = container.streams.video[0]
|
||||
audio_stream = container.streams.audio[0] if container.streams.audio else None
|
||||
frames = 0
|
||||
first_pts = None
|
||||
for packet in container.demux(video_stream):
|
||||
for frame in packet.decode():
|
||||
if first_pts is None:
|
||||
first_pts = frame.pts
|
||||
frames += 1
|
||||
return {
|
||||
"codec": video_stream.codec_context.name,
|
||||
"width": video_stream.codec_context.width,
|
||||
"height": video_stream.codec_context.height,
|
||||
"frames": frames,
|
||||
"first_pts": first_pts,
|
||||
"video_seconds": float(video_stream.duration * video_stream.time_base) if video_stream.duration else None,
|
||||
"audio_seconds": float(audio_stream.duration * audio_stream.time_base)
|
||||
if audio_stream and audio_stream.duration else None,
|
||||
"audio_codecs": [s.codec_context.name for s in container.streams.audio],
|
||||
}
|
||||
|
||||
|
||||
def test_save_to_transcode_streams_without_buffering_frames():
|
||||
"""Transcoding must not decode the whole video into memory first (~2 GiB for this source)"""
|
||||
resource = pytest.importorskip("resource") # no getrusage on Windows
|
||||
rss_scale = 1 if sys.platform == "darwin" else 1024 # ru_maxrss: bytes on macOS, KiB elsewhere
|
||||
# ru_maxrss is a lifetime peak: a heavier test running earlier would shrink the measured
|
||||
# delta and quietly defang this canary, so keep this source the biggest thing in the suite
|
||||
file_path = create_transcode_source(width=640, height=480, frames=300)
|
||||
try:
|
||||
rss_before = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale
|
||||
result = transcode_and_probe(VideoFromFile(file_path))
|
||||
rss_delta = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale - rss_before
|
||||
|
||||
assert result["codec"] == "h264"
|
||||
assert result["frames"] == 300
|
||||
assert rss_delta < 500 * 2**20, f"transcode buffered frames in RAM (peak grew {rss_delta / 2**20:.0f} MiB)"
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_honors_trim_window():
|
||||
"""start_time/duration trim applies to both video and audio on the streaming path"""
|
||||
file_path = create_transcode_source(frames=90) # 3s @ 30fps
|
||||
try:
|
||||
result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1))
|
||||
assert result["frames"] == pytest.approx(30, abs=2)
|
||||
assert result["first_pts"] == 0 # trimmed output is rebased to start at zero
|
||||
assert result["video_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_keeps_audio_of_sparse_video():
|
||||
"""Audio that runs ahead of a sparse video track (slideshows, timelapses) must be
|
||||
kept in full — it is only clamped to the video's end, never to the video cursor."""
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=30)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=48000, layout="stereo")
|
||||
for t in (0, 30, 60): # 3 frames spread over 60 seconds
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), t * 4, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = t * 15360
|
||||
frame.time_base = Fraction(1, 15360)
|
||||
container.mux(video_stream.encode(frame))
|
||||
container.mux(video_stream.encode(None))
|
||||
for offset in range(0, 48000 * 60, 1024):
|
||||
n = min(1024, 48000 * 60 - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo"
|
||||
)
|
||||
audio_frame.sample_rate = 48000
|
||||
audio_frame.pts = offset
|
||||
audio_frame.time_base = Fraction(1, 48000)
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
container.mux(audio_stream.encode(None))
|
||||
|
||||
buffer.seek(0)
|
||||
result = transcode_and_probe(VideoFromFile(buffer))
|
||||
assert result["audio_seconds"] == pytest.approx(60.0, abs=1.0)
|
||||
|
||||
|
||||
def test_save_to_transcode_vfr_audio_covers_video_span():
|
||||
"""A trim window in the sparse region of a VFR file keeps audio for the true pts span
|
||||
of the kept frames. Deriving the span as frames/average_rate undercuts it badly: the
|
||||
average is dominated by the dense region (and can be plain wrong on MediaRecorder files)."""
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=30)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=48000, layout="stereo")
|
||||
# 10 frames inside the first second, then one every 1.25 s
|
||||
for i, t in enumerate([x / 10 for x in range(10)] + [1.0, 2.25, 3.5, 4.75]):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), (i * 16) % 256, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = int(t * 15360)
|
||||
frame.time_base = Fraction(1, 15360)
|
||||
container.mux(video_stream.encode(frame))
|
||||
container.mux(video_stream.encode(None))
|
||||
for offset in range(0, 48000 * 6, 1024):
|
||||
n = min(1024, 48000 * 6 - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo"
|
||||
)
|
||||
audio_frame.sample_rate = 48000
|
||||
audio_frame.pts = offset
|
||||
audio_frame.time_base = Fraction(1, 48000)
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
container.mux(audio_stream.encode(None))
|
||||
|
||||
buffer.seek(0)
|
||||
result = transcode_and_probe(VideoFromFile(buffer, start_time=1, duration=5))
|
||||
# kept frames: 1.0/2.25/3.5/4.75 s -> rebased span 3.75 s + one nominal interval
|
||||
assert result["frames"] == 4
|
||||
assert result["audio_seconds"] == pytest.approx(4.0, abs=0.45)
|
||||
|
||||
|
||||
def test_save_to_transcode_trims_audio_in_stream_time_base_units():
|
||||
"""Matroska audio timestamps tick in 1/1000, not 1/sample_rate; trim and audio timing
|
||||
must convert through the frame's time base instead of assuming sample units. AAC audio,
|
||||
because it decodes straight to the encoder's format and hits the resampler passthrough
|
||||
that keeps the source time base on the frames."""
|
||||
file_path = create_transcode_source(frames=90, container_format="matroska", audio_codec="aac")
|
||||
try:
|
||||
result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1))
|
||||
assert result["audio_codecs"] == ["aac"]
|
||||
assert result["video_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_learns_unprobed_audio_params():
|
||||
"""mpegts is only probed a few seconds deep at open, so an audio stream whose first
|
||||
packet comes later (live captures where audio kicks in late) still has sample_rate 0
|
||||
when the transcode starts; the parameters must be learned from the stream itself."""
|
||||
sample_rate, fps, video_seconds, audio_start = 48000, 30, 13, 12
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mpegts") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=fps)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
|
||||
for i in range(video_seconds * fps):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
)
|
||||
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
|
||||
for offset in range(0, (video_seconds - audio_start) * sample_rate, 1024):
|
||||
n = min(1024, (video_seconds - audio_start) * sample_rate - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono"
|
||||
)
|
||||
audio_frame.sample_rate = sample_rate
|
||||
audio_frame.pts = audio_start * sample_rate + offset
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
for stream in (video_stream, audio_stream):
|
||||
container.mux(stream.encode(None))
|
||||
|
||||
buffer.seek(0)
|
||||
with av.open(buffer) as container:
|
||||
# the scenario requires unprobed parameters; if a future FFmpeg probes deeper,
|
||||
# push audio_start/video_seconds further out to restore it
|
||||
assert container.streams.audio[0].codec_context.sample_rate == 0
|
||||
result = transcode_and_probe(VideoFromFile(buffer))
|
||||
assert result["frames"] == video_seconds * fps
|
||||
assert result["audio_codecs"] == ["aac"]
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
|
||||
buffer.seek(0)
|
||||
trimmed_before_audio = transcode_and_probe(VideoFromFile(buffer, duration=1))
|
||||
assert trimmed_before_audio["frames"] == fps
|
||||
assert trimmed_before_audio["audio_codecs"] == []
|
||||
assert trimmed_before_audio["audio_seconds"] is None
|
||||
|
||||
buffer.seek(0)
|
||||
trimmed_crossing_audio = transcode_and_probe(VideoFromFile(buffer, start_time=11.5, duration=1))
|
||||
assert trimmed_crossing_audio["frames"] == fps
|
||||
assert trimmed_crossing_audio["audio_codecs"] == ["aac"]
|
||||
assert trimmed_crossing_audio["video_seconds"] == pytest.approx(1.0, abs=0.05)
|
||||
assert trimmed_crossing_audio["audio_seconds"] == pytest.approx(0.5, abs=0.1)
|
||||
|
||||
|
||||
def test_save_to_transcode_trimmed_fragmented_mp4_keeps_audio():
|
||||
"""Fragmented mp4 (MediaRecorder, DASH/HLS-derived files) delivers audio well behind
|
||||
video, so when the trim window's last video frame arrives the audio demuxed so far
|
||||
does not cover the window yet; the transcode must keep demuxing audio until it does
|
||||
instead of finalizing on the first audio frame it sees afterwards."""
|
||||
sample_rate, fps, seconds = 48000, 30, 6
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4", options={"movflags": "frag_keyframe+empty_moov"}) as container:
|
||||
video_stream = container.add_stream("h264", rate=fps)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
|
||||
next_audio_pts = 0
|
||||
for i in range(seconds * fps):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
)
|
||||
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
|
||||
while next_audio_pts / sample_rate <= i / fps: # feed audio alongside, like a live pipeline
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(1, 1024, dtype=torch.float32).numpy(), format="fltp", layout="mono"
|
||||
)
|
||||
audio_frame.sample_rate = sample_rate
|
||||
audio_frame.pts = next_audio_pts
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
next_audio_pts += 1024
|
||||
for stream in (video_stream, audio_stream):
|
||||
container.mux(stream.encode(None))
|
||||
|
||||
result = transcode_and_probe(VideoFromFile(buffer, start_time=0.5, duration=1.0))
|
||||
assert result["video_seconds"] == pytest.approx(1.0, abs=0.05)
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.05)
|
||||
|
||||
|
||||
def test_save_to_transcode_sparse_video_keeps_true_duration():
|
||||
"""average_rate is not a frame duration: a 3-frame video spanning 60 s averages
|
||||
0.05 fps, and padding the last frame with 1/average_rate used to extend the
|
||||
output — and the audio kept with it — about 20 s past the source span."""
|
||||
sample_rate = 48000
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=30)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
|
||||
for i, second in enumerate((0, 30, 60)):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), i * 80, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = second * 30
|
||||
frame.time_base = Fraction(1, 30)
|
||||
container.mux(video_stream.encode(frame))
|
||||
for offset in range(0, 90 * sample_rate, 1024):
|
||||
n = min(1024, 90 * sample_rate - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono"
|
||||
)
|
||||
audio_frame.sample_rate = sample_rate
|
||||
audio_frame.pts = offset
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
for stream in (video_stream, audio_stream):
|
||||
container.mux(stream.encode(None))
|
||||
|
||||
result = transcode_and_probe(VideoFromFile(buffer))
|
||||
assert result["frames"] == 3
|
||||
# the last frame keeps its true stts duration (1/30 s), not 1/average_rate (~20 s)
|
||||
assert result["video_seconds"] == pytest.approx(60.03, abs=0.05)
|
||||
assert result["audio_seconds"] == pytest.approx(60.03, abs=0.1)
|
||||
|
||||
trimmed = transcode_and_probe(VideoFromFile(buffer, duration=45))
|
||||
assert trimmed["frames"] == 2
|
||||
# a kept frame whose source duration crosses the window end is clamped to it
|
||||
assert trimmed["video_seconds"] == pytest.approx(45.0, abs=0.05)
|
||||
assert trimmed["audio_seconds"] == pytest.approx(45.0, abs=0.1)
|
||||
|
||||
|
||||
def test_save_to_transcode_clamps_final_pts_to_declared_stream_duration():
|
||||
"""Some iPhone MOVs report a video stream duration that ends before the final
|
||||
decoded frame's nominal duration. A transcode must not turn that trailing
|
||||
timestamp quirk into an extra frame interval compared to the source/remux path."""
|
||||
fps = 30
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=fps)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
for i, pts in enumerate([*range(31), 32]):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = pts
|
||||
frame.time_base = Fraction(1, fps)
|
||||
container.mux(video_stream.encode(frame))
|
||||
container.mux(video_stream.encode(None))
|
||||
|
||||
class _StreamProxy:
|
||||
def __init__(self, stream, duration):
|
||||
self._stream = stream
|
||||
self.duration = duration
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._stream, name)
|
||||
|
||||
class _StreamsProxy:
|
||||
def __init__(self, video_stream):
|
||||
self.video = [video_stream]
|
||||
self.audio = []
|
||||
|
||||
class _PacketProxy:
|
||||
def __init__(self, packet, stream):
|
||||
self._packet = packet
|
||||
self.stream = stream
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._packet, name)
|
||||
|
||||
class _ContainerProxy:
|
||||
def __init__(self, container, stream):
|
||||
self._container = container
|
||||
self._stream = stream
|
||||
self.streams = _StreamsProxy(stream)
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._container, name)
|
||||
|
||||
def demux(self, *streams):
|
||||
for packet in self._container.demux(self._stream._stream):
|
||||
yield _PacketProxy(packet, self._stream)
|
||||
|
||||
buffer.seek(0)
|
||||
output = io.BytesIO()
|
||||
with av.open(buffer) as container:
|
||||
real_stream = container.streams.video[0]
|
||||
declared_duration = 32 * int(round((1 / fps) / real_stream.time_base))
|
||||
stream = _StreamProxy(real_stream, declared_duration)
|
||||
VideoFromFile(buffer)._save_transcoded(
|
||||
_ContainerProxy(container, stream), output, VideoContainer.MP4, VideoCodec.H264, None, 8
|
||||
)
|
||||
|
||||
output.seek(0)
|
||||
with av.open(output) as container:
|
||||
video_stream = container.streams.video[0]
|
||||
frames = [f for p in container.demux(video_stream) for f in p.decode()]
|
||||
assert len(frames) == 32
|
||||
assert float(video_stream.duration * video_stream.time_base) == pytest.approx(32 / fps, abs=0.01)
|
||||
assert float(frames[-1].pts * frames[-1].time_base) == pytest.approx(31 / fps, abs=0.01)
|
||||
|
||||
|
||||
def test_save_to_transcode_irregular_vfr_keeps_span():
|
||||
"""B-frames reorder packets, and mp4 sample durations follow decode order: the dts
|
||||
timeline ends before the pts timeline, so an irregular-VFR source's tail holds fell
|
||||
out of the container (this 20.23 s span used to come out as 15.27 s, and the 10 s
|
||||
trim as 6.03 s). The transcode encodes without B-frames so every sample keeps its
|
||||
true display duration."""
|
||||
durations = [1, 1, 60, 1, 1, 120, 1, 180, 1, 1, 150, 90] # 1/30 s ticks, span 20.2333 s
|
||||
generator = torch.Generator().manual_seed(7)
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=30)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
pts = 0
|
||||
for duration in durations:
|
||||
# textured frames, so an encoder with default settings has B-frames to gain from
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.randint(0, 255, (64, 64, 3), generator=generator, dtype=torch.uint8).numpy(),
|
||||
format="rgb24",
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = pts
|
||||
frame.time_base = Fraction(1, 30)
|
||||
pts += duration
|
||||
for packet in video_stream.encode(frame):
|
||||
packet.duration = duration # exact stts in the source
|
||||
container.mux(packet)
|
||||
container.mux(video_stream.encode(None))
|
||||
|
||||
result = transcode_and_probe(VideoFromFile(buffer))
|
||||
assert result["frames"] == len(durations)
|
||||
assert result["video_seconds"] == pytest.approx(sum(durations) / 30, abs=0.05)
|
||||
|
||||
trimmed = transcode_and_probe(VideoFromFile(buffer, duration=10))
|
||||
assert trimmed["frames"] == 8 # frames at 12.167 s+ fall outside the window
|
||||
assert trimmed["video_seconds"] == pytest.approx(10.0, abs=0.05)
|
||||
|
||||
|
||||
def test_save_to_transcode_trim_survives_missing_leading_pts():
|
||||
"""A trim should survive pts-less kept frames followed by a real-pts frame past the window."""
|
||||
nulled_frames = 0
|
||||
|
||||
class _PacketProxy:
|
||||
def __init__(self, packet):
|
||||
self._packet = packet
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._packet, name)
|
||||
|
||||
@property
|
||||
def stream(self):
|
||||
return self._packet.stream
|
||||
|
||||
def decode(self):
|
||||
nonlocal nulled_frames
|
||||
frames = self._packet.decode()
|
||||
for frame in frames:
|
||||
if nulled_frames < 2:
|
||||
frame.pts = None
|
||||
nulled_frames += 1
|
||||
return frames
|
||||
|
||||
class _ContainerProxy:
|
||||
def __init__(self, real):
|
||||
self._real = real
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._real, name)
|
||||
|
||||
def demux(self, *streams):
|
||||
for packet in self._real.demux(*streams):
|
||||
yield _PacketProxy(packet)
|
||||
|
||||
file_path = create_transcode_source(frames=10, audio_streams=0)
|
||||
try:
|
||||
buffer = io.BytesIO()
|
||||
with av.open(file_path) as container:
|
||||
# 0.05 s window: both pts-less frames are kept (synthesized pts 0 and 512),
|
||||
# and the first real-pts frame (1024 ticks) already lies past end_pts (768)
|
||||
VideoFromFile(file_path, duration=0.05)._save_transcoded(
|
||||
_ContainerProxy(container), buffer, VideoContainer.MP4, VideoCodec.H264, None, 8
|
||||
)
|
||||
assert nulled_frames == 2
|
||||
buffer.seek(0)
|
||||
with av.open(buffer) as container:
|
||||
video_stream = container.streams.video[0]
|
||||
frames = [f for p in container.demux(video_stream) for f in p.decode()]
|
||||
assert len(frames) == 2
|
||||
assert float(video_stream.duration * video_stream.time_base) == pytest.approx(2 / 30, abs=0.01)
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_bakes_rotation():
|
||||
"""A 90-degree display-matrix rotation swaps the output dimensions (portrait video)"""
|
||||
file_path = create_transcode_source(width=64, height=32, rotation=True)
|
||||
try:
|
||||
result = transcode_and_probe(VideoFromFile(file_path))
|
||||
assert (result["width"], result["height"]) == (32, 64)
|
||||
assert result["frames"] == 30
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_skips_undecodable_audio():
|
||||
"""Streaming transcode keeps the decodable audio track and drops undecodable ones;
|
||||
with no decodable audio at all the output is video-only instead of crashing."""
|
||||
mixed = all_bad = None
|
||||
try:
|
||||
mixed = create_transcode_source(audio_streams=1, undecodable_audio=1)
|
||||
all_bad = create_transcode_source(audio_streams=0, undecodable_audio=2)
|
||||
result = transcode_and_probe(VideoFromFile(mixed))
|
||||
assert result["audio_codecs"] == ["aac"]
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
assert transcode_and_probe(VideoFromFile(all_bad))["audio_codecs"] == []
|
||||
finally:
|
||||
for path in (mixed, all_bad):
|
||||
if path:
|
||||
os.unlink(path)
|
||||
|
||||
@@ -1,186 +0,0 @@
|
||||
"""SeedVR2 conditioning node regression tests."""
|
||||
|
||||
import importlib
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
from comfy.ldm.seedvr.constants import SEEDVR2_LATENT_CHANNELS
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
|
||||
_SENTINEL = object()
|
||||
_TARGETS = (
|
||||
("comfy.model_management", "comfy"),
|
||||
("comfy_extras.nodes_seedvr", "comfy_extras"),
|
||||
)
|
||||
|
||||
|
||||
def _import_nodes_seedvr_isolated():
|
||||
"""Import comfy_extras.nodes_seedvr with comfy.model_management mocked."""
|
||||
priors = []
|
||||
for mod_name, parent_name in _TARGETS:
|
||||
prior_mod = sys.modules.get(mod_name, _SENTINEL)
|
||||
parent = sys.modules.get(parent_name)
|
||||
attr = mod_name.split(".")[-1]
|
||||
prior_attr = (
|
||||
getattr(parent, attr, _SENTINEL) if parent is not None else _SENTINEL
|
||||
)
|
||||
priors.append((mod_name, parent_name, attr, prior_mod, prior_attr))
|
||||
|
||||
mock_mm = MagicMock()
|
||||
for fn in (
|
||||
"xformers_enabled", "xformers_enabled_vae",
|
||||
"pytorch_attention_enabled", "pytorch_attention_enabled_vae",
|
||||
"sage_attention_enabled", "flash_attention_enabled",
|
||||
"is_intel_xpu",
|
||||
):
|
||||
getattr(mock_mm, fn).return_value = False
|
||||
tv = torch.version.__version__.split(".")
|
||||
mock_mm.torch_version_numeric = (int(tv[0]), int(tv[1]))
|
||||
mock_mm.WINDOWS = False
|
||||
sys.modules["comfy.model_management"] = mock_mm
|
||||
if sys.modules.get("comfy") is None:
|
||||
importlib.import_module("comfy")
|
||||
comfy_pkg = sys.modules.get("comfy")
|
||||
if comfy_pkg is not None:
|
||||
setattr(comfy_pkg, "model_management", mock_mm)
|
||||
nodes_seedvr = sys.modules.get("comfy_extras.nodes_seedvr") or (
|
||||
importlib.import_module("comfy_extras.nodes_seedvr")
|
||||
)
|
||||
|
||||
def _restore():
|
||||
for mod_name, parent_name, attr, prior_mod, prior_attr in priors:
|
||||
if prior_mod is _SENTINEL:
|
||||
sys.modules.pop(mod_name, None)
|
||||
else:
|
||||
sys.modules[mod_name] = prior_mod
|
||||
parent = sys.modules.get(parent_name)
|
||||
if parent is None:
|
||||
continue
|
||||
if prior_attr is _SENTINEL:
|
||||
if hasattr(parent, attr):
|
||||
delattr(parent, attr)
|
||||
else:
|
||||
setattr(parent, attr, prior_attr)
|
||||
|
||||
return nodes_seedvr, _restore
|
||||
|
||||
|
||||
class _Rope(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.freqs = nn.Parameter(torch.zeros(4))
|
||||
|
||||
|
||||
class _Block(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.rope = _Rope()
|
||||
|
||||
|
||||
class _DiffusionModel(nn.Module):
|
||||
def __init__(self, n_blocks=3, conditioning_dtype=torch.float32):
|
||||
super().__init__()
|
||||
self.blocks = nn.ModuleList([_Block() for _ in range(n_blocks)])
|
||||
self.register_buffer("positive_conditioning", torch.ones((2, 4), dtype=conditioning_dtype))
|
||||
self.register_buffer("negative_conditioning", torch.zeros((3, 4), dtype=conditioning_dtype))
|
||||
|
||||
|
||||
class _ModelInner:
|
||||
def __init__(self, diffusion_model):
|
||||
self.diffusion_model = diffusion_model
|
||||
|
||||
|
||||
class _ModelPatcher:
|
||||
def __init__(self, diffusion_model):
|
||||
self.model = _ModelInner(diffusion_model)
|
||||
|
||||
|
||||
def test_seedvr2_conditioning_schema_exposes_conditioning_outputs():
|
||||
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
|
||||
try:
|
||||
schema = nodes_seedvr.SeedVR2Conditioning.define_schema()
|
||||
assert [input_item.id for input_item in schema.inputs] == [
|
||||
"model",
|
||||
"vae_conditioning",
|
||||
]
|
||||
assert schema.inputs[1].display_name == "latent"
|
||||
assert [output.display_name for output in schema.outputs] == [
|
||||
"positive",
|
||||
"negative",
|
||||
]
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_seedvr2_conditioning_rejects_wrong_latent_channels():
|
||||
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
|
||||
try:
|
||||
patcher = _ModelPatcher(_DiffusionModel())
|
||||
vae_conditioning = {"samples": torch.zeros(1, 8, 2, 2, 2)}
|
||||
|
||||
with pytest.raises(ValueError, match=f"{SEEDVR2_LATENT_CHANNELS} channels"):
|
||||
nodes_seedvr.SeedVR2Conditioning.execute(patcher, vae_conditioning)
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_seedvr2_conditioning_returns_conditioning_deterministically():
|
||||
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
|
||||
try:
|
||||
diffusion_model = _DiffusionModel()
|
||||
patcher = _ModelPatcher(diffusion_model)
|
||||
samples = torch.arange(
|
||||
1,
|
||||
1 + SEEDVR2_LATENT_CHANNELS * 3 * 2 * 2,
|
||||
dtype=torch.float32,
|
||||
).reshape(1, SEEDVR2_LATENT_CHANNELS, 3, 2, 2)
|
||||
vae_conditioning = {"samples": samples}
|
||||
|
||||
first_positive, first_negative = (
|
||||
nodes_seedvr.SeedVR2Conditioning.execute(
|
||||
patcher,
|
||||
vae_conditioning,
|
||||
)
|
||||
)
|
||||
second_positive, second_negative = (
|
||||
nodes_seedvr.SeedVR2Conditioning.execute(
|
||||
patcher,
|
||||
vae_conditioning,
|
||||
)
|
||||
)
|
||||
|
||||
channel_last = samples.movedim(1, -1).contiguous()
|
||||
expected_condition = torch.cat(
|
||||
[
|
||||
channel_last,
|
||||
torch.ones((*channel_last.shape[:-1], 1)),
|
||||
],
|
||||
dim=-1,
|
||||
).movedim(-1, 1)
|
||||
|
||||
assert torch.equal(
|
||||
first_positive[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
assert torch.equal(
|
||||
second_positive[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
assert torch.equal(
|
||||
first_negative[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
assert torch.equal(
|
||||
second_negative[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
finally:
|
||||
restore()
|
||||
@@ -1,55 +0,0 @@
|
||||
import importlib
|
||||
import inspect
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
|
||||
def test_seedvr_node_signature_matches_schema():
|
||||
mock_mm = MagicMock()
|
||||
mock_mm.xformers_enabled.return_value = False
|
||||
mock_mm.xformers_enabled_vae.return_value = False
|
||||
mock_mm.sage_attention_enabled.return_value = False
|
||||
mock_mm.flash_attention_enabled.return_value = False
|
||||
|
||||
sentinel = object()
|
||||
prior_cpu = cli_args.cpu
|
||||
cli_args.cpu = True
|
||||
prior_module = sys.modules.get("comfy_extras.nodes_seedvr", sentinel)
|
||||
comfy_pkg = sys.modules.get("comfy")
|
||||
prior_mm_attr = getattr(comfy_pkg, "model_management", sentinel) if comfy_pkg else sentinel
|
||||
|
||||
with patch.dict(sys.modules, {"comfy.model_management": mock_mm}):
|
||||
if comfy_pkg is not None:
|
||||
setattr(comfy_pkg, "model_management", mock_mm)
|
||||
sys.modules.pop("comfy_extras.nodes_seedvr", None)
|
||||
try:
|
||||
nodes_seedvr = importlib.import_module("comfy_extras.nodes_seedvr")
|
||||
for node_cls in (nodes_seedvr.SeedVR2Preprocess, nodes_seedvr.SeedVR2PostProcessing, nodes_seedvr.SeedVR2Conditioning):
|
||||
schema_ids = [i.id for i in node_cls.define_schema().inputs]
|
||||
exec_params = [
|
||||
p for p in inspect.signature(node_cls.execute).parameters.keys()
|
||||
if p != "cls"
|
||||
]
|
||||
assert schema_ids == exec_params, (
|
||||
f"{node_cls.__name__} schema/execute drift: "
|
||||
f"schema_ids={schema_ids}, exec_params={exec_params}"
|
||||
)
|
||||
finally:
|
||||
cli_args.cpu = prior_cpu
|
||||
if prior_module is sentinel:
|
||||
sys.modules.pop("comfy_extras.nodes_seedvr", None)
|
||||
else:
|
||||
sys.modules["comfy_extras.nodes_seedvr"] = prior_module
|
||||
if comfy_pkg is not None:
|
||||
if prior_mm_attr is sentinel:
|
||||
if hasattr(comfy_pkg, "model_management"):
|
||||
delattr(comfy_pkg, "model_management")
|
||||
else:
|
||||
setattr(comfy_pkg, "model_management", prior_mm_attr)
|
||||
@@ -1,51 +0,0 @@
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
from comfy_extras import nodes_seedvr # noqa: E402
|
||||
|
||||
|
||||
def _schema_ids(items):
|
||||
return [item.id for item in items]
|
||||
|
||||
|
||||
def test_seedvr2_post_processing_schema():
|
||||
schema = nodes_seedvr.SeedVR2PostProcessing.define_schema()
|
||||
|
||||
assert _schema_ids(schema.inputs) == ["images", "original_resized_images", "color_correction_method"]
|
||||
assert schema.inputs[2].options == ["lab", "wavelet", "adain", "none"]
|
||||
assert schema.inputs[2].default == "lab"
|
||||
assert schema.outputs[0].get_io_type() == "IMAGE"
|
||||
|
||||
|
||||
def test_seedvr2_post_processing_oom_error_uses_color_correction_method(monkeypatch):
|
||||
decoded = torch.full((1, 3, 4, 4), 0.25)
|
||||
reference = torch.full((1, 3, 4, 4), 0.75)
|
||||
|
||||
def _lab(content, style):
|
||||
raise torch.cuda.OutOfMemoryError("CUDA out of memory")
|
||||
|
||||
monkeypatch.setattr(nodes_seedvr.comfy.model_management, "vae_device", lambda: torch.device("cpu"))
|
||||
monkeypatch.setattr(nodes_seedvr.comfy.model_management, "get_free_memory", lambda device: 1_000_000)
|
||||
|
||||
with patch.object(nodes_seedvr, "lab_color_transfer", _lab):
|
||||
with pytest.raises(RuntimeError) as excinfo:
|
||||
nodes_seedvr.SeedVR2PostProcessing._color_transfer_chunked(
|
||||
decoded, reference, torch.device("cpu"), "lab",
|
||||
)
|
||||
assert "color_correction_method=lab" in str(excinfo.value)
|
||||
assert " method=lab" not in str(excinfo.value)
|
||||
|
||||
|
||||
def test_seedvr2_post_processing_unknown_color_correction_method_raises():
|
||||
decoded = torch.zeros(1, 2, 4, 4, 3)
|
||||
original = torch.zeros(1, 2, 4, 4, 3)
|
||||
with pytest.raises(ValueError) as excinfo:
|
||||
nodes_seedvr.SeedVR2PostProcessing.execute(decoded, original, "bogus")
|
||||
assert "color_correction_method" in str(excinfo.value)
|
||||
@@ -1,77 +0,0 @@
|
||||
"""SeedVR2 temporal chunk/merge node regression tests."""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
from comfy.ldm.seedvr.constants import (
|
||||
BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE,
|
||||
SEEDVR2_CHUNK_GIB_PER_MPX_FRAME,
|
||||
SEEDVR2_CHUNK_RESERVED_GIB,
|
||||
SEEDVR2_CHUNK_SIGMA_GIB,
|
||||
SEEDVR2_CHUNK_SIGMA_K,
|
||||
SEEDVR2_LATENT_CHANNELS,
|
||||
)
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.model_management # noqa: E402
|
||||
from comfy_extras.nodes_seedvr import SeedVR2TemporalChunk, SeedVR2TemporalMerge, _seedvr2_chunk_crossfade_weights # noqa: E402
|
||||
|
||||
def _latent(t_latent, h=8, w=8, b=1):
|
||||
g = torch.Generator().manual_seed(7)
|
||||
return {"samples": torch.randn(b, SEEDVR2_LATENT_CHANNELS, t_latent, h, w, generator=g)}
|
||||
|
||||
def _split(latent, frames_per_chunk, temporal_overlap, chunking_mode="manual"):
|
||||
combo = {"chunking_mode": chunking_mode}
|
||||
if chunking_mode != "auto":
|
||||
combo["frames_per_chunk"] = frames_per_chunk
|
||||
return SeedVR2TemporalChunk.execute(latent, temporal_overlap, combo).args
|
||||
|
||||
def _merge(chunks, temporal_overlap):
|
||||
return SeedVR2TemporalMerge.execute(chunks, [temporal_overlap]).args[0]
|
||||
|
||||
def test_chunk_temporal_windows_and_validation():
|
||||
with pytest.raises(ValueError, match="4n\\+1"):
|
||||
_split(_latent(9), 20, 0)
|
||||
with pytest.raises(ValueError, match="5-D"):
|
||||
_split({"samples": torch.zeros(1, SEEDVR2_LATENT_CHANNELS * 9, 8, 8)}, 21, 0)
|
||||
with pytest.raises(ValueError, match="chunking_mode"):
|
||||
_split(_latent(13), 21, 0, "adaptive")
|
||||
latent = _latent(13)
|
||||
chunks, overlap = _split(latent, 21, 2) # chunk_latent=6, step=4 -> [0:6], [4:10], [8:13]
|
||||
assert overlap == 2 and [c["samples"].shape[2] for c in chunks] == [6, 6, 5]
|
||||
assert all(torch.equal(c["samples"], latent["samples"][:, :, s:e]) for c, (s, e) in zip(chunks, [(0, 6), (4, 10), (8, 13)]))
|
||||
assert len(_split(_latent(13), 21, 999)[0]) == 8 # overlap clamps to chunk_latent-1 -> step=1
|
||||
assert (r := _split(_latent(5), 21, 3)) and len(r[0]) == 1 and r[1] == 0 # t_pixel <= 21: passthrough
|
||||
|
||||
def test_chunk_auto_mode_applies_vram_law(monkeypatch):
|
||||
mpx_per_frame = (32 * 32) * (BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE ** 2) / 1e6
|
||||
free_gb = (
|
||||
SEEDVR2_CHUNK_RESERVED_GIB
|
||||
+ SEEDVR2_CHUNK_SIGMA_K * SEEDVR2_CHUNK_SIGMA_GIB
|
||||
+ 5.1 * SEEDVR2_CHUNK_GIB_PER_MPX_FRAME * mpx_per_frame
|
||||
)
|
||||
monkeypatch.setattr(comfy.model_management, "get_free_memory", lambda dev=None: free_gb * (1024 ** 3))
|
||||
assert [c["samples"].shape[2] for c in _split(_latent(13, h=32, w=32), 1, 0, "auto")[0]] == [5, 5, 3]
|
||||
assert _split(_latent(13, h=32, w=32, b=2), 1, 0, "auto")[0][0]["samples"].shape[2] == 2 # batch halves the chunk
|
||||
|
||||
def test_merge_crossfade_and_reassembly():
|
||||
latent = _latent(13)
|
||||
latent["noise_mask"] = torch.rand(1, 1, 13, 8, 8)
|
||||
latent["batch_index"] = [0]
|
||||
merged = _merge(_split(latent, 21, 0)[0], 0)
|
||||
assert torch.equal(merged["samples"], latent["samples"])
|
||||
assert "noise_mask" not in merged and merged["batch_index"] == [0]
|
||||
assert torch.allclose(_merge(_split(latent, 21, 3)[0], 3)["samples"], latent["samples"], atol=1e-6)
|
||||
w = _seedvr2_chunk_crossfade_weights(3, merged["samples"].device, merged["samples"].dtype)
|
||||
assert w[0] == 1.0 and w[-1] == 0.0 and torch.all(w[:-1] >= w[1:])
|
||||
ones, zeros = {"samples": torch.ones(1, SEEDVR2_LATENT_CHANNELS, 6, 8, 8)}, {"samples": torch.zeros(1, SEEDVR2_LATENT_CHANNELS, 6, 8, 8)}
|
||||
fused = _merge([ones, zeros], 3)["samples"] # overlap equals w: prev fades out, next fades in
|
||||
assert torch.equal(fused[:, :, 3:6], w.view(1, 1, 3, 1, 1).expand(1, SEEDVR2_LATENT_CHANNELS, 3, 8, 8))
|
||||
assert torch.equal(fused[:, :, :3], ones["samples"][:, :, :3]) and torch.equal(fused[:, :, 6:], zeros["samples"][:, :, :3])
|
||||
short = _split(latent, 21, 2)[0]
|
||||
short[0]["samples"] = short[0]["samples"][:, :, :4]
|
||||
with pytest.raises(ValueError, match="only the final chunk may be shorter"):
|
||||
_merge(short, 2)
|
||||
@@ -15,7 +15,7 @@ if not has_gpu():
|
||||
args.cpu = True
|
||||
|
||||
from comfy import ops
|
||||
from comfy.quant_ops import QUANT_ALGOS, QuantizedTensor
|
||||
from comfy.quant_ops import QuantizedTensor
|
||||
import comfy.utils
|
||||
|
||||
|
||||
@@ -283,59 +283,7 @@ class TestMixedPrecisionOps(unittest.TestCase):
|
||||
saved = model.state_dict()
|
||||
saved_conf = json.loads(saved["layer.comfy_quant"].numpy().tobytes())
|
||||
self.assertTrue(saved_conf["convrot"])
|
||||
|
||||
def test_convrot_w4a4_loads_into_params(self):
|
||||
"""ConvRot W4A4 checkpoints must load as the dedicated kitchen layout."""
|
||||
if "convrot_w4a4" not in QUANT_ALGOS:
|
||||
self.skipTest("comfy_kitchen does not provide ConvRot W4A4")
|
||||
|
||||
torch.manual_seed(456)
|
||||
layer_quant_config = {
|
||||
"layer": {
|
||||
"format": "convrot_w4a4",
|
||||
"convrot_groupsize": 256,
|
||||
"linear_dtype": "int8",
|
||||
}
|
||||
}
|
||||
weight = torch.randn(16, 256, dtype=torch.bfloat16)
|
||||
bias = torch.randn(16, dtype=torch.bfloat16)
|
||||
q_weight = QuantizedTensor.from_float(
|
||||
weight,
|
||||
"TensorCoreConvRotW4A4Layout",
|
||||
convrot_groupsize=256,
|
||||
quant_group_size=64,
|
||||
)
|
||||
state_dict = {
|
||||
"layer.weight": q_weight._qdata,
|
||||
"layer.bias": bias,
|
||||
"layer.weight_scale": q_weight._params.scale,
|
||||
}
|
||||
|
||||
state_dict, _ = comfy.utils.convert_old_quants(
|
||||
state_dict,
|
||||
metadata={"_quantization_metadata": json.dumps({"layers": layer_quant_config})},
|
||||
)
|
||||
model = torch.nn.Module()
|
||||
model.layer = ops.mixed_precision_ops({}).Linear(256, 16, device="cpu", dtype=torch.bfloat16)
|
||||
model.load_state_dict(state_dict, strict=False)
|
||||
|
||||
self.assertIsInstance(model.layer.weight, QuantizedTensor)
|
||||
self.assertEqual(model.layer.weight._layout_cls, "TensorCoreConvRotW4A4Layout")
|
||||
self.assertEqual(model.layer.weight._params.convrot_groupsize, 256)
|
||||
self.assertEqual(model.layer.weight._params.quant_group_size, 64)
|
||||
self.assertEqual(model.layer.weight._params.linear_dtype, "int8")
|
||||
|
||||
input_tensor = torch.randn(4, 256, dtype=torch.bfloat16)
|
||||
loaded_out = model.layer(input_tensor)
|
||||
ref_out = torch.nn.functional.linear(input_tensor, q_weight, bias)
|
||||
self.assertTrue(torch.equal(loaded_out, ref_out))
|
||||
|
||||
saved = model.state_dict()
|
||||
saved_conf = json.loads(saved["layer.comfy_quant"].numpy().tobytes())
|
||||
self.assertEqual(saved_conf["format"], "convrot_w4a4")
|
||||
self.assertEqual(saved_conf["convrot_groupsize"], 256)
|
||||
self.assertEqual(saved_conf["linear_dtype"], "int8")
|
||||
self.assertNotIn("quant_group_size", saved_conf)
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -2,7 +2,7 @@ from collections import defaultdict
|
||||
|
||||
import torch
|
||||
|
||||
from comfy.model_detection import detect_unet_config, model_config_from_unet, model_config_from_unet_config
|
||||
from comfy.model_detection import detect_unet_config, model_config_from_unet_config
|
||||
import comfy.supported_models
|
||||
|
||||
|
||||
@@ -73,60 +73,6 @@ def _make_flux_schnell_comfyui_sd():
|
||||
return sd
|
||||
|
||||
|
||||
def _make_seedvr2_7b_separate_mm_sd():
|
||||
return {
|
||||
"blocks.35.mlp.vid.proj_out.weight": torch.empty(3072, 1),
|
||||
"positive_conditioning": torch.empty(58, 5120),
|
||||
"negative_conditioning": torch.empty(64, 5120),
|
||||
}
|
||||
|
||||
|
||||
def _make_seedvr2_7b_shared_mm_sd():
|
||||
return {
|
||||
"blocks.35.mlp.all.proj_in_gate.weight": torch.empty(1, 1),
|
||||
"positive_conditioning": torch.empty(58, 5120),
|
||||
"negative_conditioning": torch.empty(64, 5120),
|
||||
}
|
||||
|
||||
|
||||
def _make_seedvr2_3b_shared_mm_sd():
|
||||
return {
|
||||
"blocks.31.mlp.all.proj_in_gate.weight": torch.empty(1, 1),
|
||||
"positive_conditioning": torch.empty(58, 5120),
|
||||
"negative_conditioning": torch.empty(64, 5120),
|
||||
}
|
||||
|
||||
|
||||
def _make_pid_v1_5_sd(latent_proj_channels=16):
|
||||
sd = {
|
||||
"pixel_embedder.proj.weight": torch.empty(16, 3, device="meta"),
|
||||
"lq_proj.latent_proj.0.weight": torch.empty(1024, latent_proj_channels, 3, 3, device="meta"),
|
||||
"lq_proj.pit_head.weight": torch.empty(1536, 1024, device="meta"),
|
||||
"lq_proj.gate_modules.0.content_proj.weight": torch.empty(1, 3072, device="meta"),
|
||||
"pixel_blocks.0.attn.q_norm.weight": torch.empty(72, device="meta"),
|
||||
"pixel_blocks.0.adaLN_modulation.0.weight": torch.empty(24576, 1536, device="meta"),
|
||||
"pixel_blocks.0.adaLN_modulation.0.bias": torch.empty(24576, device="meta"),
|
||||
}
|
||||
for i in range(7):
|
||||
sd[f"lq_proj.gate_modules.{i}.log_alpha"] = torch.empty((), device="meta")
|
||||
return sd
|
||||
|
||||
|
||||
def _make_joyimage_edit_plus_sd():
|
||||
sd = {
|
||||
"img_in.weight": torch.empty(4096, 16, 1, 2, 2, device="meta"),
|
||||
"condition_embedder.time_embedder.linear_1.weight": torch.empty(1, device="meta"),
|
||||
"double_blocks.0.attn.img_attn_q_norm.weight": torch.empty(128, device="meta"),
|
||||
}
|
||||
for i in range(40):
|
||||
sd[f"double_blocks.{i}.attn.img_attn_qkv.weight"] = torch.empty(1, device="meta")
|
||||
return sd
|
||||
|
||||
|
||||
def _add_model_diffusion_prefix(sd):
|
||||
return {f"model.diffusion_model.{k}": v for k, v in sd.items()}
|
||||
|
||||
|
||||
class TestModelDetection:
|
||||
"""Verify that first-match model detection selects the correct model
|
||||
based on list ordering and unet_config specificity."""
|
||||
@@ -179,116 +125,6 @@ class TestModelDetection:
|
||||
assert model_config is not None
|
||||
assert type(model_config).__name__ == "FluxSchnell"
|
||||
|
||||
def test_seedvr2_7b_separate_mm_detection_config(self):
|
||||
sd = _make_seedvr2_7b_separate_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config is not None
|
||||
assert unet_config["image_model"] == "seedvr2"
|
||||
assert unet_config["vid_dim"] == 3072
|
||||
assert unet_config["heads"] == 24
|
||||
assert unet_config["num_layers"] == 36
|
||||
assert unet_config["mm_layers"] == 36
|
||||
assert unet_config["mlp_type"] == "normal"
|
||||
assert unet_config["rope_type"] == "rope3d"
|
||||
assert unet_config["rope_dim"] == 64
|
||||
|
||||
def test_seedvr2_7b_shared_mm_detection_config(self):
|
||||
sd = _make_seedvr2_7b_shared_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config is not None
|
||||
assert unet_config["image_model"] == "seedvr2"
|
||||
assert unet_config["vid_dim"] == 3072
|
||||
assert unet_config["heads"] == 24
|
||||
assert unet_config["num_layers"] == 36
|
||||
assert unet_config["mm_layers"] == 10
|
||||
assert unet_config["mlp_type"] == "swiglu"
|
||||
assert unet_config["rope_type"] == "rope3d"
|
||||
assert unet_config["rope_dim"] == 64
|
||||
|
||||
def test_seedvr2_3b_shared_mm_detection_config(self):
|
||||
sd = _make_seedvr2_3b_shared_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config is not None
|
||||
assert unet_config["image_model"] == "seedvr2"
|
||||
assert unet_config["vid_dim"] == 2560
|
||||
assert unet_config["heads"] == 20
|
||||
assert unet_config["num_layers"] == 32
|
||||
assert unet_config["mlp_type"] == "swiglu"
|
||||
|
||||
def test_seedvr2_model_match_requires_conditioning_tensors(self):
|
||||
sd = _make_seedvr2_7b_shared_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "SeedVR2"
|
||||
|
||||
del sd["positive_conditioning"]
|
||||
assert model_config_from_unet_config(unet_config, sd) is None
|
||||
|
||||
def test_seedvr2_model_match_accepts_full_checkpoint_prefix(self):
|
||||
sd = _add_model_diffusion_prefix(_make_seedvr2_7b_shared_mm_sd())
|
||||
|
||||
assert type(model_config_from_unet(sd, "model.diffusion_model.")).__name__ == "SeedVR2"
|
||||
|
||||
def test_pid_v1_5_detection(self):
|
||||
sd = _make_pid_v1_5_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config == {
|
||||
"image_model": "pid",
|
||||
"lq_latent_channels": 16,
|
||||
"lq_hidden_dim": 1024,
|
||||
"latent_spatial_down_factor": 8,
|
||||
"lq_interval": 2,
|
||||
"lq_latent_unpatchify_factor": 1,
|
||||
"lq_conv_padding_mode": "replicate",
|
||||
"lq_gate_per_token": True,
|
||||
"pit_lq_inject": True,
|
||||
"rope_ref_h": 2048,
|
||||
"rope_ref_w": 2048,
|
||||
}
|
||||
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "PiD"
|
||||
|
||||
def test_pid_v1_5_flux2_detection(self):
|
||||
unet_config = detect_unet_config(_make_pid_v1_5_sd(latent_proj_channels=32), "")
|
||||
|
||||
assert unet_config["lq_latent_channels"] == 128
|
||||
assert unet_config["latent_spatial_down_factor"] == 16
|
||||
assert unet_config["lq_latent_unpatchify_factor"] == 2
|
||||
|
||||
def test_pid_v1_5_pixel_adaln_conversion(self):
|
||||
sd = _make_pid_v1_5_sd()
|
||||
model_config = model_config_from_unet_config(detect_unet_config(sd, ""), sd)
|
||||
processed = model_config.process_unet_state_dict(sd)
|
||||
|
||||
assert processed["pixel_blocks.0.attn.q_norm.weight"].shape == (72,)
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_msa.weight"].shape == (12288, 1536)
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_mlp.weight"].shape == (12288, 1536)
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_msa.bias"].shape == (12288,)
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_mlp.bias"].shape == (12288,)
|
||||
|
||||
def test_joyimage_edit_plus_detection(self):
|
||||
sd = _make_joyimage_edit_plus_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config == {
|
||||
"image_model": "joyimage",
|
||||
"in_channels": 16,
|
||||
"hidden_size": 4096,
|
||||
"patch_size": [1, 2, 2],
|
||||
"num_layers": 40,
|
||||
"num_attention_heads": 32,
|
||||
"text_dim": 4096,
|
||||
}
|
||||
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "JoyImage"
|
||||
|
||||
def test_incomplete_joyimage_signature_is_not_detected(self):
|
||||
sd = _make_joyimage_edit_plus_sd()
|
||||
del sd["double_blocks.0.attn.img_attn_q_norm.weight"]
|
||||
assert detect_unet_config(sd, "") is None
|
||||
|
||||
def test_unet_config_and_required_keys_combination_is_unique(self):
|
||||
"""Each model in the registry must have a unique combination of
|
||||
``unet_config`` and ``required_keys``. If two models share the same
|
||||
|
||||
@@ -1,74 +0,0 @@
|
||||
"""Regression tests for the SeedVR2 VAE forward return contract."""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
from comfy.ldm.seedvr.vae import SEEDVR2_LATENT_CHANNELS, VideoAutoencoderKL # noqa: E402
|
||||
|
||||
|
||||
_LATENT_SHAPE = (1, SEEDVR2_LATENT_CHANNELS, 2, 2, 2)
|
||||
_DECODED_SHAPE = (1, 3, 5, 16, 16)
|
||||
_INPUT_ENCODE_SHAPE = (1, 3, 5, 16, 16)
|
||||
_INPUT_DECODE_SHAPE = _LATENT_SHAPE
|
||||
|
||||
|
||||
class _StubVAE(VideoAutoencoderKL):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self._encode_out = torch.zeros(*_LATENT_SHAPE)
|
||||
self._decode_out = torch.zeros(*_DECODED_SHAPE)
|
||||
|
||||
def encode(self, x, return_dict=True):
|
||||
return self._encode_out
|
||||
|
||||
def decode_(self, z, return_dict=True):
|
||||
return self._decode_out
|
||||
|
||||
|
||||
def test_forward_encode_returns_tensor():
|
||||
vae = _StubVAE()
|
||||
x = torch.zeros(*_INPUT_ENCODE_SHAPE)
|
||||
result = vae.forward(x, mode="encode")
|
||||
assert type(result) is torch.Tensor
|
||||
assert result.shape == torch.Size(_LATENT_SHAPE)
|
||||
|
||||
|
||||
def test_forward_decode_returns_tensor():
|
||||
vae = _StubVAE()
|
||||
z = torch.zeros(*_INPUT_DECODE_SHAPE)
|
||||
result = vae.forward(z, mode="decode")
|
||||
assert type(result) is torch.Tensor
|
||||
assert result.shape == torch.Size(_DECODED_SHAPE)
|
||||
|
||||
|
||||
class _TupleReturningStubVAE(VideoAutoencoderKL):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self._encode_tensor = torch.zeros(*_LATENT_SHAPE)
|
||||
self._decode_tensor = torch.zeros(*_DECODED_SHAPE)
|
||||
|
||||
def encode(self, x, return_dict=True):
|
||||
return (self._encode_tensor,)
|
||||
|
||||
def decode_(self, z, return_dict=True):
|
||||
return (self._decode_tensor,)
|
||||
|
||||
|
||||
def test_forward_all_unwraps_one_tuple_at_each_step():
|
||||
vae = _TupleReturningStubVAE()
|
||||
x = torch.zeros(*_INPUT_ENCODE_SHAPE)
|
||||
result = vae.forward(x, mode="all")
|
||||
assert type(result) is torch.Tensor
|
||||
assert result.shape == torch.Size(_DECODED_SHAPE)
|
||||
|
||||
|
||||
def test_forward_rejects_unknown_mode():
|
||||
vae = _StubVAE()
|
||||
with pytest.raises(ValueError, match="Unknown SeedVR2 VAE forward mode"):
|
||||
vae.forward(torch.zeros(*_INPUT_ENCODE_SHAPE), mode="bogus")
|
||||
@@ -1,79 +0,0 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.sd
|
||||
import comfy.supported_models
|
||||
import comfy.ldm.seedvr.model as seedvr_model
|
||||
import comfy.ldm.seedvr.vae as seedvr_vae
|
||||
|
||||
|
||||
def test_seedvr2_fp16_manual_cast_only_for_bf16_device(monkeypatch):
|
||||
bf16_device = object()
|
||||
fp16_device = object()
|
||||
|
||||
monkeypatch.setattr(
|
||||
comfy.supported_models.comfy.model_management,
|
||||
"should_use_bf16",
|
||||
lambda device=None: device is bf16_device,
|
||||
)
|
||||
|
||||
bf16_config = comfy.supported_models.SeedVR2({"image_model": "seedvr2"})
|
||||
bf16_config.set_inference_dtype(torch.float16, None, device=bf16_device)
|
||||
assert bf16_config.manual_cast_dtype is torch.bfloat16
|
||||
|
||||
fp16_config = comfy.supported_models.SeedVR2({"image_model": "seedvr2"})
|
||||
fp16_config.set_inference_dtype(torch.float16, None, device=fp16_device)
|
||||
assert fp16_config.manual_cast_dtype is None
|
||||
|
||||
|
||||
def test_seedvr2_text_conditioning_accepts_cfg1_single_branch():
|
||||
context = torch.arange(6, dtype=torch.float32).reshape(1, 3, 2)
|
||||
|
||||
txt, txt_shape = seedvr_model.NaDiT._resolve_text_conditioning(object(), context, [0])
|
||||
|
||||
torch.testing.assert_close(txt, context.squeeze(0))
|
||||
torch.testing.assert_close(txt_shape, torch.tensor([[3]], device=context.device))
|
||||
|
||||
|
||||
def test_seedvr2_vae_decode_memory_covers_full_frame_lab_transfer():
|
||||
wrapper = seedvr_vae.VideoAutoencoderKLWrapper.__new__(seedvr_vae.VideoAutoencoderKLWrapper)
|
||||
latent_channels = seedvr_vae.SEEDVR2_LATENT_CHANNELS
|
||||
estimate = wrapper.comfy_memory_used_decode((1, latent_channels, 26, 120, 160))
|
||||
old_estimate = latent_channels * 120 * 160 * (4 * 8 * 8) * 2
|
||||
|
||||
assert estimate == 101 * 960 * 1280 * 160
|
||||
assert estimate > 15 * 1024 ** 3
|
||||
assert estimate > old_estimate * 100
|
||||
|
||||
|
||||
def test_seedvr2_vae_encode_preserves_compute_dtype(monkeypatch):
|
||||
wrapper = seedvr_vae.VideoAutoencoderKLWrapper.__new__(seedvr_vae.VideoAutoencoderKLWrapper)
|
||||
nn.Module.__init__(wrapper)
|
||||
wrapper._dummy = nn.Parameter(torch.empty(1, dtype=torch.float16))
|
||||
input_dtype = None
|
||||
|
||||
def encode(self, x):
|
||||
nonlocal input_dtype
|
||||
input_dtype = x.dtype
|
||||
return x
|
||||
|
||||
monkeypatch.setattr(seedvr_vae.VideoAutoencoderKL, "encode", encode)
|
||||
|
||||
x = torch.zeros((1, 3, 1, 8, 8), dtype=torch.float32)
|
||||
wrapper._encode_with_raw_latent(x)
|
||||
|
||||
assert input_dtype == torch.float32
|
||||
|
||||
|
||||
def test_seedvr2_vae_ops_cast_weights_to_compute_dtype():
|
||||
attention = seedvr_vae.Attention(query_dim=4, heads=1, dim_head=4).to(torch.float16)
|
||||
hidden_states = torch.zeros((1, 2, 4), dtype=torch.float32)
|
||||
|
||||
output = attention(hidden_states)
|
||||
|
||||
assert output.dtype == torch.float32
|
||||
@@ -1,169 +0,0 @@
|
||||
"""SeedVR2 internals regression tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
args.cpu = True
|
||||
|
||||
import comfy.ldm.seedvr.model as seedvr_model # noqa: E402
|
||||
import comfy.ldm.seedvr.vae as vae_mod # noqa: E402
|
||||
import comfy.ldm.modules.attention as attention # noqa: E402
|
||||
import comfy.ops as comfy_ops # noqa: E402
|
||||
from comfy.ldm.seedvr.vae import ( # noqa: E402
|
||||
causal_norm_wrapper,
|
||||
set_norm_limit,
|
||||
)
|
||||
from comfy.ldm.seedvr.attention import var_attention_optimized_split # noqa: E402
|
||||
|
||||
|
||||
_NUM_CHANNELS = 8
|
||||
_NUM_GROUPS = 4
|
||||
_TENSOR_SHAPE = (1, 8, 2, 4, 4)
|
||||
|
||||
_GROUPNORM_SUBCLASSES = [
|
||||
pytest.param(comfy_ops.disable_weight_init.GroupNorm, id="disable_weight_init"),
|
||||
pytest.param(comfy_ops.manual_cast.GroupNorm, id="manual_cast"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("groupnorm_cls", _GROUPNORM_SUBCLASSES)
|
||||
def test_seedvr_groupnorm_low_limit_uses_chunked_groupnorm_path(groupnorm_cls):
|
||||
real_group_norm = vae_mod.F.group_norm
|
||||
set_norm_limit(1e-9)
|
||||
try:
|
||||
gn = groupnorm_cls(num_channels=_NUM_CHANNELS, num_groups=_NUM_GROUPS)
|
||||
gn.eval()
|
||||
|
||||
forward_hook_calls = []
|
||||
|
||||
def _hook(module, inputs, output):
|
||||
forward_hook_calls.append(tuple(inputs[0].shape))
|
||||
|
||||
spy_calls = []
|
||||
|
||||
def _group_norm_spy(input_tensor, num_groups_arg, *args, **kwargs):
|
||||
spy_calls.append({"num_groups": int(num_groups_arg)})
|
||||
return real_group_norm(input_tensor, num_groups_arg, *args, **kwargs)
|
||||
|
||||
handle = gn.register_forward_hook(_hook)
|
||||
try:
|
||||
with patch.object(vae_mod.F, "group_norm", side_effect=_group_norm_spy):
|
||||
out_tensor = causal_norm_wrapper(gn, torch.randn(*_TENSOR_SHAPE))
|
||||
finally:
|
||||
handle.remove()
|
||||
|
||||
full_calls = len(forward_hook_calls)
|
||||
chunked_calls = sum(1 for entry in spy_calls if entry["num_groups"] < _NUM_GROUPS)
|
||||
|
||||
assert tuple(int(s) for s in out_tensor.shape) == _TENSOR_SHAPE
|
||||
assert full_calls == 0, (
|
||||
f"low-limit GroupNorm gate must NOT take the full-forward path; got full_calls={full_calls}"
|
||||
)
|
||||
assert chunked_calls > 0, (
|
||||
f"low-limit GroupNorm gate must take the chunked path; got chunked_calls={chunked_calls}"
|
||||
)
|
||||
finally:
|
||||
set_norm_limit(None)
|
||||
|
||||
|
||||
def test_seedvr2_7b_swin_attention_forward_uses_optimized_var_attention(monkeypatch):
|
||||
dim = 8
|
||||
heads = 2
|
||||
head_dim = 4
|
||||
attn = seedvr_model.NaSwinAttention(
|
||||
vid_dim=dim,
|
||||
txt_dim=dim,
|
||||
heads=heads,
|
||||
head_dim=head_dim,
|
||||
qk_bias=False,
|
||||
qk_norm=comfy_ops.disable_weight_init.RMSNorm,
|
||||
qk_norm_eps=1e-6,
|
||||
rope_type=None,
|
||||
rope_dim=head_dim,
|
||||
shared_weights=False,
|
||||
window=(2, 1, 1),
|
||||
window_method="720pwin_by_size_bysize",
|
||||
version=True,
|
||||
device="cpu",
|
||||
dtype=torch.float32,
|
||||
operations=comfy_ops.disable_weight_init,
|
||||
)
|
||||
generator = torch.Generator(device="cpu").manual_seed(11)
|
||||
vid = torch.randn(8, dim, generator=generator)
|
||||
txt = torch.randn(3, dim, generator=generator)
|
||||
vid_shape = torch.tensor([[2, 2, 2]], dtype=torch.long)
|
||||
txt_shape = torch.tensor([[3]], dtype=torch.long)
|
||||
calls = []
|
||||
|
||||
def fake_optimized_var_attention(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return kwargs["q"]
|
||||
|
||||
monkeypatch.setattr(seedvr_model, "optimized_var_attention", fake_optimized_var_attention)
|
||||
|
||||
vid_out, txt_out = attn(vid, txt, vid_shape, txt_shape, seedvr_model.Cache(disable=True))
|
||||
|
||||
assert tuple(vid_out.shape) == (8, dim)
|
||||
assert tuple(txt_out.shape) == (3, dim)
|
||||
assert len(calls) == 1
|
||||
call = calls[0]
|
||||
assert tuple(call["q"].shape) == (14, heads, head_dim)
|
||||
assert tuple(call["k"].shape) == (14, heads, head_dim)
|
||||
assert tuple(call["v"].shape) == (14, heads, head_dim)
|
||||
assert call["heads"] == heads
|
||||
assert call["skip_reshape"] is True
|
||||
assert call["skip_output_reshape"] is True
|
||||
assert call["cu_seqlens_q"] == [0, 7, 14]
|
||||
assert call["cu_seqlens_k"] == [0, 7, 14]
|
||||
|
||||
|
||||
def test_var_attention_optimized_split_calls_dense_backend_per_window(monkeypatch):
|
||||
heads = 2
|
||||
head_dim = 3
|
||||
q = torch.arange(30, dtype=torch.float32).reshape(5, heads, head_dim)
|
||||
k = q + 100
|
||||
v = q + 200
|
||||
cu = [0, 2, 5]
|
||||
calls = []
|
||||
|
||||
def fake_optimized_attention(q_arg, k_arg, v_arg, heads_arg, **kwargs):
|
||||
calls.append(
|
||||
{
|
||||
"q_shape": tuple(q_arg.shape),
|
||||
"k_shape": tuple(k_arg.shape),
|
||||
"v_shape": tuple(v_arg.shape),
|
||||
"heads": heads_arg,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
)
|
||||
return q_arg + v_arg
|
||||
|
||||
monkeypatch.setattr(attention, "optimized_attention", fake_optimized_attention)
|
||||
|
||||
out = var_attention_optimized_split(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
heads,
|
||||
cu,
|
||||
cu,
|
||||
skip_reshape=True,
|
||||
skip_output_reshape=True,
|
||||
)
|
||||
|
||||
assert tuple(out.shape) == (5, heads, head_dim)
|
||||
assert len(calls) == 2
|
||||
assert calls[0]["q_shape"] == (1, heads, 2, head_dim)
|
||||
assert calls[1]["q_shape"] == (1, heads, 3, head_dim)
|
||||
assert all(call["heads"] == heads for call in calls)
|
||||
assert all(call["kwargs"]["skip_reshape"] is True for call in calls)
|
||||
assert all(call["kwargs"]["skip_output_reshape"] is True for call in calls)
|
||||
torch.testing.assert_close(out, q + v, rtol=0, atol=0)
|
||||
|
||||
@@ -1,320 +0,0 @@
|
||||
"""SeedVR2 model, latent-format, and VAE graph regression tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
args.cpu = True
|
||||
|
||||
import comfy # noqa: E402
|
||||
import comfy.latent_formats # noqa: E402
|
||||
import comfy.ldm.seedvr.model as seedvr_model # noqa: E402
|
||||
import comfy.ldm.seedvr.vae as seedvr_vae_mod # noqa: E402
|
||||
import comfy.model_management # noqa: E402
|
||||
import comfy.ops as comfy_ops # noqa: E402
|
||||
import comfy.sample # noqa: E402
|
||||
import comfy.sd as sd_mod # noqa: E402
|
||||
import nodes as nodes_mod # noqa: E402
|
||||
from comfy.ldm.seedvr.model import NaDiT # noqa: E402
|
||||
|
||||
|
||||
_LATENT_CHANNELS = seedvr_vae_mod.SEEDVR2_LATENT_CHANNELS
|
||||
|
||||
|
||||
def _make_standin(positive_conditioning):
|
||||
class _StandIn(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.register_buffer(
|
||||
"positive_conditioning", positive_conditioning
|
||||
)
|
||||
|
||||
_resolve_text_conditioning = NaDiT._resolve_text_conditioning
|
||||
|
||||
return _StandIn()
|
||||
|
||||
|
||||
class _StubModule(nn.Module):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
|
||||
def _capture_last_layer_flags(monkeypatch, vid_dim: int, txt_in_dim: int) -> list[bool]:
|
||||
flags = []
|
||||
|
||||
class _Block(_StubModule):
|
||||
def __init__(self, *args, **kwargs):
|
||||
flags.append(kwargs["is_last_layer"])
|
||||
super().__init__()
|
||||
|
||||
monkeypatch.setattr(seedvr_model, "NaPatchIn", _StubModule)
|
||||
monkeypatch.setattr(seedvr_model, "NaPatchOut", _StubModule)
|
||||
monkeypatch.setattr(seedvr_model, "TimeEmbedding", _StubModule)
|
||||
monkeypatch.setattr(seedvr_model, "NaMMSRTransformerBlock", _Block)
|
||||
|
||||
seedvr_model.NaDiT(
|
||||
norm_eps=1e-5,
|
||||
num_layers=4,
|
||||
mlp_type="normal",
|
||||
vid_dim=vid_dim,
|
||||
txt_in_dim=txt_in_dim,
|
||||
heads=24,
|
||||
mm_layers=3,
|
||||
operations=comfy_ops.disable_weight_init,
|
||||
)
|
||||
|
||||
return flags
|
||||
|
||||
|
||||
class _Model:
|
||||
def __init__(self, latent_format):
|
||||
self._latent_format = latent_format
|
||||
|
||||
def get_model_object(self, name):
|
||||
assert name == "latent_format"
|
||||
return self._latent_format
|
||||
|
||||
|
||||
class _Patcher:
|
||||
def get_free_memory(self, device):
|
||||
return 1024 * 1024 * 1024
|
||||
|
||||
|
||||
class _EncodeWrapper(seedvr_vae_mod.VideoAutoencoderKLWrapper):
|
||||
def __init__(self, encoded):
|
||||
nn.Module.__init__(self)
|
||||
self.encoded = encoded
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.seen = []
|
||||
|
||||
def encode(self, x):
|
||||
self.seen.append(tuple(x.shape))
|
||||
return self.encoded.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
|
||||
class _DecodeWrapper(seedvr_vae_mod.VideoAutoencoderKLWrapper):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.calls = []
|
||||
|
||||
def decode(self, z, seedvr2_tiling=None):
|
||||
self.calls.append({"shape": tuple(z.shape), "seedvr2_tiling": seedvr2_tiling})
|
||||
if z.ndim == 4:
|
||||
b, tc, h, w = z.shape
|
||||
t = tc // _LATENT_CHANNELS
|
||||
else:
|
||||
b, _, t, h, w = z.shape
|
||||
return torch.zeros(b, 3, t, h * 8, w * 8, dtype=z.dtype, device=z.device)
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_public_encode_returns_tensor(monkeypatch):
|
||||
raw_latent = torch.full((1, _LATENT_CHANNELS, 1, 4, 5), 2.0)
|
||||
seen_shapes = []
|
||||
|
||||
def base_encode(self, x):
|
||||
seen_shapes.append(tuple(x.shape))
|
||||
return raw_latent.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
monkeypatch.setattr(seedvr_vae_mod.VideoAutoencoderKL, "encode", base_encode)
|
||||
|
||||
vae = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
nn.Module.__init__(vae)
|
||||
vae._dummy = nn.Parameter(torch.zeros((), dtype=torch.float32))
|
||||
|
||||
latent = vae.encode(torch.zeros(1, 3, 32, 40))
|
||||
|
||||
assert type(latent) is torch.Tensor
|
||||
assert tuple(latent.shape) == (1, _LATENT_CHANNELS, 4, 5)
|
||||
assert seen_shapes == [(1, 3, 1, 32, 40)]
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_private_encode_helper_keeps_raw_latent(monkeypatch):
|
||||
raw_latent = torch.full((1, _LATENT_CHANNELS, 1, 4, 5), 3.0)
|
||||
|
||||
def base_encode(self, x):
|
||||
return raw_latent.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
monkeypatch.setattr(seedvr_vae_mod.VideoAutoencoderKL, "encode", base_encode)
|
||||
|
||||
vae = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
nn.Module.__init__(vae)
|
||||
vae._dummy = nn.Parameter(torch.zeros((), dtype=torch.float32))
|
||||
|
||||
latent, raw = vae._encode_with_raw_latent(torch.zeros(1, 3, 32, 40))
|
||||
|
||||
assert tuple(latent.shape) == (1, _LATENT_CHANNELS, 4, 5)
|
||||
assert tuple(raw.shape) == (1, _LATENT_CHANNELS, 1, 4, 5)
|
||||
assert torch.equal(raw, raw_latent)
|
||||
|
||||
|
||||
def _make_vae(wrapper):
|
||||
vae = sd_mod.VAE.__new__(sd_mod.VAE)
|
||||
vae.first_stage_model = wrapper
|
||||
vae.device = torch.device("cpu")
|
||||
vae.output_device = torch.device("cpu")
|
||||
vae.vae_dtype = torch.float32
|
||||
vae.latent_channels = _LATENT_CHANNELS
|
||||
vae.latent_dim = 3
|
||||
vae.downscale_ratio = (lambda a: max(0, (a + 3) // 4), 8, 8)
|
||||
vae.upscale_ratio = (lambda a: max(0, a * 4 - 3), 8, 8)
|
||||
vae.output_channels = 3
|
||||
vae.disable_offload = True
|
||||
vae.extra_1d_channel = None
|
||||
vae.crop_input = False
|
||||
vae.not_video = False
|
||||
vae.handles_tiling = isinstance(wrapper, seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
vae.format_encoded = wrapper.comfy_format_encoded
|
||||
vae.patcher = _Patcher()
|
||||
vae.process_input = lambda image: image
|
||||
vae.process_output = lambda image: image.add(1.0).div(2.0).clamp(0.0, 1.0)
|
||||
vae.vae_output_dtype = lambda: torch.float32
|
||||
vae.memory_used_encode = lambda shape, dtype: 1
|
||||
vae.memory_used_decode = lambda shape, dtype: 1
|
||||
vae.throw_exception_if_invalid = lambda: None
|
||||
vae.vae_encode_crop_pixels = lambda pixels: pixels
|
||||
vae.spacial_compression_decode = lambda: 8
|
||||
vae.temporal_compression_decode = lambda: 4
|
||||
return vae
|
||||
|
||||
|
||||
def test_missing_context_falls_back_to_positive_buffer():
|
||||
pos_buffer = torch.full((58, 5120), 7.0)
|
||||
standin = _make_standin(pos_buffer)
|
||||
txt, txt_shape = standin._resolve_text_conditioning(None)
|
||||
assert txt.shape == (58, 5120)
|
||||
assert (txt == 7.0).all(), (
|
||||
"fallback path must use the positive_conditioning buffer "
|
||||
"verbatim, not a zero tensor"
|
||||
)
|
||||
assert txt_shape.shape == (1, 1)
|
||||
assert txt_shape[0, 0].item() == 58
|
||||
|
||||
|
||||
def test_seedvr2_7b_keeps_final_block_text_path(monkeypatch):
|
||||
assert _capture_last_layer_flags(monkeypatch, vid_dim=3072, txt_in_dim=3072) == [
|
||||
False,
|
||||
False,
|
||||
False,
|
||||
False,
|
||||
]
|
||||
|
||||
|
||||
def test_seedvr2_7b_rope3d_matches_wrapper_oracle():
|
||||
rope = seedvr_model.get_na_rope("rope3d", dim=64)
|
||||
generator = torch.Generator(device="cpu").manual_seed(0)
|
||||
q = torch.randn(4, 2, 128, generator=generator)
|
||||
k = torch.randn(4, 2, 128, generator=generator)
|
||||
shape = torch.tensor([[1, 2, 2]], dtype=torch.long)
|
||||
freqs = rope.get_axial_freqs(1, 2, 2).reshape(4, -1)
|
||||
|
||||
expected_q = seedvr_model._apply_seedvr2_rotary_emb(
|
||||
freqs,
|
||||
q.permute(1, 0, 2).float(),
|
||||
).to(q.dtype).permute(1, 0, 2)
|
||||
expected_k = seedvr_model._apply_seedvr2_rotary_emb(
|
||||
freqs,
|
||||
k.permute(1, 0, 2).float(),
|
||||
).to(k.dtype).permute(1, 0, 2)
|
||||
|
||||
actual_q, actual_k = rope(q.clone(), k.clone(), shape, seedvr_model.Cache(disable=True))
|
||||
|
||||
torch.testing.assert_close(actual_q, expected_q, rtol=0, atol=0)
|
||||
torch.testing.assert_close(actual_k, expected_k, rtol=0, atol=0)
|
||||
|
||||
|
||||
def test_seedvr2_forward_requires_conditioning_latents():
|
||||
model = NaDiT.__new__(NaDiT)
|
||||
x = torch.zeros(1, _LATENT_CHANNELS, 1, 4, 5)
|
||||
|
||||
with pytest.raises(ValueError, match="requires conditioning latents"):
|
||||
NaDiT.forward(model, x, timestep=torch.tensor([1.0]), context=None)
|
||||
|
||||
|
||||
def test_seedvr2_latent_format_uses_native_video_latent_shape():
|
||||
latent_format = comfy.latent_formats.SeedVR2()
|
||||
latent_image = torch.zeros(1, 1, 4, 5)
|
||||
|
||||
fixed = comfy.sample.fix_empty_latent_channels(_Model(latent_format), latent_image)
|
||||
|
||||
assert latent_format.latent_channels == _LATENT_CHANNELS
|
||||
assert latent_format.latent_dimensions == 3
|
||||
assert fixed.shape == (1, _LATENT_CHANNELS, 1, 4, 5)
|
||||
|
||||
|
||||
def test_seedvr2_model_requires_native_5d_latent():
|
||||
latent = torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5)
|
||||
assert NaDiT._check_seedvr2_video_latent(latent, _LATENT_CHANNELS, "latent") is latent
|
||||
|
||||
with pytest.raises(ValueError, match="5-D native latent"):
|
||||
NaDiT._check_seedvr2_video_latent(torch.zeros(1, _LATENT_CHANNELS * 2, 4, 5), _LATENT_CHANNELS, "latent")
|
||||
|
||||
|
||||
def test_seedvr2_encode_and_encode_tiled_preserve_native_latent_contract(monkeypatch):
|
||||
monkeypatch.setattr(sd_mod.model_management, "load_models_gpu", lambda *a, **k: None)
|
||||
|
||||
encoded = torch.full((1, _LATENT_CHANNELS, 2, 4, 5), 2.0)
|
||||
vae = _make_vae(_EncodeWrapper(encoded))
|
||||
pixels = torch.zeros(1, 5, 32, 40, 3)
|
||||
|
||||
node_output = nodes_mod.VAEEncode().encode(vae, pixels)[0]
|
||||
node_latent = node_output["samples"]
|
||||
assert set(node_output) == {"samples"}
|
||||
assert tuple(node_latent.shape) == (1, _LATENT_CHANNELS, 2, 4, 5)
|
||||
assert node_latent.dtype == torch.float32
|
||||
assert node_latent.stride()[-1] == 1
|
||||
assert torch.equal(node_latent, torch.full_like(node_latent, 2.0 * seedvr_vae_mod.BYTEDANCE_VAE_SCALING_FACTOR))
|
||||
|
||||
tiled = torch.full((1, _LATENT_CHANNELS, 2, 4, 5), 3.0)
|
||||
monkeypatch.setattr(seedvr_vae_mod, "tiled_vae", MagicMock(return_value=tiled))
|
||||
tiled_output = nodes_mod.VAEEncodeTiled().encode(
|
||||
vae,
|
||||
pixels,
|
||||
tile_size=512,
|
||||
overlap=64,
|
||||
temporal_size=16,
|
||||
temporal_overlap=4,
|
||||
)[0]
|
||||
tiled_latent = tiled_output["samples"]
|
||||
assert set(tiled_output) == {"samples"}
|
||||
assert tuple(tiled_latent.shape) == (1, _LATENT_CHANNELS, 2, 4, 5)
|
||||
assert tiled_latent.dtype == torch.float32
|
||||
assert torch.equal(tiled_latent, torch.full_like(tiled_latent, 3.0 * seedvr_vae_mod.BYTEDANCE_VAE_SCALING_FACTOR))
|
||||
|
||||
|
||||
def test_vaedecode_tiled_spatial_applies_temporal_discarded(monkeypatch):
|
||||
monkeypatch.setattr(sd_mod.model_management, "load_models_gpu", lambda *a, **k: None)
|
||||
vae = _make_vae(_DecodeWrapper())
|
||||
|
||||
nodes_mod.VAEDecodeTiled().decode(
|
||||
vae,
|
||||
{"samples": torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5)},
|
||||
tile_size=512,
|
||||
overlap=64,
|
||||
temporal_size=16,
|
||||
temporal_overlap=4,
|
||||
)
|
||||
|
||||
# Spatial inputs flow through; temporal inputs are discarded as public tiling
|
||||
# knobs, but SeedVR2's internal MemoryState causal slicing is left intact.
|
||||
assert vae.first_stage_model.calls == [
|
||||
{
|
||||
"shape": (1, _LATENT_CHANNELS, 2, 4, 5),
|
||||
"seedvr2_tiling": {
|
||||
"enable_tiling": True,
|
||||
"tile_size": (512, 512),
|
||||
"tile_overlap": (64, 64),
|
||||
"temporal_size": None,
|
||||
"temporal_overlap": None,
|
||||
},
|
||||
}
|
||||
]
|
||||
@@ -1,94 +0,0 @@
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.ldm.seedvr.vae as vae_mod # noqa: E402
|
||||
from comfy_extras import nodes_seedvr # noqa: E402
|
||||
|
||||
|
||||
_LATENT_CHANNELS = vae_mod.SEEDVR2_LATENT_CHANNELS
|
||||
|
||||
|
||||
def _make_wrapper() -> vae_mod.VideoAutoencoderKLWrapper:
|
||||
wrapper = vae_mod.VideoAutoencoderKLWrapper.__new__(
|
||||
vae_mod.VideoAutoencoderKLWrapper
|
||||
)
|
||||
nn.Module.__init__(wrapper)
|
||||
return wrapper
|
||||
|
||||
|
||||
def _fingerprint_decode_(self, z, return_dict=True):
|
||||
b = int(z.shape[0])
|
||||
t = int(z.shape[2])
|
||||
h = int(z.shape[3])
|
||||
w = int(z.shape[4])
|
||||
out = torch.empty(b, 3, t, h * 8, w * 8)
|
||||
for batch_idx in range(b):
|
||||
out[batch_idx].fill_(float(batch_idx + 1))
|
||||
return out
|
||||
|
||||
|
||||
def _decode_with_patches(wrapper, z):
|
||||
with patch.object(vae_mod.VideoAutoencoderKL, "decode_", _fingerprint_decode_):
|
||||
return wrapper.decode(z)
|
||||
|
||||
|
||||
def test_decode_b2_t3_multi_frame_batch_unchanged():
|
||||
wrapper = _make_wrapper()
|
||||
|
||||
out = _decode_with_patches(wrapper, torch.zeros(2, _LATENT_CHANNELS * 3, 2, 2))
|
||||
|
||||
assert tuple(out.shape) == (2, 3, 3, 16, 16)
|
||||
|
||||
|
||||
class _Wrapper(vae_mod.VideoAutoencoderKLWrapper):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self.calls = []
|
||||
|
||||
def parameters(self):
|
||||
return iter([torch.nn.Parameter(torch.zeros(()))])
|
||||
|
||||
def _decode_stub(self, latent):
|
||||
self.calls.append(tuple(latent.shape))
|
||||
return torch.zeros(latent.shape[0], 3, latent.shape[2], latent.shape[3] * 8, latent.shape[4] * 8)
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_decode_accepts_5d_channel_first_latents_without_preprocessor_state():
|
||||
wrapper = _Wrapper()
|
||||
|
||||
with patch.object(vae_mod.VideoAutoencoderKL, "decode_", _decode_stub):
|
||||
out = wrapper.decode(torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5))
|
||||
|
||||
assert tuple(out.shape) == (1, 3, 2, 32, 40)
|
||||
assert wrapper.calls == [(1, _LATENT_CHANNELS, 2, 4, 5)]
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_decode_rejects_wrong_rank_latents():
|
||||
wrapper = _Wrapper()
|
||||
|
||||
with pytest.raises(RuntimeError, match=r"latent input must be 4-D collapsed .* or 5-D"):
|
||||
wrapper.decode(torch.zeros(1, _LATENT_CHANNELS, 4))
|
||||
|
||||
|
||||
def _t_padded(t_in: int) -> int:
|
||||
if t_in == 1:
|
||||
return 1
|
||||
if t_in <= 4:
|
||||
return 5
|
||||
if (t_in - 1) % 4 == 0:
|
||||
return t_in
|
||||
return t_in + (4 - ((t_in - 1) % 4))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("t_in", [1, 5, 9])
|
||||
def test_t_padded_matches_cut_videos(t_in):
|
||||
dummy = torch.zeros(1, t_in, 1, 1, 1)
|
||||
assert nodes_seedvr.cut_videos(dummy).shape[1] == _t_padded(t_in)
|
||||
@@ -1,407 +0,0 @@
|
||||
from contextlib import ExitStack
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.ldm.seedvr.vae as vae_mod # noqa: E402
|
||||
import comfy.ldm.seedvr.vae as seedvr_vae_mod # noqa: E402
|
||||
import comfy.sd as sd_mod # noqa: E402
|
||||
from comfy.ldm.seedvr.vae import MemoryState, tiled_vae # noqa: E402
|
||||
|
||||
|
||||
_LATENT_CHANNELS = seedvr_vae_mod.SEEDVR2_LATENT_CHANNELS
|
||||
|
||||
|
||||
def test_runtime_decode_zero_temporal_size_preserves_model_slicing():
|
||||
class StubVAEModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.slicing_latent_min_size = 2
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.device = torch.device("cpu")
|
||||
self.use_slicing = True
|
||||
self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32))
|
||||
self.decode_min_sizes = []
|
||||
self.memory_states = []
|
||||
|
||||
def decode_(self, t_chunk):
|
||||
self.decode_min_sizes.append(self.slicing_latent_min_size)
|
||||
return vae_mod.VideoAutoencoderKL.slicing_decode(self, t_chunk)
|
||||
|
||||
def _decode(self, z, memory_state=MemoryState.DISABLED, memory_cache=None):
|
||||
self.memory_states.append(memory_state)
|
||||
b, c, d, h, w = z.shape
|
||||
return torch.zeros((b, 3, d, h * 8, w * 8), dtype=z.dtype)
|
||||
|
||||
vae = StubVAEModel()
|
||||
z = torch.zeros((1, _LATENT_CHANNELS, 5, 8, 8), dtype=torch.float32)
|
||||
|
||||
tiled_vae(
|
||||
z,
|
||||
vae,
|
||||
tile_size=(64, 64),
|
||||
tile_overlap=(0, 0),
|
||||
temporal_size=0,
|
||||
temporal_overlap=0,
|
||||
encode=False,
|
||||
)
|
||||
|
||||
assert vae.decode_min_sizes == [2]
|
||||
assert vae.memory_states == [MemoryState.INITIALIZING, MemoryState.ACTIVE]
|
||||
assert vae.slicing_latent_min_size == 2
|
||||
|
||||
|
||||
def test_zero_temporal_size_preserves_min_size_when_encode_raises():
|
||||
class RaisingVAEModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.slicing_sample_min_size = 4
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.device = torch.device("cpu")
|
||||
self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32))
|
||||
|
||||
def encode(self, t_chunk):
|
||||
raise RuntimeError("simulated encode failure")
|
||||
|
||||
vae = RaisingVAEModel()
|
||||
x = torch.zeros((1, 3, 12, 64, 64), dtype=torch.float32)
|
||||
|
||||
with pytest.raises(RuntimeError, match="simulated encode failure"):
|
||||
tiled_vae(
|
||||
x,
|
||||
vae,
|
||||
tile_size=(64, 64),
|
||||
tile_overlap=(0, 0),
|
||||
temporal_size=0,
|
||||
temporal_overlap=0,
|
||||
encode=True,
|
||||
)
|
||||
|
||||
assert vae.slicing_sample_min_size == 4
|
||||
|
||||
|
||||
def test_tiled_vae_encode_uses_tensor_return_without_indexing():
|
||||
class TensorEncodeVAEModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.slicing_sample_min_size = 4
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.device = torch.device("cpu")
|
||||
self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32))
|
||||
self.calls = []
|
||||
|
||||
def encode(self, t_chunk):
|
||||
self.calls.append(tuple(t_chunk.shape))
|
||||
b, _, _, h, w = t_chunk.shape
|
||||
return torch.ones((b, _LATENT_CHANNELS, 1, h // 8, w // 8), dtype=t_chunk.dtype)
|
||||
|
||||
vae = TensorEncodeVAEModel()
|
||||
x = torch.zeros((2, 3, 1, 64, 64), dtype=torch.float32)
|
||||
|
||||
out = tiled_vae(
|
||||
x,
|
||||
vae,
|
||||
tile_size=(64, 64),
|
||||
tile_overlap=(0, 0),
|
||||
temporal_size=0,
|
||||
temporal_overlap=0,
|
||||
encode=True,
|
||||
)
|
||||
|
||||
assert vae.calls == [(2, 3, 1, 64, 64)]
|
||||
assert tuple(out.shape) == (2, _LATENT_CHANNELS, 1, 8, 8)
|
||||
|
||||
|
||||
def test_tiled_vae_preserves_compute_dtype_with_different_parameter_dtype():
|
||||
class DummyVAE(nn.Module):
|
||||
spatial_downsample_factor = 8
|
||||
temporal_downsample_factor = 4
|
||||
slicing_sample_min_size = 8
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.device = torch.device("cpu")
|
||||
self._dummy = nn.Parameter(torch.zeros(1, dtype=torch.float16))
|
||||
self.input_dtype = None
|
||||
|
||||
def encode(self, t_chunk):
|
||||
self.input_dtype = t_chunk.dtype
|
||||
b, _, _, h, w = t_chunk.shape
|
||||
return torch.ones((b, _LATENT_CHANNELS, 1, h // 8, w // 8), dtype=t_chunk.dtype)
|
||||
|
||||
vae = DummyVAE()
|
||||
x = torch.zeros((1, 3, 1, 64, 64), dtype=torch.float32)
|
||||
|
||||
tiled_vae(x, vae, tile_size=(64, 64), tile_overlap=(16, 16), encode=True)
|
||||
|
||||
assert vae.input_dtype == torch.float32
|
||||
|
||||
|
||||
def test_tiled_vae_preserves_input_dtype_on_single_tile():
|
||||
class FloatOutputVAEModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.slicing_sample_min_size = 4
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.device = torch.device("cpu")
|
||||
self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32))
|
||||
|
||||
def encode(self, t_chunk):
|
||||
b, _, _, h, w = t_chunk.shape
|
||||
return torch.ones((b, _LATENT_CHANNELS, 1, h // 8, w // 8), dtype=torch.float32)
|
||||
|
||||
out = tiled_vae(
|
||||
torch.zeros((1, 3, 1, 64, 64), dtype=torch.float16),
|
||||
FloatOutputVAEModel(),
|
||||
tile_size=(64, 64),
|
||||
tile_overlap=(0, 0),
|
||||
temporal_size=0,
|
||||
temporal_overlap=0,
|
||||
encode=True,
|
||||
)
|
||||
|
||||
assert out.dtype == torch.float16
|
||||
|
||||
|
||||
class _SlicingDecodeVAE(nn.Module):
|
||||
def __init__(self, slicing_latent_min_size):
|
||||
super().__init__()
|
||||
self.slicing_latent_min_size = slicing_latent_min_size
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.device = torch.device("cpu")
|
||||
self.use_slicing = True
|
||||
self._dummy = nn.Parameter(torch.zeros(1, dtype=torch.float32))
|
||||
self.decode_min_sizes = []
|
||||
self.memory_states = []
|
||||
|
||||
def decode_(self, z):
|
||||
self.decode_min_sizes.append(self.slicing_latent_min_size)
|
||||
return vae_mod.VideoAutoencoderKL.slicing_decode(self, z)
|
||||
|
||||
def _decode(self, z, memory_state=MemoryState.DISABLED, memory_cache=None):
|
||||
self.memory_states.append(memory_state)
|
||||
x = z[:, :1].repeat(
|
||||
1,
|
||||
3,
|
||||
1,
|
||||
self.spatial_downsample_factor,
|
||||
self.spatial_downsample_factor,
|
||||
)
|
||||
return x
|
||||
|
||||
|
||||
def test_decode_tiled_vae_maps_temporal_args_to_latent_slicing_min_size():
|
||||
vae = _SlicingDecodeVAE(slicing_latent_min_size=2)
|
||||
z = torch.arange(
|
||||
_LATENT_CHANNELS * 5 * 8 * 8,
|
||||
dtype=torch.float32,
|
||||
).reshape(1, _LATENT_CHANNELS, 5, 8, 8)
|
||||
|
||||
tiled_vae(
|
||||
z,
|
||||
vae,
|
||||
tile_size=(64, 64),
|
||||
tile_overlap=(0, 0),
|
||||
temporal_size=12,
|
||||
temporal_overlap=4,
|
||||
encode=False,
|
||||
)
|
||||
|
||||
assert vae.decode_min_sizes == [2]
|
||||
assert vae.memory_states == [MemoryState.INITIALIZING, MemoryState.ACTIVE]
|
||||
assert vae.slicing_latent_min_size == 2
|
||||
|
||||
wrapper = vae_mod.VideoAutoencoderKLWrapper.__new__(
|
||||
vae_mod.VideoAutoencoderKLWrapper
|
||||
)
|
||||
nn.Module.__init__(wrapper)
|
||||
seedvr2_tiling = {
|
||||
"enable_tiling": True,
|
||||
"tile_size": (64, 64),
|
||||
"tile_overlap": (0, 0),
|
||||
"temporal_size": 8,
|
||||
"temporal_overlap": 7,
|
||||
}
|
||||
|
||||
captured = {}
|
||||
|
||||
def _fake_tiled_vae(latent, model, **kwargs):
|
||||
captured.update(kwargs)
|
||||
return torch.zeros(1, 3, 1, 16, 16)
|
||||
|
||||
with patch.object(vae_mod, "tiled_vae", side_effect=_fake_tiled_vae):
|
||||
wrapper.decode(torch.zeros(1, _LATENT_CHANNELS, 2, 2), seedvr2_tiling=seedvr2_tiling)
|
||||
|
||||
assert captured["temporal_overlap"] == 7
|
||||
|
||||
|
||||
def _force_oom(*a, **k):
|
||||
raise torch.cuda.OutOfMemoryError("forced OOM for dispatcher test")
|
||||
|
||||
|
||||
def _make_vae(first_stage_model, latent_channels, latent_dim):
|
||||
vae = sd_mod.VAE.__new__(sd_mod.VAE)
|
||||
vae.first_stage_model = first_stage_model
|
||||
vae.patcher = MagicMock()
|
||||
vae.patcher.get_free_memory = MagicMock(return_value=8 * 1024 * 1024 * 1024)
|
||||
vae.device = vae.output_device = torch.device("cpu")
|
||||
vae.vae_dtype = torch.float32
|
||||
vae.disable_offload = True
|
||||
vae.extra_1d_channel = None
|
||||
vae.upscale_ratio = vae.downscale_ratio = 8
|
||||
vae.upscale_index_formula = vae.downscale_index_formula = None
|
||||
vae.output_channels = 3
|
||||
vae.latent_channels = latent_channels
|
||||
vae.latent_dim = latent_dim
|
||||
vae.vae_output_dtype = lambda: torch.float32
|
||||
vae.spacial_compression_decode = lambda: 8
|
||||
vae.handles_tiling = isinstance(first_stage_model, seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
vae.format_encoded = None
|
||||
vae.process_input = lambda x: x
|
||||
vae.process_output = lambda x: x
|
||||
vae.throw_exception_if_invalid = lambda: None
|
||||
vae.memory_used_decode = lambda *a, **k: 1
|
||||
return vae
|
||||
|
||||
|
||||
def _dispatch(vae, samples, seedvr2_call, generic_call, patch_wrapper_decode):
|
||||
mm = sd_mod.model_management
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch.object(mm, "raise_non_oom", lambda e: None))
|
||||
stack.enter_context(patch.object(mm, "load_models_gpu", lambda *a, **k: None))
|
||||
stack.enter_context(patch.object(mm, "soft_empty_cache", lambda: None))
|
||||
stack.enter_context(patch.object(sd_mod.VAE, "_decode_tiled_owned", seedvr2_call))
|
||||
stack.enter_context(patch.object(sd_mod.VAE, "decode_tiled_", generic_call))
|
||||
if patch_wrapper_decode:
|
||||
stack.enter_context(patch.object(
|
||||
seedvr_vae_mod.VideoAutoencoderKLWrapper, "decode",
|
||||
side_effect=_force_oom))
|
||||
vae.decode(samples)
|
||||
|
||||
|
||||
def test_4d_seedvr2_latent_routes_to_owned_decode_tiled():
|
||||
wrapper = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(
|
||||
seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
vae = _make_vae(wrapper, latent_channels=_LATENT_CHANNELS, latent_dim=3)
|
||||
seedvr2_call = MagicMock(return_value=torch.zeros(1, 3, 9, 64, 64))
|
||||
generic_call = MagicMock(return_value=torch.zeros(1, 3, 64, 64))
|
||||
_dispatch(vae, torch.zeros(1, _LATENT_CHANNELS * 3, 8, 8), seedvr2_call, generic_call, True)
|
||||
assert seedvr2_call.call_count == 1
|
||||
assert generic_call.call_count == 0
|
||||
|
||||
|
||||
def test_4d_non_seedvr2_latent_still_routes_to_generic_decode_tiled():
|
||||
first_stage = MagicMock()
|
||||
first_stage.decode = MagicMock(side_effect=_force_oom)
|
||||
vae = _make_vae(first_stage, latent_channels=4, latent_dim=2)
|
||||
seedvr2_call = MagicMock(return_value=torch.zeros(1, 3, 9, 64, 64))
|
||||
generic_call = MagicMock(return_value=torch.zeros(1, 3, 64, 64))
|
||||
_dispatch(vae, torch.zeros(1, 4, 8, 8), seedvr2_call, generic_call, False)
|
||||
assert generic_call.call_count == 1
|
||||
assert seedvr2_call.call_count == 0
|
||||
|
||||
|
||||
def _populate_common_vae_attrs_fallback(vae):
|
||||
vae.patcher = MagicMock()
|
||||
vae.patcher.get_free_memory = MagicMock(return_value=8 * 1024 * 1024 * 1024)
|
||||
vae.device = torch.device("cpu")
|
||||
vae.output_device = torch.device("cpu")
|
||||
vae.vae_dtype = torch.float32
|
||||
vae.disable_offload = True
|
||||
vae.extra_1d_channel = None
|
||||
vae.upscale_ratio = 8
|
||||
vae.upscale_index_formula = None
|
||||
vae.output_channels = 3
|
||||
vae.latent_channels = _LATENT_CHANNELS
|
||||
vae.latent_dim = 3
|
||||
vae.downscale_ratio = 8
|
||||
vae.downscale_index_formula = None
|
||||
vae.not_video = False
|
||||
vae.crop_input = False
|
||||
vae.pad_channel_value = None
|
||||
vae.handles_tiling = isinstance(vae.first_stage_model, seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
vae.format_encoded = None
|
||||
|
||||
vae.vae_output_dtype = lambda: torch.float32
|
||||
vae.spacial_compression_encode = lambda: 8
|
||||
vae.process_input = lambda x: x
|
||||
vae.process_output = lambda x: x
|
||||
vae.throw_exception_if_invalid = lambda: None
|
||||
vae.memory_used_encode = lambda *a, **k: 1
|
||||
|
||||
|
||||
def _make_seedvr2_vae_fallback():
|
||||
vae = sd_mod.VAE.__new__(sd_mod.VAE)
|
||||
wrapper = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(
|
||||
seedvr_vae_mod.VideoAutoencoderKLWrapper
|
||||
)
|
||||
vae.first_stage_model = wrapper
|
||||
_populate_common_vae_attrs_fallback(vae)
|
||||
return vae
|
||||
|
||||
|
||||
def _make_non_seedvr2_vae_fallback():
|
||||
vae = sd_mod.VAE.__new__(sd_mod.VAE)
|
||||
vae.first_stage_model = MagicMock()
|
||||
_populate_common_vae_attrs_fallback(vae)
|
||||
return vae
|
||||
|
||||
|
||||
def _force_regular_encode_oom(*args, **kwargs):
|
||||
raise torch.cuda.OutOfMemoryError("forced OOM for dispatcher test")
|
||||
|
||||
|
||||
def test_seedvr2_3d_routes_to_owned_encode_tiled_on_oom():
|
||||
vae = _make_seedvr2_vae_fallback()
|
||||
pixel_samples = torch.zeros((1, 8, 64, 64, 3))
|
||||
|
||||
seedvr2_call = MagicMock(return_value=torch.zeros(1, _LATENT_CHANNELS, 2, 8, 8))
|
||||
generic_call = MagicMock(return_value=torch.zeros(1, _LATENT_CHANNELS, 2, 8, 8))
|
||||
|
||||
with patch.object(sd_mod.model_management, "raise_non_oom",
|
||||
lambda e: None), \
|
||||
patch.object(sd_mod.model_management, "load_models_gpu",
|
||||
lambda *a, **k: None), \
|
||||
patch.object(sd_mod.model_management, "soft_empty_cache",
|
||||
lambda: None), \
|
||||
patch.object(seedvr_vae_mod.VideoAutoencoderKLWrapper, "encode",
|
||||
side_effect=_force_regular_encode_oom), \
|
||||
patch.object(sd_mod.VAE, "_encode_tiled_owned", seedvr2_call), \
|
||||
patch.object(sd_mod.VAE, "encode_tiled_3d", generic_call):
|
||||
vae.encode(pixel_samples)
|
||||
|
||||
assert seedvr2_call.call_count == 1, (
|
||||
f"Expected _encode_tiled_owned to be called once for a SeedVR2 3D "
|
||||
f"input under OOM fallback; got {seedvr2_call.call_count} calls."
|
||||
)
|
||||
assert generic_call.call_count == 0, (
|
||||
f"encode_tiled_3d must NOT be called for a SeedVR2 input; got "
|
||||
f"{generic_call.call_count} calls."
|
||||
)
|
||||
|
||||
|
||||
def test_non_seedvr2_encode_tiled_3d_default_overlap_is_concrete():
|
||||
vae = _make_non_seedvr2_vae_fallback()
|
||||
vae.downscale_ratio = (lambda a: max(1, a // 4), 8, 8)
|
||||
vae.upscale_ratio = (lambda a: a * 4, 8, 8)
|
||||
generic_call = MagicMock(return_value=torch.zeros(1, _LATENT_CHANNELS, 2, 8, 8))
|
||||
pixel_samples = torch.zeros((1, 8, 64, 64, 3))
|
||||
|
||||
with patch.object(sd_mod.model_management, "load_models_gpu",
|
||||
lambda *a, **k: None), \
|
||||
patch.object(sd_mod.VAE, "encode_tiled_3d", generic_call):
|
||||
vae.encode_tiled(pixel_samples)
|
||||
|
||||
assert generic_call.call_args.kwargs["overlap"] == (1, 64, 64)
|
||||
@@ -818,30 +818,6 @@ class TestExecution:
|
||||
except urllib.error.HTTPError:
|
||||
pass # Expected behavior
|
||||
|
||||
def test_cached_outputs_in_job_without_client_id(self, client: ComfyClient, builder: GraphBuilder):
|
||||
g = builder
|
||||
image = g.node("StubImage", content="BLACK", height=32, width=32, batch_size=1)
|
||||
output = g.node("SaveImage", images=image.out(0))
|
||||
|
||||
# Prime the cache with a normal run.
|
||||
client.run(g)
|
||||
|
||||
# Resubmit anonymously (no client_id) so output nodes are cache hits with no websocket client.
|
||||
data = json.dumps({"prompt": g.finalize()}).encode('utf-8')
|
||||
req = urllib.request.Request(f"http://{client.server_address}/prompt", data=data)
|
||||
prompt_id = json.loads(urllib.request.urlopen(req).read())['prompt_id']
|
||||
|
||||
for _ in range(100):
|
||||
job = client.get_job(prompt_id)
|
||||
if job is not None and job['status'] not in ('pending', 'in_progress'):
|
||||
break
|
||||
time.sleep(0.1)
|
||||
else:
|
||||
raise AssertionError("Prompt did not complete in time")
|
||||
|
||||
assert job['status'] == 'completed'
|
||||
assert output.id in job['outputs'], "Cached outputs must appear in job outputs without a client_id"
|
||||
|
||||
def _create_history_item(self, client, builder):
|
||||
g = GraphBuilder(prefix="offset_test")
|
||||
input_node = g.node(
|
||||
|
||||
Reference in New Issue
Block a user