Compare commits

..

43 Commits

Author SHA1 Message Date
Jukka Seppänen 806e092ed4
Fix MageFlow on cards that don't support bf16 (#15081) 2026-07-25 21:01:51 -04:00
comfyanonymous f966a2b38c
Optimize ideogram model using comfy kitchen rms rope. (#15080) 2026-07-25 15:09:26 -04:00
Alexander Piskun fad06e5da4
[Partner Nodes] feat(Anthropic): add Claude Opus 5 to OpenRouter node (#15075)
Signed-off-by: bigcat88 <bigcat88@icloud.com>
2026-07-25 20:25:58 +03:00
Jukka Seppänen 6f6c500c15
Improve LTXV IC-lora detection (#15073) 2026-07-25 14:30:37 +03:00
Jukka Seppänen 45ffd5430b
feat: Support MageFlow (CORE-372) (#15026) 2026-07-24 23:14:01 -04:00
rattus 36aec0d086
cli_args: bump clamp to 128BGB (#15068)
Some long running chaos testing on a 512GB RAM RTX6000 pro showed that this
is a little bit too low for common template workflows switching around. The
original number was just a guess from me, so go with the scientific result
instead.
2026-07-24 19:48:52 -04:00
rattus f8a3fd9d79
upscalers: convert latent_upsampler model to DynamicVram (#15063)
These were alll non-dynamic (some non-ModelPatcher) code path calling
FreeMemory for management requiring up-front memory freeing. Convert it
to dynamic to avoid legacy free behaviour mixing into otherwise
dynamic workflows.
2026-07-24 13:34:40 -07:00
comfyanonymous 7c59a078d6
Use comfy kitchen rope functions in ltx models. (#15056) 2026-07-24 15:17:31 -04:00
Daxiong (Lin) c0ca3a5991
chore: update workflow templates to v0.11.17 (#15059) 2026-07-24 19:13:09 +08:00
comfyanonymous 0cb84e7e6e
Make Ernie use comfy kitchen rms rope (#15055) 2026-07-23 22:06:52 -04:00
Alexander Piskun feca51a854
[Partner Nodes] chore(Runway): deprecate Gen3a model (#15050)
Signed-off-by: bigcat88 <bigcat88@icloud.com>
2026-07-23 19:27:55 +03:00
Alexander Piskun 7cbe047447
[Partner Nodes] feat(ByteDance): add new "seed-audio-1.0-multilingual" model (#15034)
Signed-off-by: bigcat88 <bigcat88@icloud.com>
2026-07-23 18:32:46 +03:00
Comfy Org PR Bot a449f5f987
Bump comfyui-frontend-package to 1.47.10 (#15045) 2026-07-23 10:59:00 +02:00
comfyanonymous 2e47082c8e
Make z image/lumina 2 models use comfy kitchen rms rope. (#15036) 2026-07-22 15:34:27 -04:00
Daxiong (Lin) 54ca9193a3
chore: update workflow templates to v0.11.15 (#15030) 2026-07-23 00:49:39 +08:00
Alexander Piskun f6d30bce9a
[Partner Nodes] feat(Anthropic): add new models (#15023)
Signed-off-by: bigcat88 <bigcat88@icloud.com>
2026-07-22 17:31:26 +03:00
Alexander Piskun ba5226db96
[Partner Nodes] feat(Openrouter): add new models (#15021)
Signed-off-by: bigcat88 <bigcat88@icloud.com>
2026-07-22 17:20:18 +03:00
comfyanonymous 947c2749dd
Use optimized rms_rope function in joyai image model. (#15018) 2026-07-21 20:02:45 -07:00
cloud-code-bot[bot] 7bf8bfcd07
ci: bump cursor-review to github-workflows@964d5aa (#15017) 2026-07-21 12:00:10 -07:00
Alexander Piskun 78b43d2500
[Partner Nodes] fix(Gemini-Omni): pass videos as inline data (#15014)
Signed-off-by: bigcat88 <bigcat88@icloud.com>
2026-07-21 18:19:53 +03:00
Barish Ozbay ac3a7a654f
Add native Uni3C Controlnet support for Wan models (CORE-365) (#14946)
* Add native Uni3C controlnet support for Wan models

* Dispatch double_block patches in all Wan model variants

* Remove unused grid_sizes assignment in CameraWanModel, WanModel_S2V, HumoWanModel, and AnimateWanModel
2026-07-21 15:44:14 +03:00
TheToxin-git 593786e489
FreSca: 5D+ (ex. Anima) fix, model-agnostic iteration (#15007)
* FreSca: Make fresca work on multi dim
2026-07-21 14:43:34 +03:00
Kohaku-Blueleaf d0fec2ef7e
[Trainer,Dataset/Feature] Video processing nodes, Image Processing Node video support, trainer video support (CORE-81) (#13588) 2026-07-21 00:02:54 -04:00
Matt Miller 0384bb25f4
chore: add /AGENTS.md to CODEOWNERS (#14962)
Scope AGENTS.md review to @comfyanonymous, matching the existing /CODEOWNERS, /.ci/, and /.github/ meta-file entries.
2026-07-20 23:46:05 -04:00
comfyanonymous 35c94d6023
Fix gfx1035 not being treated like RDNA2 (#15009) 2026-07-20 23:36:03 -04:00
Jukka Seppänen ecba6f2594
feat: Support Gemma4 12B (CORE-277) (#14304) 2026-07-20 19:33:26 -04:00
comfyanonymous 6665515349
Fix wan dancer issue with batches. (#14999) 2026-07-19 15:13:49 -07:00
comfyanonymous c9602625e4
Implement regular and timestep zero reference images to krea 2 for ostris and identity edit ref loras. (#14843) 2026-07-18 17:12:18 -07:00
Alexander Piskun 83082a51c4
[Partner Nodes] fix(Google): switch to Interactions API for Omni model (#14984) 2026-07-18 19:36:34 +03:00
Daxiong (Lin) 7774301a7b
chore: update workflow templates to v0.11.12 (#14986) 2026-07-18 23:44:22 +08:00
comfyanonymous 4800e78518
More comfy-kitchen int8 optimizations. (#14980) 2026-07-18 00:53:15 -04:00
Alexander Piskun b08e6cf35f
[Partner Nodes] feat(HeyGen): add Avatar, Talking Photo, Create Avatar, Video Translate and TTS nodes (#14958)
* [Partner Nodes] feat(HeyGen): add Avatar, Talking Photo, Create Avatar, Video Translate and TTS nodes

Signed-off-by: bigcat88 <bigcat88@icloud.com>

* [Partner Nodes] fix(HeyGen): display only Avatars supported by engine

Signed-off-by: bigcat88 <bigcat88@icloud.com>

* [Partner Nodes] remove 4K option

---------

Signed-off-by: bigcat88 <bigcat88@icloud.com>
Co-authored-by: Alexis Rolland <alexisrolland@hotmail.com>
2026-07-17 16:49:16 -04:00
Daxiong (Lin) 1d1099bea0
chore: update workflow templates to v0.11.11 (#14973) 2026-07-18 00:53:15 +08:00
Comfy Org PR Bot c67f95607f
chore(openapi): sync shared API contract from cloud@4acc59a (#14947) 2026-07-17 18:29:45 +02:00
Alexander Piskun 8edea4a65d
[Partner Nodes] feat(Google): add Gemini 3.5 Flash LLM model (#14972)
Co-authored-by: Alexis Rolland <alexisrolland@hotmail.com>
2026-07-17 19:24:34 +03:00
comfyanonymous 0f42ba5146
Support anima lllite control models. (#14954)
Put them in the models/model_patches folder. Use the new AnimaLLLiteApply node.
2026-07-17 07:36:21 -07:00
comfyanonymous 71b73e3b2b
Speed up anima a bit. (#14953) 2026-07-16 19:44:02 -07:00
comfyanonymous 6a8ff7a929
Various comfy kitchen optimizations and fixes. (#14963) 2026-07-16 19:43:12 -07:00
Alexander Piskun 285a98944c
[Partner Nodes] feat(OpenAI): add GPT5.6 models (#14957)
Signed-off-by: bigcat88 <bigcat88@icloud.com>
2026-07-16 15:35:07 +03:00
彼彼 03978e1e81
[feat]Add JoyImageEdit native model support (#14428) 2026-07-15 23:48:28 -04:00
comfyanonymous 678d42c90e
Update AGENTS.md (#14955) 2026-07-15 20:09:59 -07:00
Alexis Rolland 87d23b8176
[Partner Nodes] feat(client): send ComfyUI Job Id in request headers (#14934) 2026-07-15 16:38:03 +08:00
Alexander Piskun cc6b352511
fix(Video): stream the video transcode instead of buffering every frame in RAM (CORE-353) (CORE-351) (#14813) 2026-07-15 15:23:43 +08:00
56 changed files with 4652 additions and 357 deletions

View File

@ -23,9 +23,9 @@ jobs:
# SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up # SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up
# upstream changes; keep `workflows_ref` matching so prompts/scripts load # upstream changes; keep `workflows_ref` matching so prompts/scripts load
# from the same commit as the workflow definition. # from the same commit as the workflow definition.
uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@047ca48febe3a6647608ed2e0c4331b491cb9d6a # github-workflows#9 uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@964d5aad37cbfb57c5b23961d42c2fd85868bf1d # github-workflows main (964d5aa)
with: with:
workflows_ref: 047ca48febe3a6647608ed2e0c4331b491cb9d6a workflows_ref: 964d5aad37cbfb57c5b23961d42c2fd85868bf1d
diff_excludes: >- diff_excludes: >-
:!**/.claude/** :!**/.claude/**
:!**/dist/** :!**/dist/**

View File

@ -19,6 +19,9 @@
better to remove a broken feature path than keep a complicated partial fix. better to remove a broken feature path than keep a complicated partial fix.
- Preserve existing APIs, node names, model-loading behavior, file layout, and - Preserve existing APIs, node names, model-loading behavior, file layout, and
workflow compatibility unless the change is explicitly about replacing them. 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 - Code must look hand-written for this repository. Changes that read like
generic AI-generated code will be rejected automatically: unnecessary helper generic AI-generated code will be rejected automatically: unnecessary helper
layers, vague names, boilerplate comments, defensive branches without a real layers, vague names, boilerplate comments, defensive branches without a real
@ -96,6 +99,13 @@
unless they are read by current code and change current behavior. Remove unless they are read by current code and change current behavior. Remove
pass-through or stored-but-unused values instead of preserving upstream or pass-through or stored-but-unused values instead of preserving upstream or
deprecated API baggage. 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 - If an implementation needs auxiliary values for its own workflow, expose them
through a private helper or a clearly named implementation-specific method through a private helper or a clearly named implementation-specific method
instead of overloading the public method's return contract. instead of overloading the public method's return contract.
@ -154,6 +164,10 @@
`comfy-kitchen` helpers where they already solve the problem. `comfy-kitchen` helpers where they already solve the problem.
- Use optimized comfy-kitchen ops in places where they improve performance - Use optimized comfy-kitchen ops in places where they improve performance
without changing the expected dtype, device, memory, or interface behavior. 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. - All models should use the optimized attention function selected by ComfyUI.
Treat optimized backend functions, dispatch helpers, and capability-selected Treat optimized backend functions, dispatch helpers, and capability-selected
callables as opaque. Higher-level code must not inspect function identity, callables as opaque. Higher-level code must not inspect function identity,
@ -176,6 +190,12 @@
- Model detection code that inspects linear weight shapes should only use the - Model detection code that inspects linear weight shapes should only use the
first dimension. The second dimension may be half the original size for first dimension. The second dimension may be half the original size for
NVFP4 or other 4-bit quantized models. 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 - Avoid adding `einops` usage in core inference code. Use native torch tensor
ops such as `reshape`, `view`, `permute`, `transpose`, `flatten`, `unflatten`, ops such as `reshape`, `view`, `permute`, `transpose`, `flatten`, `unflatten`,
`unsqueeze`, and `squeeze` instead. `unsqueeze`, and `squeeze` instead.
@ -192,11 +212,23 @@
methods for scalar or structural calculations. methods for scalar or structural calculations.
- Avoid unnecessary casts and transfers. Preserve the intended compute dtype, - Avoid unnecessary casts and transfers. Preserve the intended compute dtype,
storage dtype, bias dtype, and original tensor shape metadata. 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 - 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 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 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 forward; the model path should consume and return the native latent shape for
that model family. 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 - 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 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 defensive or convenience casts in model code; it is better for invalid dtype
@ -260,6 +292,15 @@
- Model implementations should add the minimal number of ComfyUI nodes required - 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 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. 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 - 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 workflow convenience unless the node is explicitly an output node. Existing
models, latents, conditioning, or other inputs should flow directly to the models, latents, conditioning, or other inputs should flow directly to the

View File

@ -1,5 +1,6 @@
* @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai * @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai
/CODEOWNERS @comfyanonymous /CODEOWNERS @comfyanonymous
/AGENTS.md @comfyanonymous
/.ci/ @comfyanonymous /.ci/ @comfyanonymous
/.github/ @comfyanonymous /.github/ @comfyanonymous

View File

@ -112,7 +112,7 @@ parser.add_argument("--preview-method", type=LatentPreviewMethod, default=Latent
parser.add_argument("--preview-size", type=int, default=512, help="Sets the maximum preview size for sampler nodes.") parser.add_argument("--preview-size", type=int, default=512, help="Sets the maximum preview size for sampler nodes.")
cache_group = parser.add_mutually_exclusive_group() cache_group = parser.add_mutually_exclusive_group()
cache_group.add_argument("--cache-ram", nargs='*', type=float, default=[], metavar="GB", help="Use RAM pressure caching with the specified headroom thresholds. This is the default caching mode. The first value sets the active-cache threshold; the optional second value sets the inactive-cache/pin threshold. Defaults when no values are provided: active 10%% of system RAM (min 2GB, max 10GB), inactive 100%% of system RAM (max 96GB).") cache_group.add_argument("--cache-ram", nargs='*', type=float, default=[], metavar="GB", help="Use RAM pressure caching with the specified headroom thresholds. This is the default caching mode. The first value sets the active-cache threshold; the optional second value sets the inactive-cache/pin threshold. Defaults when no values are provided: active 10%% of system RAM (min 2GB, max 10GB), inactive 100%% of system RAM (max 128GB).")
cache_group.add_argument("--cache-classic", action="store_true", help="Use the old style (aggressive) caching.") cache_group.add_argument("--cache-classic", action="store_true", help="Use the old style (aggressive) caching.")
cache_group.add_argument("--cache-lru", type=int, default=0, help="Use LRU caching with a maximum of N node results cached. May use more RAM/VRAM.") cache_group.add_argument("--cache-lru", type=int, default=0, help="Use LRU caching with a maximum of N node results cached. May use more RAM/VRAM.")
cache_group.add_argument("--cache-none", action="store_true", help="Reduced RAM/VRAM usage at the expense of executing every node for each run.") cache_group.add_argument("--cache-none", action="store_true", help="Reduced RAM/VRAM usage at the expense of executing every node for each run.")

278
comfy/ldm/anima/lllite.py Normal file
View File

@ -0,0 +1,278 @@
import re
import torch
from torch import nn
import torch.nn.functional as F
import comfy.ops
import comfy.utils
MODULE_PATTERN = re.compile(r"lllite_dit_blocks_(\d+)_(self_attn_[qkv]_proj|cross_attn_q_proj|mlp_layer1)$")
def _group_norm(channels, device=None, dtype=None, operations=None):
groups = 8
while groups > 1 and channels % groups != 0:
groups //= 2
return operations.GroupNorm(groups, channels, device=device, dtype=dtype)
class AnimaLLLiteResBlock(nn.Module):
def __init__(self, channels, device=None, dtype=None, operations=None):
super().__init__()
self.norm1 = _group_norm(channels, device=device, dtype=dtype, operations=operations)
self.conv1 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype)
self.norm2 = _group_norm(channels, device=device, dtype=dtype, operations=operations)
self.conv2 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype)
def forward(self, x):
h = self.conv1(F.silu(self.norm1(x)))
h = self.conv2(F.silu(self.norm2(h)))
return x + h
class AnimaLLLiteASPP(nn.Module):
def __init__(self, channels, dilations, device=None, dtype=None, operations=None):
super().__init__()
branches = []
for dilation in dilations:
if dilation == 1:
conv = operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype)
else:
conv = operations.Conv2d(channels, channels, kernel_size=3, padding=dilation, dilation=dilation, device=device, dtype=dtype)
branches.append(nn.Sequential(conv, _group_norm(channels, device=device, dtype=dtype, operations=operations), nn.SiLU()))
self.branches = nn.ModuleList(branches)
self.global_pool = nn.AdaptiveAvgPool2d(1)
self.global_conv = nn.Sequential(
operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype),
_group_norm(channels, device=device, dtype=dtype, operations=operations),
nn.SiLU(),
)
self.proj = nn.Sequential(
operations.Conv2d(channels * (len(dilations) + 1), channels, kernel_size=1, device=device, dtype=dtype),
_group_norm(channels, device=device, dtype=dtype, operations=operations),
nn.SiLU(),
)
def forward(self, x):
height, width = x.shape[-2:]
outputs = [branch(x) for branch in self.branches]
pooled = self.global_conv(self.global_pool(x))
outputs.append(F.interpolate(pooled, size=(height, width), mode="bilinear", align_corners=False))
return self.proj(torch.cat(outputs, dim=1))
class AnimaLLLiteConditioning(nn.Module):
def __init__(self, cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations, device=None, dtype=None, operations=None):
super().__init__()
half_dim = cond_dim // 2
self.conv1 = operations.Conv2d(cond_in_channels, half_dim, kernel_size=4, stride=4, device=device, dtype=dtype)
self.norm1 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations)
self.conv2 = operations.Conv2d(half_dim, half_dim, kernel_size=3, padding=1, device=device, dtype=dtype)
self.norm2 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations)
self.conv3 = operations.Conv2d(half_dim, cond_dim, kernel_size=4, stride=4, device=device, dtype=dtype)
self.norm3 = _group_norm(cond_dim, device=device, dtype=dtype, operations=operations)
self.resblocks = nn.ModuleList([
AnimaLLLiteResBlock(cond_dim, device=device, dtype=dtype, operations=operations)
for _ in range(cond_resblocks)
])
self.aspp = AnimaLLLiteASPP(cond_dim, aspp_dilations, device=device, dtype=dtype, operations=operations) if aspp_dilations else None
self.proj = operations.Conv2d(cond_dim, cond_emb_dim, kernel_size=1, device=device, dtype=dtype)
self.out_norm = operations.LayerNorm(cond_emb_dim, device=device, dtype=dtype)
def forward(self, x):
x = F.silu(self.norm1(self.conv1(x)))
x = F.silu(self.norm2(self.conv2(x)))
x = F.silu(self.norm3(self.conv3(x)))
for block in self.resblocks:
x = block(x)
if self.aspp is not None:
x = self.aspp(x)
x = self.proj(x).flatten(2).transpose(1, 2).contiguous()
return self.out_norm(x)
class AnimaLLLiteModule(nn.Module):
def __init__(self, in_dim, cond_emb_dim, mlp_dim, device=None, dtype=None, operations=None):
super().__init__()
self.down = operations.Linear(in_dim, mlp_dim, device=device, dtype=dtype)
self.mid = operations.Linear(mlp_dim + cond_emb_dim, mlp_dim, device=device, dtype=dtype)
self.cond_to_film = operations.Linear(cond_emb_dim, 2 * mlp_dim, device=device, dtype=dtype)
self.up = operations.Linear(mlp_dim, in_dim, device=device, dtype=dtype)
self.depth_embed = nn.Parameter(torch.empty(cond_emb_dim, device=device, dtype=dtype), requires_grad=False)
def forward(self, x, cond_emb, strength):
original_shape = x.shape
if x.ndim == 5:
x = x.flatten(1, 3)
if x.shape[0] != cond_emb.shape[0]:
if x.shape[0] % cond_emb.shape[0] != 0:
raise ValueError(f"Anima LLLite batch mismatch: model input batch {x.shape[0]}, control batch {cond_emb.shape[0]}")
cond_emb = cond_emb.repeat(x.shape[0] // cond_emb.shape[0], 1, 1)
if x.shape[1] != cond_emb.shape[1]:
raise ValueError(f"Anima LLLite sequence mismatch: model input has {x.shape[1]} tokens, control has {cond_emb.shape[1]}")
cond_local = cond_emb + comfy.ops.cast_to_input(self.depth_embed, cond_emb)
hidden = F.silu(self.down(x))
gamma, beta = self.cond_to_film(cond_local).chunk(2, dim=-1)
hidden = self.mid(torch.cat((cond_local, hidden), dim=-1))
hidden = F.silu(hidden * (1 + gamma) + beta)
x = x + self.up(hidden) * strength
if len(original_shape) == 5:
x = x.reshape(original_shape)
return x
class AnimaLLLite(nn.Module):
def __init__(self, state_dict, metadata, device=None, dtype=None, operations=None):
super().__init__()
metadata = metadata or {}
version = metadata.get("lllite.version", "2")
if version != "2":
raise ValueError(f"Unsupported Anima LLLite version {version!r}; only named-key v2 checkpoints are supported")
module_names = sorted({key.split(".", 1)[0] for key in state_dict if key.startswith("lllite_dit_blocks_")})
if not module_names:
raise ValueError("Anima LLLite checkpoint has no lllite_dit_blocks_* modules")
cond_in_channels = state_dict["lllite_conditioning1.conv1.weight"].shape[1]
cond_dim = state_dict["lllite_conditioning1.conv3.weight"].shape[0]
cond_emb_dim = state_dict["lllite_conditioning1.proj.weight"].shape[0]
resblock_ids = {int(key.split(".")[2]) for key in state_dict if key.startswith("lllite_conditioning1.resblocks.")}
cond_resblocks = max(resblock_ids) + 1 if resblock_ids else 0
use_aspp = any(key.startswith("lllite_conditioning1.aspp.") for key in state_dict)
dilation_string = metadata.get("lllite.aspp_dilations", "1,2,4,8")
aspp_dilations = tuple(int(value) for value in dilation_string.split(",") if value.strip()) if use_aspp else ()
self.cond_in_channels = cond_in_channels
self.inpaint_masked_input = metadata.get("lllite.inpaint_masked_input", "false").lower() == "true"
self.lllite_conditioning1 = AnimaLLLiteConditioning(
cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations,
device=device, dtype=dtype, operations=operations,
)
self.module_names = set()
self.block_count = 0
self.model_dim = None
for name in module_names:
match = MODULE_PATTERN.fullmatch(name)
if match is None:
raise ValueError(f"Unsupported Anima LLLite module name: {name}")
down_shape = state_dict[f"{name}.down.weight"].shape
mlp_dim, in_dim = down_shape
module_cond_dim = state_dict[f"{name}.cond_to_film.weight"].shape[1]
if module_cond_dim != cond_emb_dim:
raise ValueError(f"Anima LLLite conditioning dimension mismatch in {name}: {module_cond_dim} != {cond_emb_dim}")
if self.model_dim is None:
self.model_dim = in_dim
elif self.model_dim != in_dim:
raise ValueError(f"Anima LLLite model dimension mismatch in {name}: {in_dim} != {self.model_dim}")
self.add_module(name, AnimaLLLiteModule(in_dim, cond_emb_dim, mlp_dim, device=device, dtype=dtype, operations=operations))
self.module_names.add(name)
self.block_count = max(self.block_count, int(match.group(1)) + 1)
def encode_conditioning(self, image):
return self.lllite_conditioning1(image)
def apply(self, x, cond_emb, block_index, target, strength):
name = f"lllite_dit_blocks_{block_index}_{target}"
if name not in self.module_names:
return x
return self.get_submodule(name)(x, cond_emb, strength)
class AnimaLLLitePatch:
def __init__(self, model_patch, image, mask, strength, sigma_start, sigma_end):
self.model_patch = model_patch
self.image = image
self.mask = mask
self.strength = strength
self.sigma_start = sigma_start
self.sigma_end = sigma_end
def __call__(self, args):
x = args["x"]
transformer_options = args["transformer_options"]
if self.strength == 0.0:
return args
sigmas = transformer_options.get("sigmas")
if sigmas is not None:
sigma = float(sigmas.max().item())
if not self.sigma_end <= sigma <= self.sigma_start:
return args
if x.shape[2] != 1:
raise ValueError(f"Anima LLLite only supports T=1, got T={x.shape[2]}")
target_height = x.shape[-2] * 8
target_width = x.shape[-1] * 8
image = comfy.utils.common_upscale(
self.image.movedim(-1, 1), target_width, target_height, "bicubic", crop="center"
).clamp(0.0, 1.0)
image = image.to(device=x.device, dtype=x.dtype) * 2.0 - 1.0
if self.model_patch.model.cond_in_channels == 4:
mask = self.mask
if mask.ndim == 3:
mask = mask.unsqueeze(1)
if mask.ndim != 4 or mask.shape[1] != 1:
raise ValueError(f"Anima LLLite mask must have one channel, got shape {tuple(mask.shape)}")
mask = comfy.utils.common_upscale(
mask.float(), target_width, target_height, "nearest-exact", crop="center"
)
if mask.shape[0] != image.shape[0]:
if image.shape[0] % mask.shape[0] != 0:
raise ValueError(
f"Anima LLLite mask batch {mask.shape[0]} cannot be broadcast to image batch {image.shape[0]}"
)
mask = mask.repeat(image.shape[0] // mask.shape[0], 1, 1, 1)
mask = (mask >= 0.5).to(device=x.device, dtype=x.dtype)
if self.model_patch.model.inpaint_masked_input:
image = image * (mask < 0.5).to(image.dtype)
image = torch.cat((image, mask * 2.0 - 1.0), dim=1)
cond_emb = self.model_patch.model.encode_conditioning(image)
transformer_options["model_patch_data"][self] = cond_emb
return args
def to(self, device_or_dtype):
return self
def models(self):
return [self.model_patch]
class AnimaLLLiteAttentionPatch:
def __init__(self, patch, targets):
self.patch = patch
self.targets = targets
def __call__(self, q, k, v, pe=None, attn_mask=None, extra_options=None):
cond_emb = extra_options["model_patch_data"].get(self.patch)
if cond_emb is None:
return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask}
block_index = extra_options["block_index"]
values = {"q": q, "k": k, "v": v}
for value_name, target in self.targets.items():
values[value_name] = self.patch.model_patch.model.apply(
values[value_name], cond_emb, block_index, target, self.patch.strength
)
return {"q": values["q"], "k": values["k"], "v": values["v"], "pe": pe, "attn_mask": attn_mask}
class AnimaLLLiteMLPPatch:
def __init__(self, patch):
self.patch = patch
def __call__(self, args):
cond_emb = args["transformer_options"]["model_patch_data"].get(self.patch)
if cond_emb is None:
return args
args["x"] = self.patch.model_patch.model.apply(
args["x"], cond_emb, args["transformer_options"]["block_index"], "mlp_layer1", self.patch.strength
)
return args

View File

@ -14,6 +14,7 @@ from torchvision import transforms
import comfy.patcher_extension import comfy.patcher_extension
from comfy.ldm.modules.attention import optimized_attention from comfy.ldm.modules.attention import optimized_attention
import comfy.ldm.common_dit import comfy.ldm.common_dit
import comfy.ops
import comfy.quant_ops import comfy.quant_ops
@ -148,11 +149,29 @@ class Attention(nn.Module):
x: torch.Tensor, x: torch.Tensor,
context: Optional[torch.Tensor] = None, context: Optional[torch.Tensor] = None,
rope_emb: Optional[torch.Tensor] = None, rope_emb: Optional[torch.Tensor] = None,
transformer_options: Optional[dict] = {},
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
q = self.q_proj(x)
context = x if context is None else context context = x if context is None else context
k = self.k_proj(context) q_input = x
v = self.v_proj(context) k_input = context
v_input = context
transformer_patches = transformer_options.get("patches", {})
patch_name = "attn1_patch" if self.is_selfattn else "attn2_patch"
if patch_name in transformer_patches:
extra_options = transformer_options.copy()
extra_options["n_heads"] = self.n_heads
extra_options["dim_head"] = self.head_dim
for patch in transformer_patches[patch_name]:
out = patch(q_input, k_input, v_input, pe=rope_emb, attn_mask=None, extra_options=extra_options)
q_input = out.get("q", q_input)
k_input = out.get("k", k_input)
v_input = out.get("v", v_input)
rope_emb = out.get("pe", rope_emb)
q = self.q_proj(q_input)
k = self.k_proj(k_input)
v = self.v_proj(v_input)
q, k, v = map( q, k, v = map(
lambda t: rearrange(t, "b ... (h d) -> b ... h d", h=self.n_heads, d=self.head_dim), lambda t: rearrange(t, "b ... (h d) -> b ... h d", h=self.n_heads, d=self.head_dim),
(q, k, v), (q, k, v),
@ -161,11 +180,16 @@ class Attention(nn.Module):
def apply_norm_and_rotary_pos_emb( def apply_norm_and_rotary_pos_emb(
q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, rope_emb: Optional[torch.Tensor] q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, rope_emb: Optional[torch.Tensor]
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
q = self.q_norm(q)
k = self.k_norm(k)
v = self.v_norm(v) v = self.v_norm(v)
if self.is_selfattn and rope_emb is not None: # only apply to self-attention! if self.is_selfattn and rope_emb is not None: # only apply to self-attention!
q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rope_emb) q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, q, offloadable=True)
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, k, offloadable=True)
q, k = comfy.quant_ops.ck.rms_rope_split_half(q, k, rope_emb, q_scale, k_scale, self.q_norm.eps)
comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream)
comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream)
else:
q = self.q_norm(q)
k = self.k_norm(k)
return q, k, v return q, k, v
q, k, v = apply_norm_and_rotary_pos_emb(q, k, v, rope_emb) q, k, v = apply_norm_and_rotary_pos_emb(q, k, v, rope_emb)
@ -188,7 +212,7 @@ class Attention(nn.Module):
x (Tensor): The query tensor of shape [B, Mq, K] x (Tensor): The query tensor of shape [B, Mq, K]
context (Optional[Tensor]): The key tensor of shape [B, Mk, K] or use x as context [self attention] if None context (Optional[Tensor]): The key tensor of shape [B, Mk, K] or use x as context [self attention] if None
""" """
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb) q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb, transformer_options=transformer_options)
return self.compute_attention(q, k, v, transformer_options=transformer_options) return self.compute_attention(q, k, v, transformer_options=transformer_options)
@ -555,8 +579,14 @@ class Block(nn.Module):
self.layer_norm_mlp, self.layer_norm_mlp,
scale_mlp_B_T_1_1_D, scale_mlp_B_T_1_1_D,
shift_mlp_B_T_1_1_D, shift_mlp_B_T_1_1_D,
) ).to(compute_dtype)
result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D.to(compute_dtype)) patches = transformer_options.get("patches", {})
if "mlp_patch" in patches:
args = {"x": normalized_x_B_T_H_W_D, "transformer_options": transformer_options}
for patch in patches["mlp_patch"]:
args = patch(args)
normalized_x_B_T_H_W_D = args["x"]
result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D)
x_B_T_H_W_D = torch.addcmul(x_B_T_H_W_D, gate_mlp_B_T_1_1_D.to(residual_dtype), result_B_T_H_W_D.to(residual_dtype)) x_B_T_H_W_D = torch.addcmul(x_B_T_H_W_D, gate_mlp_B_T_1_1_D.to(residual_dtype), result_B_T_H_W_D.to(residual_dtype))
return x_B_T_H_W_D return x_B_T_H_W_D
@ -863,11 +893,22 @@ class MiniTrainDIT(nn.Module):
x_B_T_H_W_D.shape == extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape x_B_T_H_W_D.shape == extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape
), f"{x_B_T_H_W_D.shape} != {extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape}" ), f"{x_B_T_H_W_D.shape} != {extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape}"
transformer_options = kwargs.get("transformer_options", {})
patches = transformer_options.get("patches", {})
if "post_input" in patches:
transformer_options = transformer_options.copy()
transformer_options["model_patch_data"] = {}
if "post_input" in patches:
for patch in patches["post_input"]:
out = patch({"img": x_B_T_H_W_D, "x": x_B_C_T_H_W, "transformer_options": transformer_options})
x_B_T_H_W_D = out["img"]
block_kwargs = { block_kwargs = {
"rope_emb_L_1_1_D": rope_emb_L_1_1_D.unsqueeze(1).unsqueeze(0), "rope_emb_L_1_1_D": rope_emb_L_1_1_D.unsqueeze(1).unsqueeze(0),
"adaln_lora_B_T_3D": adaln_lora_B_T_3D, "adaln_lora_B_T_3D": adaln_lora_B_T_3D,
"extra_per_block_pos_emb": extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D, "extra_per_block_pos_emb": extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D,
"transformer_options": kwargs.get("transformer_options", {}), "transformer_options": transformer_options,
} }
# The residual stream for this model has large values. To make fp16 compute_dtype work, we keep the residual stream # The residual stream for this model has large values. To make fp16 compute_dtype work, we keep the residual stream
@ -877,7 +918,8 @@ class MiniTrainDIT(nn.Module):
if x_B_T_H_W_D.dtype == torch.float16: if x_B_T_H_W_D.dtype == torch.float16:
x_B_T_H_W_D = x_B_T_H_W_D.float() x_B_T_H_W_D = x_B_T_H_W_D.float()
for block in self.blocks: for block_index, block in enumerate(self.blocks):
transformer_options["block_index"] = block_index
x_B_T_H_W_D = block( x_B_T_H_W_D = block(
x_B_T_H_W_D, x_B_T_H_W_D,
t_embedding_B_T_D, t_embedding_B_T_D,

View File

@ -5,6 +5,7 @@ import torch.nn.functional as F
from comfy.ldm.modules.attention import optimized_attention from comfy.ldm.modules.attention import optimized_attention
import comfy.model_management import comfy.model_management
import comfy.ops
import comfy.quant_ops import comfy.quant_ops
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor: def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
@ -111,11 +112,17 @@ class ErnieImageAttention(nn.Module):
query = q_flat.view(B, S, self.heads, self.head_dim) query = q_flat.view(B, S, self.heads, self.head_dim)
key = k_flat.view(B, S, self.heads, self.head_dim) key = k_flat.view(B, S, self.heads, self.head_dim)
query = self.norm_q(query) if image_rotary_emb is not None and not comfy.model_management.in_training:
key = self.norm_k(key) q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, query, offloadable=True)
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, key, offloadable=True)
if image_rotary_emb is not None: query, key = comfy.quant_ops.ck.rms_rope_split_half(query, key, image_rotary_emb, q_scale, k_scale, self.norm_q.eps)
query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb) comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream)
comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream)
else:
query = self.norm_q(query)
key = self.norm_k(key)
if image_rotary_emb is not None:
query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb)
q_flat = query.reshape(B, S, -1) q_flat = query.reshape(B, S, -1)
k_flat = key.reshape(B, S, -1) k_flat = key.reshape(B, S, -1)

View File

@ -12,10 +12,13 @@ import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
import comfy.model_management
import comfy.ops
import comfy.patcher_extension import comfy.patcher_extension
import comfy.quant_ops
from comfy.ldm.lumina.model import FeedForward from comfy.ldm.lumina.model import FeedForward
from comfy.ldm.modules.attention import optimized_attention_masked from comfy.ldm.modules.attention import optimized_attention_masked
from comfy.text_encoders.llama import apply_rope, precompute_freqs_cis from comfy.text_encoders.llama import precompute_freqs_cis
# Per-token role indicators # Per-token role indicators
SEQUENCE_PADDING_INDICATOR = -1 SEQUENCE_PADDING_INDICATOR = -1
@ -25,6 +28,22 @@ LLM_TOKEN_INDICATOR = 3
IMAGE_POSITION_OFFSET = 65536 IMAGE_POSITION_OFFSET = 65536
def _split_half_rope_matrix(freqs_cis):
cos, sin, neg_sin = freqs_cis
half_dim = sin.shape[-1]
matrix = torch.stack(
(cos[..., :half_dim], neg_sin, sin, cos[..., half_dim:]), dim=-1
)
return matrix.reshape(*matrix.shape[:-1], 2, 2).unsqueeze(2)
def _apply_rope_split_half1(x, freqs_cis):
x_dtype = x.dtype
x = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(freqs_cis.dtype)
output = freqs_cis[..., 0] * x[..., 0] + freqs_cis[..., 1] * x[..., 1]
return output.movedim(-1, -2).reshape(*x.shape[:-3], -1).to(x_dtype)
class Ideogram4Attention(nn.Module): class Ideogram4Attention(nn.Module):
def __init__(self, hidden_size, num_heads, eps=1e-5, dtype=None, device=None, operations=None): def __init__(self, hidden_size, num_heads, eps=1e-5, dtype=None, device=None, operations=None):
super().__init__() super().__init__()
@ -42,16 +61,23 @@ class Ideogram4Attention(nn.Module):
qkv = self.qkv(x).view(batch_size, seq_len, 3, self.num_heads, self.head_dim) qkv = self.qkv(x).view(batch_size, seq_len, 3, self.num_heads, self.head_dim)
q, k, v = qkv.unbind(dim=2) q, k, v = qkv.unbind(dim=2)
q = self.norm_q(q) if comfy.model_management.in_training:
k = self.norm_k(k) q = _apply_rope_split_half1(self.norm_q(q), freqs_cis)
k = _apply_rope_split_half1(self.norm_k(k), freqs_cis)
else:
q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, q, offloadable=True)
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, k, offloadable=True)
q, k = comfy.quant_ops.ck.rms_rope_split_half(
q, k, freqs_cis, q_scale, k_scale, self.norm_q.eps
)
comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream)
comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream)
# (B, heads, L, head_dim) # (B, heads, L, head_dim)
q = q.transpose(1, 2) q = q.transpose(1, 2)
k = k.transpose(1, 2) k = k.transpose(1, 2)
v = v.transpose(1, 2) v = v.transpose(1, 2)
q, k = apply_rope(q, k, freqs_cis)
out = optimized_attention_masked(q, k, v, self.num_heads, attn_mask, skip_reshape=True, transformer_options=transformer_options) out = optimized_attention_masked(q, k, v, self.num_heads, attn_mask, skip_reshape=True, transformer_options=transformer_options)
return self.o(out) return self.o(out)
@ -181,6 +207,7 @@ class Ideogram4Transformer(nn.Module):
self.head_dim, position_ids[0].transpose(0, 1), self.rope_theta, self.head_dim, position_ids[0].transpose(0, 1), self.rope_theta,
rope_dims=self.mrope_section, interleaved_mrope=True, device=position_ids.device, rope_dims=self.mrope_section, interleaved_mrope=True, device=position_ids.device,
) )
freqs_cis = _split_half_rope_matrix(freqs_cis)
if attn_mask is not None and attn_mask.dtype == torch.bool: if attn_mask is not None and attn_mask.dtype == torch.bool:
attn_mask = torch.zeros_like(attn_mask, dtype=h.dtype).masked_fill_(~attn_mask, -torch.finfo(h.dtype).max) attn_mask = torch.zeros_like(attn_mask, dtype=h.dtype).masked_fill_(~attn_mask, -torch.finfo(h.dtype).max)

454
comfy/ldm/joyimage/model.py Normal file
View File

@ -0,0 +1,454 @@
# 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))
txt_q = self.txt_attn_q_norm(txt_q)
txt_k = self.txt_attn_k_norm(txt_k)
img_q_scale, _, img_q_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_q_norm, img_q, offloadable=True)
img_k_scale, _, img_k_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_k_norm, img_k, offloadable=True)
img_q, img_k = comfy_kitchen.rms_rope(
img_q,
img_k,
image_rotary_emb,
img_q_scale,
img_k_scale,
self.img_attn_q_norm.eps,
)
comfy.ops.uncast_bias_weight(self.img_attn_q_norm, img_q_scale, None, img_q_offload_stream)
comfy.ops.uncast_bias_weight(self.img_attn_k_norm, img_k_scale, None, img_k_offload_stream)
joint_q = torch.cat([img_q, txt_q], dim=1)
joint_k = torch.cat([img_k, txt_k], dim=1)
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]

View File

@ -15,6 +15,7 @@ from einops import rearrange
import comfy.model_management import comfy.model_management
import comfy.patcher_extension import comfy.patcher_extension
import comfy.ldm.common_dit import comfy.ldm.common_dit
import comfy.utils
from comfy.ldm.flux.layers import EmbedND, timestep_embedding from comfy.ldm.flux.layers import EmbedND, timestep_embedding
from comfy.ldm.flux.math import apply_rope from comfy.ldm.flux.math import apply_rope
from comfy.ldm.modules.attention import optimized_attention_masked from comfy.ldm.modules.attention import optimized_attention_masked
@ -73,11 +74,20 @@ class Attention(nn.Module):
self.wo = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype) self.wo = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype)
def forward(self, x, freqs=None, mask=None, transformer_options={}): def forward(self, x, freqs=None, mask=None, transformer_options={}):
transformer_patches = transformer_options.get("patches", {})
extra_options = transformer_options.copy()
q, k, v, gate = self.wq(x), self.wk(x), self.wv(x), self.gate(x) q, k, v, gate = self.wq(x), self.wk(x), self.wv(x), self.gate(x)
q = rearrange(q, "B L (H D) -> B H L D", H=self.heads) q = rearrange(q, "B L (H D) -> B H L D", H=self.heads)
k = rearrange(k, "B L (H D) -> B H L D", H=self.kvheads) k = rearrange(k, "B L (H D) -> B H L D", H=self.kvheads)
v = rearrange(v, "B L (H D) -> B H L D", H=self.kvheads) v = rearrange(v, "B L (H D) -> B H L D", H=self.kvheads)
q, k = self.qknorm(q, k) q, k = self.qknorm(q, k)
if "block_index" in transformer_options and "attn1_patch" in transformer_patches:
for p in transformer_patches["attn1_patch"]:
out = p(q, k, v, pe=freqs, attn_mask=mask, extra_options=extra_options)
q, k, v = out.get("q", q), out.get("k", k), out.get("v", v)
freqs, mask = out.get("pe", freqs), out.get("attn_mask", mask)
if freqs is not None: if freqs is not None:
q, k = apply_rope(q, k, freqs) q, k = apply_rope(q, k, freqs)
if self.kvheads != self.heads: if self.kvheads != self.heads:
@ -86,6 +96,11 @@ class Attention(nn.Module):
v = v.repeat_interleave(rep, dim=1) v = v.repeat_interleave(rep, dim=1)
out = optimized_attention_masked(q, k, v, self.heads, mask=mask, skip_reshape=True, out = optimized_attention_masked(q, k, v, self.heads, mask=mask, skip_reshape=True,
transformer_options=transformer_options) transformer_options=transformer_options)
if "block_index" in transformer_options and "attn1_output_patch" in transformer_patches:
for p in transformer_patches["attn1_output_patch"]:
out = p(out, extra_options)
return self.wo(out * F.sigmoid(gate)) return self.wo(out * F.sigmoid(gate))
@ -158,8 +173,44 @@ class SingleStreamBlock(nn.Module):
self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations) self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations)
self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations) self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations)
def forward(self, x, vec, freqs, mask=None, transformer_options={}): def forward(self, x, vec, freqs, mask=None, timestep_zero_index=None, transformer_options={}):
prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec) prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec)
if timestep_zero_index is not None:
bs = x.shape[0]
ref_prescale = prescale[bs:]
ref_preshift = preshift[bs:]
ref_pregate = pregate[bs:]
ref_postscale = postscale[bs:]
ref_postshift = postshift[bs:]
ref_postgate = postgate[bs:]
prescale = prescale[:bs]
preshift = preshift[:bs]
pregate = pregate[:bs]
postscale = postscale[:bs]
postshift = postshift[:bs]
postgate = postgate[:bs]
pre = self.prenorm(x)
pre[:, :timestep_zero_index].mul_(1 + prescale).add_(preshift)
pre[:, timestep_zero_index:].mul_(1 + ref_prescale).add_(ref_preshift)
attn = self.attn(pre, freqs, mask, transformer_options=transformer_options)
del pre
attn[:, :timestep_zero_index].mul_(pregate)
attn[:, timestep_zero_index:].mul_(ref_pregate)
x = x + attn
del attn
post = self.postnorm(x)
post[:, :timestep_zero_index].mul_(1 + postscale).add_(postshift)
post[:, timestep_zero_index:].mul_(1 + ref_postscale).add_(ref_postshift)
mlp = self.mlp(post)
del post
mlp[:, :timestep_zero_index].mul_(postgate)
mlp[:, timestep_zero_index:].mul_(ref_postgate)
x = x + mlp
del mlp
return x
x = x + pregate * self.attn((1 + prescale) * self.prenorm(x) + preshift, freqs, mask, transformer_options=transformer_options) x = x + pregate * self.attn((1 + prescale) * self.prenorm(x) + preshift, freqs, mask, transformer_options=transformer_options)
x = x + postgate * self.mlp((1 + postscale) * self.postnorm(x) + postshift) x = x + postgate * self.mlp((1 + postscale) * self.postnorm(x) + postshift)
return x return x
@ -181,7 +232,7 @@ class LastLayer(nn.Module):
class SingleStreamDiT(nn.Module): class SingleStreamDiT(nn.Module):
def __init__(self, features=6144, tdim=256, txtdim=2560, heads=48, kvheads=12, multiplier=4, def __init__(self, features=6144, tdim=256, txtdim=2560, heads=48, kvheads=12, multiplier=4,
layers=28, patch=2, channels=16, bias=False, theta=1e3, txtlayers=12, layers=28, patch=2, channels=16, bias=False, theta=1e3, txtlayers=12,
txtheads=20, txtkvheads=20, image_model=None, txtheads=20, txtkvheads=20, default_ref_method=None, image_model=None,
device=None, dtype=None, operations=None, **kwargs): device=None, dtype=None, operations=None, **kwargs):
super().__init__() super().__init__()
self.dtype = dtype self.dtype = dtype
@ -191,6 +242,7 @@ class SingleStreamDiT(nn.Module):
self.heads = heads self.heads = heads
self.txtdim = txtdim self.txtdim = txtdim
self.txtlayers = txtlayers self.txtlayers = txtlayers
self.default_ref_method = default_ref_method
headdim = features // heads headdim = features // heads
axes = [headdim - 12 * (headdim // 16), 6 * (headdim // 16), 6 * (headdim // 16)] axes = [headdim - 12 * (headdim // 16), 6 * (headdim // 16), 6 * (headdim // 16)]
@ -221,61 +273,110 @@ class SingleStreamDiT(nn.Module):
operations.Linear(features, features * 6, device=device, dtype=dtype), operations.Linear(features, features * 6, device=device, dtype=dtype),
) )
def forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs): def forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
return comfy.patcher_extension.WrapperExecutor.new_class_executor( return comfy.patcher_extension.WrapperExecutor.new_class_executor(
self._forward, self._forward,
self, self,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options), comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options),
).execute(x, timesteps, context, attention_mask, transformer_options, **kwargs) ).execute(x, timesteps, context, attention_mask, ref_latents, transformer_options, **kwargs)
def _forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs): def process_img(self, x, index=0):
patch = self.patch
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch))
h, w = x.shape[-2] // patch, x.shape[-1] // patch
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch)
img_ids = torch.zeros(h, w, 3, device=x.device, dtype=torch.float32)
img_ids[..., 0] = index
img_ids[..., 1] = torch.arange(h, device=x.device, dtype=torch.float32)[:, None]
img_ids[..., 2] = torch.arange(w, device=x.device, dtype=torch.float32)[None, :]
return img, img_ids.reshape(1, h * w, 3).repeat(x.shape[0], 1, 1), h, w
def _forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
transformer_options = transformer_options.copy()
temporal = x.ndim == 5 temporal = x.ndim == 5
if temporal: if temporal:
b5, c5, t5, h5, w5 = x.shape b5, c5, t5, h5, w5 = x.shape
x = x.reshape(b5 * t5, c5, h5, w5) x = x.reshape(b5 * t5, c5, h5, w5)
bs, c, H_orig, W_orig = x.shape bs, _, h_orig, w_orig = x.shape
patch = self.patch patch = self.patch
# Pad the latent up to a multiple of patch (as Flux/Lumina/QwenImage do); crop back at the end.
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch))
H, W = x.shape[-2], x.shape[-1]
h_, w_ = H // patch, W // patch
# context arrives as (B, seq, txtlayers*txtdim); reshape to (B, txtlayers, seq, txtdim). # context arrives as (B, seq, txtlayers*txtdim); reshape to (B, txtlayers, seq, txtdim).
context = self._unpack_context(context) context = self._unpack_context(context)
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch) img, imgpos, h_, w_ = self.process_img(x)
img_tokens = img.shape[1]
timestep_zero_index = None
ref_method = kwargs.get("ref_latents_method", self.default_ref_method)
if ref_method is not None and ref_latents is not None and len(ref_latents) > 0:
ref_tokens = []
ref_pos = []
ref_num_tokens = []
for index, ref in enumerate(ref_latents, 1):
if ref.ndim == 5:
rb, rc, rt, rh5, rw5 = ref.shape
ref = ref.reshape(rb * rt, rc, rh5, rw5)
ref = comfy.utils.repeat_to_batch_size(ref, bs)
kontext, kontext_ids, _, _ = self.process_img(ref, index=index)
ref_tokens.append(kontext)
ref_pos.append(kontext_ids)
ref_num_tokens.append(kontext.shape[1])
img = torch.cat([img] + ref_tokens, dim=1)
imgpos = torch.cat([imgpos] + ref_pos, dim=1)
del ref_tokens, ref_pos
if ref_method == "index_timestep_zero":
timestep_zero_index = img_tokens
transformer_options["reference_image_num_tokens"] = ref_num_tokens
img = self.first(img) img = self.first(img)
t = self.tmlp(timestep_embedding(timesteps, self.tdim).unsqueeze(1).to(img.dtype)) t = self.tmlp(timestep_embedding(timesteps, self.tdim).unsqueeze(1).to(img.dtype))
tvec = self.tproj(t) tvec = self.tproj(t)
if timestep_zero_index is not None:
t0 = self.tmlp(timestep_embedding(torch.zeros_like(timesteps), self.tdim).unsqueeze(1).to(img.dtype))
tvec = torch.cat((tvec, self.tproj(t0)), dim=0)
context = self.txtfusion(context, mask=None, transformer_options=transformer_options) context = self.txtfusion(context, mask=None, transformer_options=transformer_options)
context = self.txtmlp(context) context = self.txtmlp(context)
txtlen, imglen = context.shape[1], img.shape[1] txtlen = context.shape[1]
device = context.device
txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32)
patches = transformer_options.get("patches", {})
if "post_input" in patches:
for p in patches["post_input"]:
out = p({"img": img, "txt": context, "img_ids": imgpos, "txt_ids": txtpos, "transformer_options": transformer_options})
img, context = out["img"], out["txt"]
imgpos, txtpos = out["img_ids"], out["txt_ids"]
combined = torch.cat((context, img), dim=1) combined = torch.cat((context, img), dim=1)
del context, img
if timestep_zero_index is not None:
timestep_zero_index += txtlen
# Position ids: text at 0, image at (0, h_idx, w_idx). # Position ids: text at 0, image at (0, h_idx, w_idx).
device = combined.device
txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32)
imgids = torch.zeros(h_, w_, 3, device=device, dtype=torch.float32)
imgids[..., 1] = torch.arange(h_, device=device, dtype=torch.float32)[:, None]
imgids[..., 2] = torch.arange(w_, device=device, dtype=torch.float32)[None, :]
imgpos = imgids.reshape(1, h_ * w_, 3).repeat(bs, 1, 1)
pos = torch.cat((txtpos, imgpos), dim=1) pos = torch.cat((txtpos, imgpos), dim=1)
del txtpos, imgpos
freqs = self.pe_embedder(pos) freqs = self.pe_embedder(pos)
del pos
for block in self.blocks: transformer_options["total_blocks"] = len(self.blocks)
combined = block(combined, tvec, freqs, None, transformer_options=transformer_options) transformer_options["block_type"] = "single"
transformer_options["img_slice"] = [txtlen, combined.shape[1]]
for i, block in enumerate(self.blocks):
transformer_options["block_index"] = i
combined = block(combined, tvec, freqs, None, timestep_zero_index=timestep_zero_index, transformer_options=transformer_options)
final = self.last(combined, t) final = self.last(combined, t)
out = final[:, txtlen:txtlen + imglen, :] del combined
out = final[:, txtlen:txtlen + img_tokens, :]
out = rearrange(out, "b (h w) (c ph pw) -> b c (h ph) (w pw)", out = rearrange(out, "b (h w) (c ph pw) -> b c (h ph) (w pw)",
h=h_, w=w_, ph=patch, pw=patch, c=self.channels) h=h_, w=w_, ph=patch, pw=patch, c=self.channels)
out = out[:, :, :H_orig, :W_orig] # crop padding back off out = out[:, :, :h_orig, :w_orig] # crop padding back off
if temporal: if temporal:
out = out.reshape(b5, t5, self.channels, H_orig, W_orig).movedim(1, 2) out = out.reshape(b5, t5, self.channels, h_orig, w_orig).movedim(1, 2)
return out return out
def _unpack_context(self, context): def _unpack_context(self, context):

View File

@ -6,9 +6,8 @@ import torch
from comfy.ldm.lightricks.model import ( from comfy.ldm.lightricks.model import (
CrossAttention, CrossAttention,
FeedForward, FeedForward,
freqs_cis_matrix,
generate_freq_grid_np, generate_freq_grid_np,
interleaved_freqs_cis,
split_freqs_cis,
) )
from torch import nn from torch import nn
@ -244,12 +243,15 @@ class Embeddings1DConnector(nn.Module):
expected_freqs = dim // 2 expected_freqs = dim // 2
current_freqs = freqs.shape[-1] current_freqs = freqs.shape[-1]
pad_size = expected_freqs - current_freqs pad_size = expected_freqs - current_freqs
cos_freq, sin_freq = split_freqs_cis(
freqs, pad_size, self.num_attention_heads
)
else: else:
cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem) pad_size = dim % n_elem
return cos_freq.to(dtype=out_dtype), sin_freq.to(dtype=out_dtype), self.split_rope return freqs_cis_matrix(
freqs,
pad_size,
self.split_rope,
self.num_attention_heads,
out_dtype,
)
def forward( def forward(
self, self,

View File

@ -97,11 +97,11 @@ class SpatialRationalResampler(nn.Module):
For dims==3, work per-frame for spatial scaling (temporal axis untouched). For dims==3, work per-frame for spatial scaling (temporal axis untouched).
""" """
def __init__(self, mid_channels: int, scale: float): def __init__(self, mid_channels: int, scale: float, operations):
super().__init__() super().__init__()
self.scale = float(scale) self.scale = float(scale)
self.num, self.den = _rational_for_scale(self.scale) self.num, self.den = _rational_for_scale(self.scale)
self.conv = nn.Conv2d( self.conv = operations.Conv2d(
mid_channels, (self.num**2) * mid_channels, kernel_size=3, padding=1 mid_channels, (self.num**2) * mid_channels, kernel_size=3, padding=1
) )
self.pixel_shuffle = PixelShuffleND(2, upscale_factors=(self.num, self.num)) self.pixel_shuffle = PixelShuffleND(2, upscale_factors=(self.num, self.num))
@ -119,18 +119,18 @@ class SpatialRationalResampler(nn.Module):
class ResBlock(nn.Module): class ResBlock(nn.Module):
def __init__( def __init__(
self, channels: int, mid_channels: Optional[int] = None, dims: int = 3 self, channels: int, operations, mid_channels: Optional[int] = None, dims: int = 3
): ):
super().__init__() super().__init__()
if mid_channels is None: if mid_channels is None:
mid_channels = channels mid_channels = channels
Conv = nn.Conv2d if dims == 2 else nn.Conv3d Conv = operations.Conv2d if dims == 2 else operations.Conv3d
self.conv1 = Conv(channels, mid_channels, kernel_size=3, padding=1) self.conv1 = Conv(channels, mid_channels, kernel_size=3, padding=1)
self.norm1 = nn.GroupNorm(32, mid_channels) self.norm1 = operations.GroupNorm(32, mid_channels)
self.conv2 = Conv(mid_channels, channels, kernel_size=3, padding=1) self.conv2 = Conv(mid_channels, channels, kernel_size=3, padding=1)
self.norm2 = nn.GroupNorm(32, channels) self.norm2 = operations.GroupNorm(32, channels)
self.activation = nn.SiLU() self.activation = nn.SiLU()
def forward(self, x: torch.Tensor) -> torch.Tensor: def forward(self, x: torch.Tensor) -> torch.Tensor:
@ -159,6 +159,7 @@ class LatentUpsampler(nn.Module):
def __init__( def __init__(
self, self,
operations,
in_channels: int = 128, in_channels: int = 128,
mid_channels: int = 512, mid_channels: int = 512,
num_blocks_per_stage: int = 4, num_blocks_per_stage: int = 4,
@ -179,34 +180,34 @@ class LatentUpsampler(nn.Module):
self.spatial_scale = float(spatial_scale) self.spatial_scale = float(spatial_scale)
self.rational_resampler = rational_resampler self.rational_resampler = rational_resampler
Conv = nn.Conv2d if dims == 2 else nn.Conv3d Conv = operations.Conv2d if dims == 2 else operations.Conv3d
self.initial_conv = Conv(in_channels, mid_channels, kernel_size=3, padding=1) self.initial_conv = Conv(in_channels, mid_channels, kernel_size=3, padding=1)
self.initial_norm = nn.GroupNorm(32, mid_channels) self.initial_norm = operations.GroupNorm(32, mid_channels)
self.initial_activation = nn.SiLU() self.initial_activation = nn.SiLU()
self.res_blocks = nn.ModuleList( self.res_blocks = nn.ModuleList(
[ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)] [ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)]
) )
if spatial_upsample and temporal_upsample: if spatial_upsample and temporal_upsample:
self.upsampler = nn.Sequential( self.upsampler = nn.Sequential(
nn.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1), operations.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1),
PixelShuffleND(3), PixelShuffleND(3),
) )
elif spatial_upsample: elif spatial_upsample:
if rational_resampler: if rational_resampler:
self.upsampler = SpatialRationalResampler( self.upsampler = SpatialRationalResampler(
mid_channels=mid_channels, scale=self.spatial_scale mid_channels=mid_channels, scale=self.spatial_scale, operations=operations
) )
else: else:
self.upsampler = nn.Sequential( self.upsampler = nn.Sequential(
nn.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1), operations.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1),
PixelShuffleND(2), PixelShuffleND(2),
) )
elif temporal_upsample: elif temporal_upsample:
self.upsampler = nn.Sequential( self.upsampler = nn.Sequential(
nn.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1), operations.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1),
PixelShuffleND(1), PixelShuffleND(1),
) )
else: else:
@ -215,11 +216,14 @@ class LatentUpsampler(nn.Module):
) )
self.post_upsample_res_blocks = nn.ModuleList( self.post_upsample_res_blocks = nn.ModuleList(
[ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)] [ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)]
) )
self.final_conv = Conv(mid_channels, in_channels, kernel_size=3, padding=1) self.final_conv = Conv(mid_channels, in_channels, kernel_size=3, padding=1)
def get_dtype(self):
return getattr(self.initial_conv, "weight_comfy_model_dtype", self.initial_conv.weight.dtype)
def forward(self, latent: torch.Tensor) -> torch.Tensor: def forward(self, latent: torch.Tensor) -> torch.Tensor:
b, c, f, h, w = latent.shape b, c, f, h, w = latent.shape
@ -266,7 +270,7 @@ class LatentUpsampler(nn.Module):
return x return x
@classmethod @classmethod
def from_config(cls, config): def from_config(cls, config, operations):
return cls( return cls(
in_channels=config.get("in_channels", 4), in_channels=config.get("in_channels", 4),
mid_channels=config.get("mid_channels", 128), mid_channels=config.get("mid_channels", 128),
@ -276,6 +280,7 @@ class LatentUpsampler(nn.Module):
temporal_upsample=config.get("temporal_upsample", False), temporal_upsample=config.get("temporal_upsample", False),
spatial_scale=config.get("spatial_scale", 2.0), spatial_scale=config.get("spatial_scale", 2.0),
rational_resampler=config.get("rational_resampler", False), rational_resampler=config.get("rational_resampler", False),
operations=operations,
) )
def config(self): def config(self):

View File

@ -12,6 +12,8 @@ from torch import nn
import comfy.patcher_extension import comfy.patcher_extension
import comfy.ldm.modules.attention import comfy.ldm.modules.attention
import comfy.ldm.common_dit import comfy.ldm.common_dit
import comfy.model_management
import comfy.quant_ops
from .symmetric_patchifier import SymmetricPatchifier, latent_to_pixel_coords from .symmetric_patchifier import SymmetricPatchifier, latent_to_pixel_coords
@ -322,40 +324,42 @@ class FeedForward(nn.Module):
return self.net(x) return self.net(x)
def apply_rotary_emb(input_tensor, freqs_cis): def apply_rotary_emb(input_tensor, freqs_cis):
cos_freqs, sin_freqs = freqs_cis[0], freqs_cis[1] rotation_matrix, split_pe = freqs_cis
split_pe = freqs_cis[2] if len(freqs_cis) > 2 else False original_shape = input_tensor.shape
return ( input_tensor = input_tensor.reshape(
apply_split_rotary_emb(input_tensor, cos_freqs, sin_freqs) input_tensor.shape[0], input_tensor.shape[1], rotation_matrix.shape[2], -1
if split_pe else
apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs)
) )
def apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs): # TODO: remove duplicate funcs and pick the best/fastest one if comfy.model_management.in_training:
t_dup = rearrange(input_tensor, "... (d r) -> ... d r", r=2) if split_pe:
t1, t2 = t_dup.unbind(dim=-1) t = input_tensor.reshape(*input_tensor.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2)
t_dup = torch.stack((-t2, t1), dim=-1) else:
input_tensor_rot = rearrange(t_dup, "... d r -> ... (d r)") t = input_tensor.reshape(*input_tensor.shape[:-1], -1, 1, 2)
t = t.to(rotation_matrix.dtype)
output = rotation_matrix[..., 0] * t[..., 0] + rotation_matrix[..., 1] * t[..., 1]
if split_pe:
output = output.movedim(-1, -2)
output = output.reshape(input_tensor.shape).type_as(input_tensor)
elif split_pe:
output = comfy.quant_ops.ck.apply_rope_split_half1(input_tensor, rotation_matrix)
else:
output = comfy.quant_ops.ck.apply_rope1(input_tensor, rotation_matrix)
return output.reshape(original_shape)
out = input_tensor * cos_freqs + input_tensor_rot * sin_freqs def apply_rotary_emb_qk(q, k, freqs_cis):
if comfy.model_management.in_training:
return apply_rotary_emb(q, freqs_cis), apply_rotary_emb(k, freqs_cis)
return out rotation_matrix, split_pe = freqs_cis
q_shape = q.shape
def apply_split_rotary_emb(input_tensor, cos, sin): k_shape = k.shape
needs_reshape = False q = q.reshape(q.shape[0], q.shape[1], rotation_matrix.shape[2], -1)
if input_tensor.ndim != 4 and cos.ndim == 4: k = k.reshape(k.shape[0], k.shape[1], rotation_matrix.shape[2], -1)
B, H, T, _ = cos.shape if split_pe:
input_tensor = input_tensor.reshape(B, T, H, -1).swapaxes(1, 2) q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rotation_matrix)
needs_reshape = True else:
split_input = rearrange(input_tensor, "... (d r) -> ... d r", d=2) q, k = comfy.quant_ops.ck.apply_rope(q, k, rotation_matrix)
first_half_input = split_input[..., :1, :] return q.reshape(q_shape), k.reshape(k_shape)
second_half_input = split_input[..., 1:, :]
output = split_input * cos.unsqueeze(-2)
first_half_output = output[..., :1, :]
second_half_output = output[..., 1:, :]
first_half_output.addcmul_(-sin.unsqueeze(-2), second_half_input)
second_half_output.addcmul_(sin.unsqueeze(-2), first_half_input)
output = rearrange(output, "... d r -> ... (d r)")
return output.swapaxes(1, 2).reshape(B, T, -1) if needs_reshape else output
class GuideAttentionMask: class GuideAttentionMask:
@ -461,9 +465,13 @@ class CrossAttention(nn.Module):
q = self.q_norm(q) q = self.q_norm(q)
k = self.k_norm(k) k = self.k_norm(k)
# These norms span all heads, so the per-head RMS+RoPE kernel is not equivalent.
if pe is not None: if pe is not None:
q = apply_rotary_emb(q, pe) if k_pe is None and q.shape == k.shape:
k = apply_rotary_emb(k, pe if k_pe is None else k_pe) q, k = apply_rotary_emb_qk(q, k, pe)
else:
q = apply_rotary_emb(q, pe)
k = apply_rotary_emb(k, pe if k_pe is None else k_pe)
if mask is None: if mask is None:
out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options) out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options)
@ -653,36 +661,23 @@ def generate_freqs(indices, indices_grid, max_pos, use_middle_indices_grid):
) )
return freqs return freqs
def interleaved_freqs_cis(freqs, pad_size): def freqs_cis_matrix(freqs, pad_size, split_mode, num_attention_heads, out_dtype):
cos_freq = freqs.cos().repeat_interleave(2, dim=-1) cos_freq = freqs.cos().to(out_dtype)
sin_freq = freqs.sin().repeat_interleave(2, dim=-1) sin_freq = freqs.sin().to(out_dtype)
if pad_size != 0: if pad_size:
cos_padding = torch.ones_like(cos_freq[:, :, : pad_size]) matrix_pad_size = pad_size if split_mode else pad_size // 2
sin_padding = torch.zeros_like(cos_freq[:, :, : pad_size]) cos_padding = torch.ones_like(cos_freq[:, :, :matrix_pad_size])
cos_freq = torch.cat([cos_padding, cos_freq], dim=-1) sin_padding = torch.zeros_like(sin_freq[:, :, :matrix_pad_size])
sin_freq = torch.cat([sin_padding, sin_freq], dim=-1) cos_freq = torch.cat((cos_padding, cos_freq), dim=-1)
return cos_freq, sin_freq sin_freq = torch.cat((sin_padding, sin_freq), dim=-1)
def split_freqs_cis(freqs, pad_size, num_attention_heads): B, T, _ = cos_freq.shape
cos_freq = freqs.cos() cos_freq = cos_freq.reshape(B, T, num_attention_heads, -1)
sin_freq = freqs.sin() sin_freq = sin_freq.reshape(B, T, num_attention_heads, -1)
rotation_matrix = torch.stack(
if pad_size != 0: (cos_freq, -sin_freq, sin_freq, cos_freq), dim=-1
cos_padding = torch.ones_like(cos_freq[:, :, :pad_size]) )
sin_padding = torch.zeros_like(sin_freq[:, :, :pad_size]) return rotation_matrix.reshape(*rotation_matrix.shape[:-1], 2, 2), split_mode
cos_freq = torch.concatenate([cos_padding, cos_freq], axis=-1)
sin_freq = torch.concatenate([sin_padding, sin_freq], axis=-1)
# Reshape freqs to be compatible with multi-head attention
B , T, half_HD = cos_freq.shape
cos_freq = cos_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads)
sin_freq = sin_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads)
cos_freq = torch.swapaxes(cos_freq, 1, 2) # (B,H,T,D//2)
sin_freq = torch.swapaxes(sin_freq, 1, 2) # (B,H,T,D//2)
return cos_freq, sin_freq
class LTXBaseModel(torch.nn.Module, ABC): class LTXBaseModel(torch.nn.Module, ABC):
""" """
@ -885,12 +880,17 @@ class LTXBaseModel(torch.nn.Module, ABC):
expected_freqs = dim // 2 expected_freqs = dim // 2
current_freqs = freqs.shape[-1] current_freqs = freqs.shape[-1]
pad_size = expected_freqs - current_freqs pad_size = expected_freqs - current_freqs
cos_freq, sin_freq = split_freqs_cis(freqs, pad_size, num_attention_heads)
else: else:
# 2 because of cos and sin by 3 for (t, x, y), 1 for temporal only # 2 because of cos and sin by 3 for (t, x, y), 1 for temporal only
n_elem = 2 * indices_grid.shape[1] n_elem = 2 * indices_grid.shape[1]
cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem) pad_size = dim % n_elem
return cos_freq.to(out_dtype), sin_freq.to(out_dtype), split_mode return freqs_cis_matrix(
freqs,
pad_size,
split_mode,
num_attention_heads,
out_dtype,
)
def _prepare_positional_embeddings(self, pixel_coords, frame_rate, x_dtype): def _prepare_positional_embeddings(self, pixel_coords, frame_rate, x_dtype):
"""Prepare positional embeddings.""" """Prepare positional embeddings."""

View File

@ -6,6 +6,9 @@ import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
import comfy.ldm.common_dit import comfy.ldm.common_dit
import comfy.model_management
import comfy.ops
import comfy.quant_ops
from comfy.ldm.modules.diffusionmodules.mmdit import TimestepEmbedder from comfy.ldm.modules.diffusionmodules.mmdit import TimestepEmbedder
from comfy.ldm.modules.attention import optimized_attention_masked from comfy.ldm.modules.attention import optimized_attention_masked
@ -97,6 +100,7 @@ class JointAttention(nn.Module):
self.n_local_kv_heads = self.n_kv_heads self.n_local_kv_heads = self.n_kv_heads
self.n_rep = self.n_local_heads // self.n_local_kv_heads self.n_rep = self.n_local_heads // self.n_local_kv_heads
self.head_dim = dim // n_heads self.head_dim = dim // n_heads
self.qk_norm = qk_norm
self.qkv = operation_settings.get("operations").Linear( self.qkv = operation_settings.get("operations").Linear(
dim, dim,
@ -151,10 +155,21 @@ class JointAttention(nn.Module):
xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim) xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)
xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim) xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)
xq = self.q_norm(xq) if self.qk_norm and not comfy.model_management.in_training:
xk = self.k_norm(xk) q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, xq, offloadable=True)
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, xk, offloadable=True)
xq, xk = apply_rope(xq, xk, freqs_cis) epsilon = self.q_norm.eps if self.q_norm.eps is not None else torch.finfo(torch.float32).eps
if self.n_local_heads == self.n_local_kv_heads:
xq, xk = comfy.quant_ops.ck.rms_rope(xq, xk, freqs_cis, q_scale, k_scale, epsilon)
else:
xq = comfy.quant_ops.ck.rms_rope1(xq, freqs_cis, q_scale, epsilon)
xk = comfy.quant_ops.ck.rms_rope1(xk, freqs_cis, k_scale, epsilon)
comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream)
comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream)
else:
xq = self.q_norm(xq)
xk = self.k_norm(xk)
xq, xk = apply_rope(xq, xk, freqs_cis)
n_rep = self.n_local_heads // self.n_local_kv_heads n_rep = self.n_local_heads // self.n_local_kv_heads
if n_rep >= 1: if n_rep >= 1:

View File

@ -0,0 +1,186 @@
# Mage-Flow (https://github.com/microsoft/Mage) native-resolution MMDiT (MIT)
# Architecture is a 12-layer variant of the Qwen-Image double-stream block with
# patch_size=1 (no 2x2 packing), unrotated text tokens and a bf16-rounded
# timestep frequency table.
import math
import torch
import torch.nn as nn
from typing import Optional, Tuple
from comfy.ldm.lightricks.model import TimestepEmbedding
from comfy.ldm.flux.layers import EmbedND
from comfy.ldm.qwen_image.model import QwenImageTransformerBlock, LastLayer
import comfy.patcher_extension
class MageTimestepProjEmbeddings(nn.Module):
def __init__(self, embedding_dim, dtype=None, device=None, operations=None):
super().__init__()
self.timestep_embedder = TimestepEmbedding(
in_channels=256, time_embed_dim=embedding_dim,
dtype=dtype, device=device, operations=operations
)
def forward(self, timestep, hidden_states):
half_dim = 128
exponent = -math.log(10000) * torch.arange(half_dim, dtype=torch.float32, device=timestep.device) / half_dim
emb = torch.exp(exponent).to(timestep.dtype)
emb = timestep[:, None].float() * emb[None, :]
emb = 1000.0 * emb
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) # flip_sin_to_cos
return self.timestep_embedder(emb.to(dtype=hidden_states.dtype))
class MageFlowTransformer2DModel(nn.Module):
def __init__(
self,
in_channels: int = 128,
out_channels: Optional[int] = 128,
num_layers: int = 12,
attention_head_dim: int = 128,
num_attention_heads: int = 24,
joint_attention_dim: int = 2560,
axes_dims_rope: Tuple[int, int, int] = (16, 56, 56),
image_model=None,
dtype=None,
device=None,
operations=None,
):
super().__init__()
self.dtype = dtype
self.patch_size = 1
self.in_channels = in_channels
self.out_channels = out_channels or in_channels
self.inner_dim = num_attention_heads * attention_head_dim
self.pe_embedder = EmbedND(dim=attention_head_dim, theta=10000, axes_dim=list(axes_dims_rope))
self.time_text_embed = MageTimestepProjEmbeddings(embedding_dim=self.inner_dim, dtype=dtype, device=device, operations=operations)
self.txt_norm = operations.RMSNorm(joint_attention_dim, eps=1e-6, dtype=dtype, device=device)
self.img_in = operations.Linear(in_channels, self.inner_dim, dtype=dtype, device=device)
self.txt_in = operations.Linear(joint_attention_dim, self.inner_dim, dtype=dtype, device=device)
self.transformer_blocks = nn.ModuleList([
QwenImageTransformerBlock(
dim=self.inner_dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
dtype=dtype,
device=device,
operations=operations
)
for _ in range(num_layers)
])
self.norm_out = LastLayer(self.inner_dim, self.inner_dim, dtype=dtype, device=device, operations=operations)
self.proj_out = operations.Linear(self.inner_dim, self.out_channels, bias=True, dtype=dtype, device=device)
def process_img(self, x, index=0):
# patch_size=1: tokens are raw latent pixels, no 2x2 packing.
bs, c, h, w = x.shape
hidden_states = x.movedim(1, -1).reshape(bs, h * w, c)
img_ids = torch.zeros((h, w, 3), device=x.device)
# Frame axis: positive image index (0 = target, 1..N = reference images).
img_ids[:, :, 0] = index
# Mage scale_rope centering: positions [-ceil(n/2), floor(n/2)), i.e.
# offset by (n - n//2). Differs from Qwen-Image's -(n//2) for odd sizes.
img_ids[:, :, 1] = img_ids[:, :, 1] + torch.arange(h, device=x.device)[:, None] - (h - h // 2)
img_ids[:, :, 2] = img_ids[:, :, 2] + torch.arange(w, device=x.device)[None, :] - (w - w // 2)
return hidden_states, img_ids.reshape(h * w, 3).unsqueeze(0).expand(bs, -1, -1), (h, w)
def forward(self, x, timestep, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
self._forward,
self,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options)
).execute(x, timestep, context, attention_mask, ref_latents, transformer_options, **kwargs)
def _forward(self, x, timestep, context, attention_mask=None, ref_latents=None, transformer_options={}, control=None, **kwargs):
if attention_mask is not None and not torch.is_floating_point(attention_mask):
attention_mask = (attention_mask - 1).to(x.dtype) * torch.finfo(x.dtype).max
hidden_states, img_ids, orig_shape = self.process_img(x)
num_embeds = hidden_states.shape[1]
if ref_latents is not None:
ref_num_tokens = []
index = 0
for ref in ref_latents:
index += 1
kontext, kontext_ids, _ = self.process_img(ref, index=index)
hidden_states = torch.cat([hidden_states, kontext], dim=1)
img_ids = torch.cat([img_ids, kontext_ids], dim=1)
ref_num_tokens.append(kontext.shape[1])
transformer_options = transformer_options.copy()
transformer_options["reference_image_num_tokens"] = ref_num_tokens
# Text tokens are not rotated in Mage-Flow: RoPE at position 0 is the
# identity rotation.
txt_ids = torch.zeros((x.shape[0], context.shape[1], 3), device=x.device)
hidden_states = self.img_in(hidden_states)
context = self.txt_norm(context)
context = self.txt_in(context)
temb = self.time_text_embed(timestep, hidden_states)
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
if "post_input" in patches:
for p in patches["post_input"]:
out = p({"img": hidden_states, "txt": context, "img_ids": img_ids, "txt_ids": txt_ids, "transformer_options": transformer_options})
hidden_states = out["img"]
context = out["txt"]
img_ids = out["img_ids"]
txt_ids = out["txt_ids"]
ids = torch.cat((txt_ids, img_ids), dim=1)
image_rotary_emb = self.pe_embedder(ids).contiguous()
del ids, txt_ids, img_ids
transformer_options["total_blocks"] = len(self.transformer_blocks)
transformer_options["block_type"] = "double"
for i, block in enumerate(self.transformer_blocks):
transformer_options["block_index"] = i
if ("double_block", i) in blocks_replace:
def block_wrap(args):
out = {}
out["txt"], out["img"] = block(hidden_states=args["img"], encoder_hidden_states=args["txt"], encoder_hidden_states_mask=attention_mask, temb=args["vec"], image_rotary_emb=args["pe"], transformer_options=args["transformer_options"])
return out
out = blocks_replace[("double_block", i)]({"img": hidden_states, "txt": context, "vec": temb, "pe": image_rotary_emb, "transformer_options": transformer_options}, {"original_block": block_wrap})
hidden_states = out["img"]
context = out["txt"]
else:
context, hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=context,
encoder_hidden_states_mask=attention_mask,
temb=temb,
image_rotary_emb=image_rotary_emb,
transformer_options=transformer_options,
)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": hidden_states, "txt": context, "x": x, "block_index": i, "transformer_options": transformer_options})
hidden_states = out["img"]
context = out["txt"]
if control is not None: # Controlnet
control_i = control.get("input")
if i < len(control_i):
add = control_i[i]
if add is not None:
hidden_states[:, :add.shape[1]] += add
hidden_states = self.norm_out(hidden_states, temb)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states[:, :num_embeds]
h, w = orig_shape
return hidden_states.reshape(x.shape[0], h, w, self.out_channels).movedim(-1, 1)

477
comfy/ldm/mage_flow/vae.py Normal file
View File

@ -0,0 +1,477 @@
# Mage-VAE (https://github.com/microsoft/Mage) (MIT)
# Symmetric one-step diffusion codec: DConvEncoder (image -> 128ch latent) and
# DConvDenoiser + CoD Decoder (latent -> image). 16x downsample, latents in the
# Flux.2-VAE-anchored space (no patch packing, no BN normalization).
# Both encode and decode are single forward passes at t=0.
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
import comfy.ops
from comfy.ldm.modules.diffusionmodules.model import vae_attention
ops = comfy.ops.disable_weight_init
def nonlinearity(x):
return torch.nn.functional.silu(x)
def Normalize(in_channels):
return ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
def modulate(x, shift, scale):
if x.dim() == 4:
b, c = x.shape[:2]
return x * (1 + scale.view(b, c, 1, 1)) + shift.view(b, c, 1, 1)
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
class LayerNorm2d(ops.LayerNorm):
def __init__(self, num_channels, eps=1e-6, affine=True):
super().__init__(num_channels, eps=eps, elementwise_affine=affine)
def forward(self, x):
x = x.permute(0, 2, 3, 1).contiguous()
x = super().forward(x)
return x.permute(0, 3, 1, 2).contiguous()
class TimestepEmbedder(nn.Module):
"""DConv-style timestep MLP (max_period=10000, freq_size=256)."""
def __init__(self, hidden_size, frequency_embedding_size=256):
super().__init__()
self.mlp = nn.Sequential(
ops.Linear(frequency_embedding_size, hidden_size, bias=True),
nn.SiLU(),
ops.Linear(hidden_size, hidden_size, bias=True),
)
self.frequency_embedding_size = frequency_embedding_size
@staticmethod
def timestep_embedding(t, dim, max_period=10000):
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(0, half, dtype=torch.float32) / half
).to(t.device)
args = t[:, None].float() * freqs[None]
emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
emb = torch.cat([emb, torch.zeros_like(emb[:, :1])], dim=-1)
return emb
def forward(self, t, dtype):
emb = self.timestep_embedding(t, self.frequency_embedding_size)
return self.mlp(emb.to(dtype))
class BottleneckPatchEmbed(nn.Module):
"""Image patch embed concatenated with a per-patch conditioning vector."""
def __init__(self, patch_size=16, in_chans=3, pca_dim=128, embed_dim=384, bias=True):
super().__init__()
self.proj1 = ops.Conv2d(in_chans, pca_dim, kernel_size=patch_size, stride=patch_size, bias=False)
self.proj2 = ops.Conv2d(pca_dim + embed_dim, embed_dim, kernel_size=1, bias=bias)
def forward(self, x, cond):
return self.proj2(torch.cat([self.proj1(x), cond], dim=1))
class DiCoBlock(nn.Module):
"""DConv block with adaLN modulation."""
def __init__(self, hidden_size, mlp_ratio=4.0):
super().__init__()
self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True)
self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
self.ca = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
ops.Conv2d(hidden_size, hidden_size, 1, bias=True),
nn.Sigmoid(),
)
ffn = int(mlp_ratio * hidden_size)
self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True)
self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True)
self.norm1 = LayerNorm2d(hidden_size, affine=False)
self.norm2 = LayerNorm2d(hidden_size, affine=False)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
ops.Linear(hidden_size, 6 * hidden_size, bias=True),
)
def forward(self, inp, c):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1)
x = modulate(self.norm1(inp), shift_msa, scale_msa)
x = F.gelu(self.conv2(self.conv1(x)))
x = x * self.ca(x)
x = self.conv3(x)
x = inp + gate_msa[..., None, None] * x
x = x + gate_mlp[..., None, None] * self.conv5(
F.gelu(self.conv4(modulate(self.norm2(x), shift_mlp, scale_mlp)))
)
return x
class EncoderDiCoBlock(nn.Module):
"""DiCoBlock without adaLN, for the encoder head pathway."""
def __init__(self, hidden_size, mlp_ratio=4.0):
super().__init__()
self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True)
self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
self.ca = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
ops.Conv2d(hidden_size, hidden_size, 1, bias=True),
nn.Sigmoid(),
)
ffn = int(mlp_ratio * hidden_size)
self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True)
self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True)
self.norm1 = LayerNorm2d(hidden_size)
self.norm2 = LayerNorm2d(hidden_size)
def forward(self, inp):
x = self.norm1(inp)
x = F.gelu(self.conv2(self.conv1(x)))
x = x * self.ca(x)
x = self.conv3(x)
x = inp + x
return x + self.conv5(F.gelu(self.conv4(self.norm2(x))))
class NerfEmbedder(nn.Module):
"""Patch-position embedder used by the DConv decoder x-pathway."""
def __init__(self, in_channels, hidden_size_input, max_freqs=8):
super().__init__()
self.max_freqs = max_freqs
self.embedder = nn.Sequential(
ops.Linear(in_channels + max_freqs ** 2, hidden_size_input, bias=True),
)
def fetch_pos(self, patch_size, device, dtype):
pos = torch.linspace(0, 1, patch_size, device=device, dtype=dtype)
pos_y, pos_x = torch.meshgrid(pos, pos, indexing="ij")
pos_x = pos_x.reshape(-1, 1, 1)
pos_y = pos_y.reshape(-1, 1, 1)
freqs = torch.linspace(0, self.max_freqs, self.max_freqs, dtype=dtype, device=device)
fx = freqs[None, :, None]
fy = freqs[None, None, :]
coeffs = (1 + fx * fy) ** -1
dct_x = torch.cos(pos_x * fx * torch.pi)
dct_y = torch.cos(pos_y * fy * torch.pi)
return (dct_x * dct_y * coeffs).view(1, -1, self.max_freqs ** 2)
def forward(self, x):
B, P2, _ = x.shape
ps = int(P2 ** 0.5)
dct = self.fetch_pos(ps, x.device, x.dtype).expand(B, -1, -1)
return self.embedder(torch.cat([x, dct], dim=-1))
class NerfFinalLayer(nn.Module):
def __init__(self, hidden_size, out_channels):
super().__init__()
self.norm = ops.RMSNorm(hidden_size, eps=1e-6)
self.linear = ops.Linear(hidden_size, out_channels, bias=True)
def forward(self, x):
return self.linear(self.norm(x))
class MLPResBlock(nn.Module):
def __init__(self, channels):
super().__init__()
self.in_ln = ops.LayerNorm(channels, eps=1e-6)
self.mlp = nn.Sequential(
ops.Linear(channels, channels, bias=True),
nn.SiLU(),
ops.Linear(channels, channels, bias=True),
)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
ops.Linear(channels, 3 * channels, bias=True),
)
def forward(self, x, y):
shift, scale, gate = self.adaLN_modulation(y).chunk(3, dim=-1)
h = self.in_ln(x) * (1 + scale) + shift
return x + gate * self.mlp(h)
class SimpleMLPAdaLN(nn.Module):
"""Final small MLP that maps NerfEmbedder features to per-patch RGB."""
def __init__(self, in_channels, model_channels, out_channels, z_channels, num_res_blocks, patch_size):
super().__init__()
self.in_channels = in_channels
self.model_channels = model_channels
self.out_channels = out_channels
self.num_res_blocks = num_res_blocks
self.patch_size = patch_size
self.cond_embed = ops.Linear(z_channels, patch_size ** 2 * model_channels)
self.input_proj = ops.Linear(in_channels, model_channels)
self.res_blocks = nn.ModuleList(MLPResBlock(model_channels) for _ in range(num_res_blocks))
def forward(self, x, c):
x = self.input_proj(x)
c = self.cond_embed(c).reshape(c.shape[0], self.patch_size ** 2, -1)
for block in self.res_blocks:
x = block(x, c)
return x
class ResnetBlock(nn.Module):
"""GroupNorm + Conv ResBlock used by the CoD Decoder."""
def __init__(self, *, in_channels, out_channels=None):
super().__init__()
out_channels = out_channels or in_channels
self.in_channels = in_channels
self.out_channels = out_channels
self.norm1 = Normalize(in_channels)
self.conv1 = ops.Conv2d(in_channels, out_channels, 3, padding=1)
self.norm2 = Normalize(out_channels)
self.conv2 = ops.Conv2d(out_channels, out_channels, 3, padding=1)
if in_channels != out_channels:
self.nin_shortcut = ops.Conv2d(in_channels, out_channels, 1)
def forward(self, x):
h = self.conv1(nonlinearity(self.norm1(x)))
h = self.conv2(nonlinearity(self.norm2(h)))
if self.in_channels != self.out_channels:
x = self.nin_shortcut(x)
return x + h
class AttnBlock(nn.Module):
"""Patched (windowed) self-attention used by the CoD Decoder."""
def __init__(self, in_channels, patch_size=32):
super().__init__()
self.in_channels = in_channels
self.patch_size = patch_size
self.norm = Normalize(in_channels)
self.q = ops.Conv2d(in_channels, in_channels, 1)
self.k = ops.Conv2d(in_channels, in_channels, 1)
self.v = ops.Conv2d(in_channels, in_channels, 1)
self.proj_out = ops.Conv2d(in_channels, in_channels, 1)
# VAE attention selection: full-precision backends only (no sage/quantized attention)
self.optimized_attention = vae_attention()
def forward(self, x):
h_ = self.norm(x)
Q = self.q(h_)
K = self.k(h_)
V = self.v(h_)
d = self.patch_size
b, c, H, W = Q.shape
pad_h = (d - H % d) % d
pad_w = (d - W % d) % d
if pad_h or pad_w:
Q = F.pad(Q, (0, pad_w, 0, pad_h), mode="replicate")
K = F.pad(K, (0, pad_w, 0, pad_h), mode="replicate")
V = F.pad(V, (0, pad_w, 0, pad_h), mode="replicate")
_, _, H_pad, W_pad = Q.shape
nph, npw = H_pad // d, W_pad // d
np_ = nph * npw
def to_patches(t):
return (t.reshape(b, c, nph, d, npw, d)
.permute(0, 2, 4, 1, 3, 5)
.reshape(b * np_, c, d * d))
# [b*np, c, d*d]: attention over the d*d spatial positions of each window
Q = to_patches(Q)
K = to_patches(K)
V = to_patches(V)
h_ = self.optimized_attention(Q, K, V)
h_ = h_.reshape(b, nph, npw, c, d, d).permute(0, 3, 1, 4, 2, 5).reshape(b, c, H_pad, W_pad)
if pad_h or pad_w:
h_ = h_[:, :, :H, :W]
return x + self.proj_out(h_)
class CoDDecoder(nn.Module):
"""CoD Decoder: latent -> conditioning features for the denoiser (ds=16, light)."""
def __init__(self, out_ch=384, z_ch=128):
super().__init__()
self.conv_in = ops.Conv2d(z_ch, out_ch, kernel_size=3, stride=1, padding=1)
self.block = nn.Sequential(
ResnetBlock(in_channels=out_ch, out_channels=out_ch),
AttnBlock(out_ch, patch_size=32),
ResnetBlock(in_channels=out_ch, out_channels=out_ch),
AttnBlock(out_ch, patch_size=32),
ResnetBlock(in_channels=out_ch, out_channels=out_ch),
)
self.norm_out = Normalize(out_ch)
self.conv_out = ops.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1)
self.ada = nn.Identity()
def forward(self, z):
h = self.block(self.conv_in(z))
h = self.conv_out(nonlinearity(self.norm_out(h)))
return self.ada(h)
class DConvEncoder(nn.Module):
"""DConvEncoder: image -> packed (mean, logvar) latent."""
def __init__(
self,
z_ch=128,
hidden_size=384,
num_blocks=21,
patch_size=16,
mlp_ratio=4.0,
head_size=768,
num_head_blocks=2,
out_ch_mult=2,
):
super().__init__()
self.z_ch = z_ch
self.patch_size = patch_size
self.patch_cond_embed = ops.Conv2d(3, head_size, kernel_size=patch_size, stride=patch_size, bias=True)
self.head_blocks = nn.ModuleList([
EncoderDiCoBlock(head_size, mlp_ratio=mlp_ratio) for _ in range(num_head_blocks)
])
self.proj_down = ops.Conv2d(head_size, hidden_size, kernel_size=1, bias=True)
self.z_proj = ops.Conv2d(z_ch, hidden_size, kernel_size=1, bias=True)
self.fuse_proj = ops.Conv2d(hidden_size * 2, hidden_size, kernel_size=1, bias=True)
self.t_embedder = TimestepEmbedder(hidden_size)
self.blocks = nn.ModuleList([
DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_blocks)
])
self.norm_out = LayerNorm2d(hidden_size)
self.proj_out = ops.Conv2d(hidden_size, z_ch * out_ch_mult, kernel_size=1, bias=True)
def forward_pred(self, z_t, t, y):
cond = self.patch_cond_embed(y)
for block in self.head_blocks:
cond = block(cond)
cond = self.proj_down(cond)
s = self.fuse_proj(torch.cat([cond, self.z_proj(z_t)], dim=1))
c = self.t_embedder(t.view(-1), y.dtype)
for block in self.blocks:
s = block(s, c)
return self.proj_out(self.norm_out(s))
class YEmbedder(nn.Module):
"""Holds only the CoD decoder (the original Flux2-VAE encoder side is dropped at load)."""
def __init__(self, ch=384, z_ch=128):
super().__init__()
self.decoder = CoDDecoder(out_ch=ch, z_ch=z_ch)
class DConvDenoiser(nn.Module):
"""One-step DConv denoiser: latent (via cond) + zero noise -> reconstructed image."""
def __init__(
self,
patch_size=16,
in_channels=3,
hidden_size=384,
hidden_size_x=32,
mlp_ratio=4.0,
num_blocks=24,
num_cond_blocks=21,
bottleneck_dim=128,
):
super().__init__()
self.in_channels = in_channels
self.patch_size = patch_size
self.hidden_size = hidden_size
self.num_cond_blocks = num_cond_blocks
self.t_embedder = TimestepEmbedder(hidden_size)
self.y_embedder_x = ops.Conv2d(hidden_size, hidden_size_x * patch_size ** 2, 1, 1, 0)
self.x_embedder = NerfEmbedder(in_channels + hidden_size_x, hidden_size_x, max_freqs=8)
self.s_embedder = BottleneckPatchEmbed(patch_size, in_channels, bottleneck_dim, hidden_size, bias=True)
self.blocks = nn.ModuleList([
DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_cond_blocks)
])
self.dec_net = SimpleMLPAdaLN(
in_channels=hidden_size_x,
model_channels=hidden_size_x,
out_channels=in_channels,
z_channels=hidden_size,
num_res_blocks=num_blocks - num_cond_blocks,
patch_size=patch_size,
)
self.final_layer = NerfFinalLayer(hidden_size_x, in_channels)
self.y_embedder = YEmbedder(ch=hidden_size, z_ch=bottleneck_dim)
def forward(self, x, t, cond):
b, _, h, w = x.shape
c = self.t_embedder(t.view(-1), x.dtype)
s = self.s_embedder(x, cond)
for block in self.blocks:
s = block(s, c)
length = s.shape[-2] * s.shape[-1]
s = s.permute(0, 2, 3, 1).reshape(-1, self.hidden_size)
x = torch.nn.functional.unfold(x, kernel_size=self.patch_size, stride=self.patch_size)
x = torch.cat([x, self.y_embedder_x(cond).flatten(2)], dim=1)
x = x.reshape(b, -1, self.patch_size ** 2, length).permute(0, 3, 2, 1).flatten(0, 1)
x = self.x_embedder(x)
x = self.dec_net(x, s)
x = self.final_layer(x)
x = x.transpose(1, 2).reshape(b, length, -1)
return torch.nn.functional.fold(
x.transpose(1, 2).contiguous(), (h, w),
kernel_size=self.patch_size, stride=self.patch_size,
)
class MageVAE(nn.Module):
"""
Encode: DConvEncoder (one-step at t=0) -> posterior mean [B, 128, H/16, W/16]
Decode: DConvDenoiser + CoD Decoder -> image [B, 3, H, W] in [-1, 1]
"""
latent_channels = 128
downsample_factor = 16
def __init__(self):
super().__init__()
self.dconv_encoder = DConvEncoder()
self.decoder_model = DConvDenoiser()
def encode(self, x):
B, _, H, W = x.shape
ps = self.dconv_encoder.patch_size
z_t = torch.zeros(B, self.dconv_encoder.z_ch, H // ps, W // ps, device=x.device, dtype=x.dtype)
t = torch.zeros(B, device=x.device, dtype=x.dtype)
out = self.dconv_encoder.forward_pred(z_t, t, x)
return out[:, : self.latent_channels] # posterior mean (sample_posterior=False)
def decode(self, z):
cond = self.decoder_model.y_embedder.decoder(z)
B = z.shape[0]
H = z.shape[2] * self.downsample_factor
W = z.shape[3] * self.downsample_factor
noise = torch.zeros(B, 3, H, W, device=z.device, dtype=z.dtype)
t = torch.zeros(B, device=z.device, dtype=z.dtype)
return self.decoder_model.forward(noise, t, cond)

View File

@ -552,6 +552,7 @@ class WanModel(torch.nn.Module):
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8] List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
""" """
# embeddings # embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype) x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:] grid_sizes = x.shape[2:]
transformer_options["grid_sizes"] = grid_sizes transformer_options["grid_sizes"] = grid_sizes
@ -564,11 +565,13 @@ class WanModel(torch.nn.Module):
e0 = self.time_projection(e).unflatten(2, (6, self.dim)) e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None full_ref = None
img_offset = 0
if self.ref_conv is not None: if self.ref_conv is not None:
full_ref = kwargs.get("reference_latent", None) full_ref = kwargs.get("reference_latent", None)
if full_ref is not None: if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2) full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1) x = torch.concat((full_ref, x), dim=1)
img_offset = full_ref.shape[1]
# In-context reference (Bernini) # In-context reference (Bernini)
context_latents = kwargs.get("context_latents", None) context_latents = kwargs.get("context_latents", None)
@ -589,6 +592,7 @@ class WanModel(torch.nn.Module):
context_img_len = clip_fea.shape[-2] context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {}) patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {}) blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks) transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double" transformer_options["block_type"] = "double"
@ -604,6 +608,11 @@ class WanModel(torch.nn.Module):
else: else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
# head # head
x = self.head(x, e) x = self.head(x, e)
@ -777,6 +786,7 @@ class VaceWanModel(WanModel):
**kwargs, **kwargs,
): ):
# embeddings # embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype) x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:] grid_sizes = x.shape[2:]
transformer_options["grid_sizes"] = grid_sizes transformer_options["grid_sizes"] = grid_sizes
@ -807,6 +817,7 @@ class VaceWanModel(WanModel):
x_orig = x x_orig = x
patches_replace = transformer_options.get("patches_replace", {}) patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {}) blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks) transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double" transformer_options["block_type"] = "double"
@ -822,6 +833,11 @@ class VaceWanModel(WanModel):
else: else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
ii = self.vace_layers_mapping.get(i, None) ii = self.vace_layers_mapping.get(i, None)
if ii is not None: if ii is not None:
for iii in range(len(c)): for iii in range(len(c)):
@ -887,6 +903,7 @@ class CameraWanModel(WanModel):
**kwargs, **kwargs,
): ):
# embeddings # embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype) x = self.patch_embedding(x.float()).to(x.dtype)
if self.control_adapter is not None and camera_conditions is not None: if self.control_adapter is not None and camera_conditions is not None:
x = x + self.control_adapter(camera_conditions).to(x.dtype) x = x + self.control_adapter(camera_conditions).to(x.dtype)
@ -909,6 +926,7 @@ class CameraWanModel(WanModel):
context_img_len = clip_fea.shape[-2] context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {}) patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {}) blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks) transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double" transformer_options["block_type"] = "double"
@ -924,6 +942,11 @@ class CameraWanModel(WanModel):
else: else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
# head # head
x = self.head(x, e) x = self.head(x, e)
@ -1335,6 +1358,7 @@ class WanModel_S2V(WanModel):
# embeddings # embeddings
bs, _, time, height, width = x.shape bs, _, time, height, width = x.shape
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype) x = self.patch_embedding(x.float()).to(x.dtype)
if control_video is not None: if control_video is not None:
x = x + self.cond_encoder(control_video) x = x + self.cond_encoder(control_video)
@ -1379,6 +1403,7 @@ class WanModel_S2V(WanModel):
context = self.text_embedding(context) context = self.text_embedding(context)
patches_replace = transformer_options.get("patches_replace", {}) patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {}) blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks) transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double" transformer_options["block_type"] = "double"
@ -1393,6 +1418,12 @@ class WanModel_S2V(WanModel):
x = out["img"] x = out["img"]
else: else:
x = block(x, e=e0, freqs=freqs, context=context, transformer_options=transformer_options) x = block(x, e=e0, freqs=freqs, context=context, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
if audio_emb is not None: if audio_emb is not None:
x = self.audio_injector(x, i, audio_emb, audio_emb_global, seq_len) x = self.audio_injector(x, i, audio_emb, audio_emb_global, seq_len)
# head # head
@ -1599,6 +1630,7 @@ class HumoWanModel(WanModel):
bs, _, time, height, width = x.shape bs, _, time, height, width = x.shape
# embeddings # embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype) x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:] grid_sizes = x.shape[2:]
x = x.flatten(2).transpose(1, 2) x = x.flatten(2).transpose(1, 2)
@ -1630,6 +1662,7 @@ class HumoWanModel(WanModel):
audio = None audio = None
patches_replace = transformer_options.get("patches_replace", {}) patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {}) blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks) transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double" transformer_options["block_type"] = "double"
@ -1645,6 +1678,11 @@ class HumoWanModel(WanModel):
else: else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, audio=audio, transformer_options=transformer_options) x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, audio=audio, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
# head # head
x = self.head(x, e) x = self.head(x, e)
@ -1660,8 +1698,14 @@ class SCAILWanModel(WanModel):
def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs): def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs):
x_input = x
img_offset = 0
if reference_latent is not None: if reference_latent is not None:
x = torch.cat((reference_latent, x), dim=2) x = torch.cat((reference_latent, x), dim=2)
img_offset = (reference_latent.shape[2] // self.patch_size[0]) * \
(reference_latent.shape[3] // self.patch_size[1]) * \
(reference_latent.shape[4] // self.patch_size[2])
# embeddings # embeddings
x = self.patch_embedding(x.float()).to(x.dtype) x = self.patch_embedding(x.float()).to(x.dtype)
@ -1697,6 +1741,7 @@ class SCAILWanModel(WanModel):
context_img_len = clip_fea.shape[-2] context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {}) patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {}) blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks) transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double" transformer_options["block_type"] = "double"
@ -1712,6 +1757,11 @@ class SCAILWanModel(WanModel):
else: else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
# head # head
x = self.head(x, e) x = self.head(x, e)

View File

@ -493,6 +493,7 @@ class AnimateWanModel(WanModel):
**kwargs, **kwargs,
): ):
# embeddings # embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype) x = self.patch_embedding(x.float()).to(x.dtype)
x, motion_vec = self.after_patch_embedding(x, pose_latents, face_pixel_values) x, motion_vec = self.after_patch_embedding(x, pose_latents, face_pixel_values)
grid_sizes = x.shape[2:] grid_sizes = x.shape[2:]
@ -505,11 +506,13 @@ class AnimateWanModel(WanModel):
e0 = self.time_projection(e).unflatten(2, (6, self.dim)) e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None full_ref = None
img_offset = 0
if self.ref_conv is not None: if self.ref_conv is not None:
full_ref = kwargs.get("reference_latent", None) full_ref = kwargs.get("reference_latent", None)
if full_ref is not None: if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2) full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1) x = torch.concat((full_ref, x), dim=1)
img_offset = full_ref.shape[1]
# context # context
context = self.text_embedding(context) context = self.text_embedding(context)
@ -522,6 +525,7 @@ class AnimateWanModel(WanModel):
context_img_len = clip_fea.shape[-2] context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {}) patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {}) blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks) transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double" transformer_options["block_type"] = "double"
@ -537,6 +541,11 @@ class AnimateWanModel(WanModel):
else: else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
if i % 5 == 0 and motion_vec is not None: if i % 5 == 0 and motion_vec is not None:
x = x + self.face_adapter.fuser_blocks[i // 5](x, motion_vec) x = x + self.face_adapter.fuser_blocks[i // 5](x, motion_vec)

View File

@ -111,6 +111,7 @@ class WanDancerModel(WanModel):
def forward_orig(self, x, t, context, clip_fea=None, clip_fea_ref=None, freqs=None, audio_embed=None, fps=30, audio_inject_scale=1.0, transformer_options={}, **kwargs): def forward_orig(self, x, t, context, clip_fea=None, clip_fea_ref=None, freqs=None, audio_embed=None, fps=30, audio_inject_scale=1.0, transformer_options={}, **kwargs):
# embeddings # embeddings
x_input = x
if int(fps + 0.5) != 30: if int(fps + 0.5) != 30:
x = self.patch_embedding_global(x.float()).to(x.dtype) x = self.patch_embedding_global(x.float()).to(x.dtype)
else: else:
@ -128,11 +129,13 @@ class WanDancerModel(WanModel):
e0 = self.time_projection(e).unflatten(2, (6, self.dim)) e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None full_ref = None
img_offset = 0
if self.ref_conv is not None: # model has the weight, but this wasn't used in the original pipeline if self.ref_conv is not None: # model has the weight, but this wasn't used in the original pipeline
full_ref = kwargs.get("reference_latent", None) full_ref = kwargs.get("reference_latent", None)
if full_ref is not None: if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2) full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1) x = torch.concat((full_ref, x), dim=1)
img_offset = full_ref.shape[1]
# context # context
context = self.text_embedding(context) context = self.text_embedding(context)
@ -163,6 +166,7 @@ class WanDancerModel(WanModel):
context_img_len += clip_fea_ref.shape[-2] context_img_len += clip_fea_ref.shape[-2]
patches_replace = transformer_options.get("patches_replace", {}) patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {}) blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks) transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double" transformer_options["block_type"] = "double"
@ -177,6 +181,12 @@ class WanDancerModel(WanModel):
x = out["img"] x = out["img"]
else: else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
if audio_emb is not None: if audio_emb is not None:
x = self.music_injector(x, i, audio_emb, audio_emb_global=None, seq_len=seq_len, scale=audio_inject_scale) x = self.music_injector(x, i, audio_emb, audio_emb_global=None, seq_len=seq_len, scale=audio_inject_scale)

149
comfy/ldm/wan/uni3c.py Normal file
View File

@ -0,0 +1,149 @@
# Uni3C controlnet for Wan 2.1: https://github.com/ewrfcas/Uni3C
# Converted from the original diffusers based implementation.
import torch
import torch.nn as nn
from comfy.ldm.flux.layers import EmbedND
from .model import WanSelfAttention
class Uni3CLayerNormZero(nn.Module):
def __init__(
self,
conditioning_dim,
embedding_dim,
eps=1e-5,
device=None, dtype=None, operations=None
):
super().__init__()
self.silu = nn.SiLU()
self.linear = operations.Linear(conditioning_dim, 3 * embedding_dim, device=device, dtype=dtype)
self.norm = operations.LayerNorm(embedding_dim, eps=eps, elementwise_affine=True, device=device, dtype=dtype)
def forward(self, x, temb):
shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1)
x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
return x, gate[:, None, :]
class Uni3CAttentionBlock(nn.Module):
def __init__(
self,
dim,
ffn_dim,
num_heads,
time_embed_dim=5120,
eps=1e-6,
device=None, dtype=None, operations=None
):
super().__init__()
operation_settings = {"operations": operations, "device": device, "dtype": dtype}
self.norm1 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
self.self_attn = WanSelfAttention(dim, num_heads, qk_norm=True, eps=eps, operation_settings=operation_settings)
self.norm2 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
self.ffn = nn.Sequential(
operations.Linear(dim, ffn_dim, device=device, dtype=dtype), nn.GELU(approximate='tanh'),
operations.Linear(ffn_dim, dim, device=device, dtype=dtype))
def forward(self, x, temb, freqs):
norm_x, gate_msa = self.norm1(x, temb)
x = x + gate_msa * self.self_attn(norm_x, freqs)
norm_x, gate_ff = self.norm2(x, temb)
x = x + gate_ff * self.ffn(norm_x)
return x
class MaskCamEmbed(nn.Module):
def __init__(
self,
add_channels=7,
mid_channels=256,
conv_out_dim=5120,
device=None, dtype=None, operations=None
):
super().__init__()
self.mask_padding = [0, 0, 0, 0, 3, 0] # first frame conditioning
self.mask_proj = nn.Sequential(
operations.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8), device=device, dtype=dtype),
operations.GroupNorm(mid_channels // 8, mid_channels, device=device, dtype=dtype),
nn.SiLU())
self.mask_zero_proj = operations.Conv3d(mid_channels, conv_out_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), device=device, dtype=dtype)
def forward(self, add_inputs):
add_padded = torch.nn.functional.pad(add_inputs, self.mask_padding, mode="constant", value=0)
add_embeds = self.mask_proj(add_padded)
add_embeds = self.mask_zero_proj(add_embeds)
add_embeds = add_embeds.flatten(2).transpose(1, 2)
return add_embeds
class WanUni3CControlnet(nn.Module):
def __init__(
self,
in_channels=36,
conv_out_dim=5120,
dim=1024,
ffn_dim=8192,
num_heads=16,
num_layers=20,
time_embed_dim=5120,
out_proj_dim=5120,
add_channels=7,
mid_channels=256,
device=None, dtype=None, operations=None
):
super().__init__()
patch_size = (1, 2, 2)
self.num_layers = num_layers
self.controlnet_patch_embedding = operations.Conv3d(
in_channels, conv_out_dim, kernel_size=patch_size, stride=patch_size, device=device, dtype=torch.float32)
self.controlnet_mask_embedding = MaskCamEmbed(add_channels, mid_channels, conv_out_dim, device=device, dtype=dtype, operations=operations)
if conv_out_dim != dim:
self.proj_in = operations.Linear(conv_out_dim, dim, device=device, dtype=dtype)
else:
self.proj_in = nn.Identity()
self.controlnet_blocks = nn.ModuleList([
Uni3CAttentionBlock(dim, ffn_dim, num_heads, time_embed_dim, device=device, dtype=dtype, operations=operations)
for _ in range(num_layers)])
self.proj_out = nn.ModuleList([
operations.Linear(dim, out_proj_dim, device=device, dtype=dtype)
for _ in range(num_layers)])
head_dim = dim // num_heads
self.rope_embedder = EmbedND(dim=head_dim, theta=10000.0, axes_dim=[head_dim - 4 * (head_dim // 6), 2 * (head_dim // 6), 2 * (head_dim // 6)])
def rope_encode(self, t_len, h_len, w_len, device=None, dtype=None):
img_ids = torch.zeros((t_len, h_len, w_len, 3), device=device, dtype=dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.arange(t_len, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.arange(h_len, device=device, dtype=dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.arange(w_len, device=device, dtype=dtype).reshape(1, 1, -1)
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
freqs = self.rope_embedder(img_ids).movedim(1, 2)
return freqs
def process_input(self, control_input, render_mask=None, camera_embedding=None):
# render_mask/camera_embedding are the checkpoint's extra conditioning path, not wired up yet
hidden = self.controlnet_patch_embedding(control_input.float()).to(control_input.dtype)
t_len, h_len, w_len = hidden.shape[2:]
freqs = self.rope_encode(t_len, h_len, w_len, device=hidden.device, dtype=hidden.dtype)
hidden = hidden.flatten(2).transpose(1, 2)
add_inputs = None
if camera_embedding is not None and render_mask is not None:
add_inputs = torch.cat([render_mask, camera_embedding], dim=1)
elif render_mask is not None:
add_inputs = render_mask
if add_inputs is not None:
hidden = hidden + self.controlnet_mask_embedding(add_inputs.to(hidden.dtype))
hidden = self.proj_in(hidden)
return hidden, freqs
def forward_block(self, block_index, hidden, temb, freqs):
hidden = self.controlnet_blocks[block_index](hidden, temb, freqs)
residual = self.proj_out[block_index](hidden)
return hidden, residual

View File

@ -58,6 +58,8 @@ import comfy.ldm.omnigen.omnigen2
import comfy.ldm.seedvr.model import comfy.ldm.seedvr.model
import comfy.ldm.boogu.model import comfy.ldm.boogu.model
import comfy.ldm.qwen_image.model import comfy.ldm.qwen_image.model
import comfy.ldm.mage_flow.model
import comfy.ldm.joyimage.model
import comfy.ldm.ideogram4.model import comfy.ldm.ideogram4.model
import comfy.ldm.krea2.model import comfy.ldm.krea2.model
import comfy.ldm.kandinsky5.model import comfy.ldm.kandinsky5.model
@ -2023,11 +2025,11 @@ class WAN22_WanDancer(WAN21):
fps = kwargs.get("fps", None) fps = kwargs.get("fps", None)
if fps is not None: if fps is not None:
out['fps'] = comfy.conds.CONDRegular(torch.FloatTensor([fps])) out['fps'] = comfy.conds.CONDConstant(fps)
audio_inject_scale = kwargs.get("audio_inject_scale", None) audio_inject_scale = kwargs.get("audio_inject_scale", None)
if audio_inject_scale is not None: if audio_inject_scale is not None:
out['audio_inject_scale'] = comfy.conds.CONDRegular(torch.FloatTensor([audio_inject_scale])) out['audio_inject_scale'] = comfy.conds.CONDConstant(audio_inject_scale)
return out return out
class Hunyuan3Dv2(BaseModel): class Hunyuan3Dv2(BaseModel):
@ -2226,10 +2228,7 @@ class Omnigen2(BaseModel):
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
ref_latents = kwargs.get("reference_latents", None) ref_latents = kwargs.get("reference_latents", None)
if ref_latents is not None: if ref_latents is not None:
latents = [] out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents])
for lat in ref_latents:
latents.append(self.process_latent_in(lat))
out['ref_latents'] = comfy.conds.CONDList(latents)
return out return out
def extra_conds_shapes(self, **kwargs): def extra_conds_shapes(self, **kwargs):
@ -2245,8 +2244,8 @@ class Boogu(Omnigen2):
self.memory_usage_factor_conds = ("ref_latents",) self.memory_usage_factor_conds = ("ref_latents",)
class QwenImage(BaseModel): class QwenImage(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLUX, device=None): def __init__(self, model_config, model_type=ModelType.FLUX, device=None, unet_model=comfy.ldm.qwen_image.model.QwenImageTransformer2DModel):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.qwen_image.model.QwenImageTransformer2DModel) super().__init__(model_config, model_type, device=device, unet_model=unet_model)
self.memory_usage_factor_conds = ("ref_latents",) self.memory_usage_factor_conds = ("ref_latents",)
def extra_conds(self, **kwargs): def extra_conds(self, **kwargs):
@ -2276,6 +2275,43 @@ class QwenImage(BaseModel):
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16]) out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
return out return out
class MageFlow(QwenImage):
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.mage_flow.model.MageFlowTransformer2DModel)
def process_timestep(self, timestep, **kwargs):
# Mage runs in bf16 and rounds its timestep frequency table to the timestep dtype, keep that on fp32 devices.
return timestep.to(torch.bfloat16)
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, 128, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 128])
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): class Ideogram4(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLOW, device=None): 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) super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.ideogram4.model.Ideogram4Transformer2DModel)
@ -2294,12 +2330,30 @@ class Ideogram4(BaseModel):
class Krea2(BaseModel): class Krea2(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLUX, device=None): def __init__(self, model_config, model_type=ModelType.FLUX, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.krea2.model.SingleStreamDiT) super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.krea2.model.SingleStreamDiT)
self.memory_usage_factor_conds = ("ref_latents",)
def extra_conds(self, **kwargs): def extra_conds(self, **kwargs):
out = super().extra_conds(**kwargs) out = super().extra_conds(**kwargs)
cross_attn = kwargs.get("cross_attn", None) cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None: if cross_attn is not None:
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
ref_latents = kwargs.get("reference_latents", None)
if ref_latents is not None:
latents = []
for lat in ref_latents:
latents.append(self.process_latent_in(lat))
out['ref_latents'] = comfy.conds.CONDList(latents)
ref_latents_method = kwargs.get("reference_latents_method", None)
if ref_latents_method is not None:
out['ref_latents_method'] = comfy.conds.CONDConstant(ref_latents_method)
return out
def extra_conds_shapes(self, **kwargs):
out = {}
ref_latents = kwargs.get("reference_latents", None)
if ref_latents is not None:
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
return out return out
class HunyuanImage21(BaseModel): class HunyuanImage21(BaseModel):

View File

@ -884,6 +884,13 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
"selected_layer_index": selected_layer_index, "selected_layer_index": selected_layer_index,
} }
if '{}txt_norm.weight'.format(key_prefix) in state_dict_keys and '{}proj_out.weight'.format(key_prefix) in state_dict_keys and state_dict['{}txt_norm.weight'.format(key_prefix)].shape[0] == 2560 and state_dict['{}proj_out.weight'.format(key_prefix)].shape[0] == 128: # Mage-Flow (Qwen Image txt_norm/proj_out are 3584/64)
dit_config = {}
dit_config["image_model"] = "mage_flow"
dit_config["in_channels"] = 128
dit_config["num_layers"] = count_blocks(state_dict_keys, '{}transformer_blocks.'.format(key_prefix) + '{}.')
return dit_config
if '{}txt_norm.weight'.format(key_prefix) in state_dict_keys: # Qwen Image if '{}txt_norm.weight'.format(key_prefix) in state_dict_keys: # Qwen Image
dit_config = {} dit_config = {}
dit_config["image_model"] = "qwen_image" dit_config["image_model"] = "qwen_image"
@ -1058,6 +1065,25 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
dit_config["image_model"] = "SAM31" dit_config["image_model"] = "SAM31"
return dit_config 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: if '{}input_blocks.0.0.weight'.format(key_prefix) not in state_dict_keys:
return None return None

View File

@ -473,7 +473,7 @@ except:
SUPPORT_FP8_OPS = args.supports_fp8_compute SUPPORT_FP8_OPS = args.supports_fp8_compute
AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"] AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1035", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]
AMD_ENABLE_MIOPEN_ENV = 'COMFYUI_ENABLE_MIOPEN' AMD_ENABLE_MIOPEN_ENV = 'COMFYUI_ENABLE_MIOPEN'
try: try:

View File

@ -17,6 +17,7 @@ import comfy.ldm.wan.vae
import comfy.ldm.wan.vae2_2 import comfy.ldm.wan.vae2_2
import comfy.ldm.hunyuan3d.vae import comfy.ldm.hunyuan3d.vae
import comfy.ldm.seedvr.vae import comfy.ldm.seedvr.vae
import comfy.ldm.mage_flow.vae
import comfy.ldm.triposplat.vae import comfy.ldm.triposplat.vae
import comfy.ldm.ace.vae.music_dcae_pipeline import comfy.ldm.ace.vae.music_dcae_pipeline
import comfy.ldm.cogvideo.vae import comfy.ldm.cogvideo.vae
@ -60,6 +61,7 @@ import comfy.text_encoders.qwen_image
import comfy.text_encoders.hunyuan_image import comfy.text_encoders.hunyuan_image
import comfy.text_encoders.z_image import comfy.text_encoders.z_image
import comfy.text_encoders.krea2 import comfy.text_encoders.krea2
import comfy.text_encoders.mage_flow
import comfy.text_encoders.ideogram4 import comfy.text_encoders.ideogram4
import comfy.text_encoders.ovis import comfy.text_encoders.ovis
import comfy.text_encoders.kandinsky5 import comfy.text_encoders.kandinsky5
@ -76,6 +78,7 @@ import comfy.text_encoders.gemma4
import comfy.text_encoders.cogvideo import comfy.text_encoders.cogvideo
import comfy.text_encoders.sa3 import comfy.text_encoders.sa3
import comfy.text_encoders.gpt_oss import comfy.text_encoders.gpt_oss
import comfy.text_encoders.joyimage
import comfy.model_patcher import comfy.model_patcher
import comfy.lora import comfy.lora
@ -566,6 +569,17 @@ class VAE:
self.upscale_index_formula = (4, 8, 8) self.upscale_index_formula = (4, 8, 8)
self.process_input = lambda image: image * 2.0 - 1.0 self.process_input = lambda image: image * 2.0 - 1.0
self.crop_input = False self.crop_input = False
elif "student.dconv_encoder.proj_out.weight" in sd: # Mage-VAE (one-step diffusion codec, Flux2-anchored 128ch/16x latents)
sd = comfy.utils.state_dict_prefix_replace(sd, {"student.dconv_encoder.": "dconv_encoder.", "pipeline.": "decoder_model."})
# Drop the unused Flux2-VAE anchor encoder carried in the checkpoint.
sd = {k: v for k, v in sd.items() if not k.startswith("decoder_model.y_embedder.encoder.") and not k.startswith("decoder_model.y_embedder.bottleneck.")}
self.first_stage_model = comfy.ldm.mage_flow.vae.MageVAE()
self.latent_channels = 128
self.downscale_ratio = 16
self.upscale_ratio = 16
self.working_dtypes = [torch.bfloat16, torch.float32]
self.memory_used_encode = lambda shape, dtype: (400 * shape[2] * shape[3]) * model_management.dtype_size(dtype)
self.memory_used_decode = lambda shape, dtype: (1000 * shape[2] * shape[3] * 16 * 16) * model_management.dtype_size(dtype)
elif "decoder.conv_in.weight" in sd: elif "decoder.conv_in.weight" in sd:
if sd['decoder.conv_in.weight'].shape[1] == 64: 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} 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}
@ -1377,6 +1391,8 @@ class CLIPType(Enum):
IDEOGRAM4 = 30 IDEOGRAM4 = 30
BOOGU = 31 BOOGU = 31
KREA2 = 32 KREA2 = 32
JOYIMAGE = 33
MAGE = 34
@ -1432,6 +1448,7 @@ class TEModel(Enum):
GPT_OSS_20B = 33 GPT_OSS_20B = 33
QWEN3VL_4B = 34 QWEN3VL_4B = 34
QWEN3VL_8B = 35 QWEN3VL_8B = 35
GEMMA_4_12B = 36
def detect_te_model(sd): def detect_te_model(sd):
@ -1461,6 +1478,9 @@ def detect_te_model(sd):
if 'model.layers.0.post_feedforward_layernorm.weight' in sd: if 'model.layers.0.post_feedforward_layernorm.weight' in sd:
if 'model.layers.59.self_attn.q_norm.weight' in sd: if 'model.layers.59.self_attn.q_norm.weight' in sd:
return TEModel.GEMMA_4_31B return TEModel.GEMMA_4_31B
# Gemma4 12B Unified: 48 layers, encoder-free; global layers drop v_proj (attention_k_eq_v).
if 'model.layers.47.self_attn.q_norm.weight' in sd and 'model.layers.5.self_attn.v_proj.weight' not in sd:
return TEModel.GEMMA_4_12B
if 'model.layers.41.self_attn.q_norm.weight' in sd and 'model.layers.47.self_attn.q_norm.weight' not in sd: if 'model.layers.41.self_attn.q_norm.weight' in sd and 'model.layers.47.self_attn.q_norm.weight' not in sd:
return TEModel.GEMMA_4_E4B return TEModel.GEMMA_4_E4B
if 'model.layers.34.self_attn.q_norm.weight' in sd and 'model.layers.41.self_attn.q_norm.weight' not in sd: if 'model.layers.34.self_attn.q_norm.weight' in sd and 'model.layers.41.self_attn.q_norm.weight' not in sd:
@ -1616,10 +1636,11 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
clip_target.clip = comfy.text_encoders.sa3.SAT5GemmaModel clip_target.clip = comfy.text_encoders.sa3.SAT5GemmaModel
clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer
tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None) tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None)
elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B): elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B, TEModel.GEMMA_4_12B):
variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B,
TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B,
TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B}[te_model] TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B,
TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model]
clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant) clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant)
clip_target.tokenizer = variant.tokenizer clip_target.tokenizer = variant.tokenizer
tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None) tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None)
@ -1706,6 +1727,14 @@ 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_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.clip = comfy.text_encoders.krea2.te(**llama_detect(clip_data))
clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer
elif clip_type == CLIPType.MAGE and te_model == TEModel.QWEN3VL_4B: # Mage-Flow: full Qwen3-VL-4B, last hidden state, Qwen-Image-style templates.
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.mage_flow.te(**llama_detect(clip_data))
clip_target.tokenizer = comfy.text_encoders.mage_flow.MageFlowTokenizer
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. 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" 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) clip_target.clip = comfy.text_encoders.flux.klein_te(**llama_detect(clip_data), model_type=klein_model_type)

View File

@ -27,6 +27,8 @@ import comfy.text_encoders.z_image
import comfy.text_encoders.ideogram4 import comfy.text_encoders.ideogram4
import comfy.text_encoders.boogu import comfy.text_encoders.boogu
import comfy.text_encoders.krea2 import comfy.text_encoders.krea2
import comfy.text_encoders.mage_flow
import comfy.text_encoders.joyimage
import comfy.text_encoders.anima import comfy.text_encoders.anima
import comfy.text_encoders.ace15 import comfy.text_encoders.ace15
import comfy.text_encoders.longcat_image import comfy.text_encoders.longcat_image
@ -1882,6 +1884,35 @@ class Krea2(supported_models_base.BASE):
hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref)) hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref))
return supported_models_base.ClipTarget(comfy.text_encoders.krea2.Krea2Tokenizer, comfy.text_encoders.krea2.te(**hunyuan_detect)) return supported_models_base.ClipTarget(comfy.text_encoders.krea2.Krea2Tokenizer, comfy.text_encoders.krea2.te(**hunyuan_detect))
class MageFlow(supported_models_base.BASE):
unet_config = {
"image_model": "mage_flow",
}
sampling_settings = {
"multiplier": 1.0,
"shift": 6.0,
}
memory_usage_factor = 6.5
unet_extra_config = {}
latent_format = latent_formats.Flux2
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):
out = model_base.MageFlow(self, device=device)
return out
def clip_target(self, state_dict={}):
pref = self.text_encoder_key_prefix[0]
hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref))
return supported_models_base.ClipTarget(comfy.text_encoders.mage_flow.MageFlowTokenizer, comfy.text_encoders.mage_flow.te(**hunyuan_detect))
class QwenImage(supported_models_base.BASE): class QwenImage(supported_models_base.BASE):
unet_config = { unet_config = {
"image_model": "qwen_image", "image_model": "qwen_image",
@ -1911,6 +1942,38 @@ class QwenImage(supported_models_base.BASE):
hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen25_7b.transformer.".format(pref)) 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)) 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): class HunyuanImage21(HunyuanVideo):
unet_config = { unet_config = {
"image_model": "hunyuan_video", "image_model": "hunyuan_video",
@ -2388,7 +2451,9 @@ models = [
ACEStep15, ACEStep15,
Omnigen2, Omnigen2,
Boogu, Boogu,
MageFlow,
QwenImage, QwenImage,
JoyImage,
Ideogram4, Ideogram4,
Krea2, Krea2,
Flux2, Flux2,

View File

@ -1,11 +1,15 @@
import torch import torch
import torch.nn as nn import torch.nn as nn
import torchaudio.functional as AF
import torchvision.transforms.functional as TVF
import numpy as np import numpy as np
from tokenizers import Tokenizer
from dataclasses import dataclass from dataclasses import dataclass
import math import math
from comfy import sd1_clip from comfy import sd1_clip
import comfy.model_management import comfy.model_management
import comfy.ops
from comfy.ldm.modules.attention import optimized_attention_for_device from comfy.ldm.modules.attention import optimized_attention_for_device
from comfy.rmsnorm import rms_norm from comfy.rmsnorm import rms_norm
from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding
@ -21,6 +25,10 @@ GEMMA4_VISION_CONFIG = {"hidden_size": 768, "image_size": 896, "intermediate_siz
GEMMA4_VISION_31B_CONFIG = {"hidden_size": 1152, "image_size": 896, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 16, "head_dim": 72, "rms_norm_eps": 1e-6, "position_embedding_size": 10240, "pooling_kernel_size": 3} GEMMA4_VISION_31B_CONFIG = {"hidden_size": 1152, "image_size": 896, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 16, "head_dim": 72, "rms_norm_eps": 1e-6, "position_embedding_size": 10240, "pooling_kernel_size": 3}
GEMMA4_AUDIO_CONFIG = {"hidden_size": 1024, "num_hidden_layers": 12, "num_attention_heads": 8, "intermediate_size": 4096, "conv_kernel_size": 5, "attention_chunk_size": 12, "attention_context_left": 13, "attention_context_right": 0, "attention_logit_cap": 50.0, "output_proj_dims": 1536, "rms_norm_eps": 1e-6, "residual_weight": 0.5} GEMMA4_AUDIO_CONFIG = {"hidden_size": 1024, "num_hidden_layers": 12, "num_attention_heads": 8, "intermediate_size": 4096, "conv_kernel_size": 5, "attention_chunk_size": 12, "attention_context_left": 13, "attention_context_right": 0, "attention_logit_cap": 50.0, "output_proj_dims": 1536, "rms_norm_eps": 1e-6, "residual_weight": 0.5}
# Encoder-free (gemma4_unified) multimodal embedders: raw patches/waveform projected directly into LM space.
GEMMA4_UNIFIED_VISION_CONFIG = {"model_patch_size": 48, "patch_size": 16, "pooling_kernel_size": 3, "mm_embed_dim": 3840, "mm_posemb_size": 1120, "output_proj_dims": 3840, "rms_norm_eps": 1e-6}
GEMMA4_UNIFIED_AUDIO_CONFIG = {"audio_samples_per_token": 640, "output_proj_dims": 640, "rms_norm_eps": 1e-6}
@dataclass @dataclass
class Gemma4Config: class Gemma4Config:
vocab_size: int = 262144 vocab_size: int = 262144
@ -35,6 +43,9 @@ class Gemma4Config:
transformer_type: str = "gemma4" transformer_type: str = "gemma4"
head_dim = 256 head_dim = 256
global_head_dim = 512 global_head_dim = 512
num_global_key_value_heads = None
attention_k_eq_v = False
vision_bidirectional = False
rms_norm_add = False rms_norm_add = False
mlp_activation = "gelu_pytorch_tanh" mlp_activation = "gelu_pytorch_tanh"
qkv_bias = False qkv_bias = False
@ -51,6 +62,7 @@ class Gemma4Config:
num_kv_shared_layers: int = 18 num_kv_shared_layers: int = 18
use_double_wide_mlp: bool = False use_double_wide_mlp: bool = False
stop_tokens = [1, 50, 106] stop_tokens = [1, 50, 106]
suppress_tokens = []
vision_config = GEMMA4_VISION_CONFIG vision_config = GEMMA4_VISION_CONFIG
audio_config = GEMMA4_AUDIO_CONFIG audio_config = GEMMA4_AUDIO_CONFIG
mm_tokens_per_image = 280 mm_tokens_per_image = 280
@ -72,12 +84,30 @@ class Gemma4_31B_Config(Gemma4Config):
num_hidden_layers: int = 60 num_hidden_layers: int = 60
num_attention_heads: int = 32 num_attention_heads: int = 32
num_key_value_heads: int = 16 num_key_value_heads: int = 16
vision_bidirectional = True
sliding_attention = [1024, 1024, 1024, 1024, 1024, False] sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
hidden_size_per_layer_input: int = 0 hidden_size_per_layer_input: int = 0
num_kv_shared_layers: int = 0 num_kv_shared_layers: int = 0
audio_config = None audio_config = None
vision_config = GEMMA4_VISION_31B_CONFIG vision_config = GEMMA4_VISION_31B_CONFIG
@dataclass
class Gemma4_12B_Config(Gemma4Config):
hidden_size: int = 3840
intermediate_size: int = 15360
num_hidden_layers: int = 48
num_attention_heads: int = 16
num_key_value_heads: int = 8
num_global_key_value_heads = 1
attention_k_eq_v = True
vision_bidirectional = True
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
hidden_size_per_layer_input: int = 0
num_kv_shared_layers: int = 0
audio_config = GEMMA4_UNIFIED_AUDIO_CONFIG
vision_config = GEMMA4_UNIFIED_VISION_CONFIG
suppress_tokens = [258883, 258882]
# unfused RoPE as addcmul_ RoPE diverges from reference code # unfused RoPE as addcmul_ RoPE diverges from reference code
def _apply_rotary_pos_emb(x, freqs_cis): def _apply_rotary_pos_emb(x, freqs_cis):
@ -89,17 +119,18 @@ def _apply_rotary_pos_emb(x, freqs_cis):
return out return out
class Gemma4Attention(nn.Module): class Gemma4Attention(nn.Module):
def __init__(self, config, head_dim, device=None, dtype=None, ops=None): def __init__(self, config, head_dim, num_kv_heads=None, k_eq_v=False, device=None, dtype=None, ops=None):
super().__init__() super().__init__()
self.num_heads = config.num_attention_heads self.num_heads = config.num_attention_heads
self.num_kv_heads = config.num_key_value_heads self.num_kv_heads = num_kv_heads if num_kv_heads is not None else config.num_key_value_heads
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.head_dim = head_dim self.head_dim = head_dim
self.inner_size = self.num_heads * head_dim self.inner_size = self.num_heads * head_dim
self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype) self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype)
self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype) self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype) # k_eq_v: V reuses the K projection (no separate v_proj weight)
self.v_proj = None if k_eq_v else ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype) self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype)
self.q_norm = None self.q_norm = None
@ -133,7 +164,10 @@ class Gemma4Attention(nn.Module):
shareable_kv = None shareable_kv = None
else: else:
xk = self.k_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim) xk = self.k_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim) if self.v_proj is not None:
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
else:
xv = xk # k_eq_v: V is the raw K projection (before k_norm/RoPE)
if self.k_norm is not None: if self.k_norm is not None:
xk = self.k_norm(xk) xk = self.k_norm(xk)
xv = rms_norm(xv) xv = rms_norm(xv)
@ -186,7 +220,10 @@ class TransformerBlockGemma4(nn.Module):
head_dim = config.head_dim if self.sliding_attention else config.global_head_dim head_dim = config.head_dim if self.sliding_attention else config.global_head_dim
self.self_attn = Gemma4Attention(config, head_dim=head_dim, device=device, dtype=dtype, ops=ops) # k_eq_v only on global layers, which then use num_global_key_value_heads
k_eq_v = config.attention_k_eq_v and not self.sliding_attention
num_kv_heads = config.num_global_key_value_heads if k_eq_v else config.num_key_value_heads
self.self_attn = Gemma4Attention(config, head_dim=head_dim, num_kv_heads=num_kv_heads, k_eq_v=k_eq_v, device=device, dtype=dtype, ops=ops)
num_kv_shared = config.num_kv_shared_layers num_kv_shared = config.num_kv_shared_layers
first_kv_shared = config.num_hidden_layers - num_kv_shared first_kv_shared = config.num_hidden_layers - num_kv_shared
@ -203,9 +240,9 @@ class TransformerBlockGemma4(nn.Module):
self.per_layer_input_gate = ops.Linear(config.hidden_size, self.hidden_size_per_layer_input, bias=False, device=device, dtype=dtype) self.per_layer_input_gate = ops.Linear(config.hidden_size, self.hidden_size_per_layer_input, bias=False, device=device, dtype=dtype)
self.per_layer_projection = ops.Linear(self.hidden_size_per_layer_input, config.hidden_size, bias=False, device=device, dtype=dtype) self.per_layer_projection = ops.Linear(self.hidden_size_per_layer_input, config.hidden_size, bias=False, device=device, dtype=dtype)
self.post_per_layer_input_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype) self.post_per_layer_input_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype)
self.register_buffer("layer_scalar", torch.ones(1, device=device, dtype=dtype))
else: # layer_scalar exists on every gemma4 variant, independent of per-layer input
self.layer_scalar = None self.register_buffer("layer_scalar", torch.empty(1, device=device, dtype=dtype))
def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, per_layer_input=None, shared_kv=None): def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, per_layer_input=None, shared_kv=None):
sliding_window = None sliding_window = None
@ -244,8 +281,7 @@ class TransformerBlockGemma4(nn.Module):
x = self.post_per_layer_input_norm(x) x = self.post_per_layer_input_norm(x)
x = residual + x x = residual + x
if self.layer_scalar is not None: x = x * comfy.ops.cast_to_input(self.layer_scalar, x)
x = x * self.layer_scalar
return x, present_key_value, shareable_kv return x, present_key_value, shareable_kv
@ -334,6 +370,19 @@ class Gemma4Transformer(nn.Module):
causal_mask.masked_fill_(torch.ones_like(causal_mask, dtype=torch.bool).triu_(1), min_val) causal_mask.masked_fill_(torch.ones_like(causal_mask, dtype=torch.bool).triu_(1), min_val)
mask = mask + causal_mask if mask is not None else causal_mask mask = mask + causal_mask if mask is not None else causal_mask
# Bidirectional attention within each image soft-token block (prefill only; text/audio stay causal).
if self.config.vision_bidirectional and past_len == 0 and embeds_info:
block_ids = torch.full((seq_len,), -1, dtype=torch.long, device=x.device)
group = 0
for info in embeds_info:
if info.get("type") == "image":
start = info["index"]
block_ids[start:start + info["size"]] = group
group += 1
if group > 0:
same_block = (block_ids[:, None] == block_ids[None, :]) & (block_ids[:, None] >= 0)
mask = mask.masked_fill(same_block, 0.0)
# Per-layer inputs # Per-layer inputs
per_layer_inputs = None per_layer_inputs = None
if self.hidden_size_per_layer_input: if self.hidden_size_per_layer_input:
@ -354,8 +403,24 @@ class Gemma4Transformer(nn.Module):
shared_global_kv = None # KV from last non-shared global layer shared_global_kv = None # KV from last non-shared global layer
intermediate = None intermediate = None
all_intermediate = None
only_layers = None
if intermediate_output is not None:
if isinstance(intermediate_output, list):
all_intermediate = []
only_layers = {len(self.layers) + layer if layer < 0 else layer for layer in intermediate_output}
elif intermediate_output == "all":
all_intermediate = []
intermediate_output = None
elif intermediate_output < 0:
intermediate_output = len(self.layers) + intermediate_output
next_key_values = [] next_key_values = []
for i, layer in enumerate(self.layers): for i, layer in enumerate(self.layers):
if all_intermediate is not None:
if only_layers is None or (i in only_layers):
all_intermediate.append(x.unsqueeze(1).clone())
past_kv = past_key_values[i] if past_key_values is not None and len(past_key_values) > 0 else None past_kv = past_key_values[i] if past_key_values is not None and len(past_key_values) > 0 else None
layer_kwargs = {} layer_kwargs = {}
@ -385,7 +450,18 @@ class Gemma4Transformer(nn.Module):
if self.norm is not None: if self.norm is not None:
x = self.norm(x) x = self.norm(x)
if len(next_key_values) > 0: if all_intermediate is not None:
if only_layers is None or (len(self.layers) in only_layers):
all_intermediate.append(x.unsqueeze(1).clone())
if len(all_intermediate) > 0:
intermediate = torch.cat(all_intermediate, dim=1)
if intermediate is not None and final_layer_norm_intermediate and self.norm is not None:
intermediate = self.norm(intermediate)
# Only hand back the KV cache when caching was actually requested; SDClipModel reads
# outputs[2] as the pooled output.
if past_key_values is not None and len(next_key_values) > 0:
return x, intermediate, next_key_values return x, intermediate, next_key_values
return x, intermediate return x, intermediate
@ -404,6 +480,8 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module):
cap = self.model.config.final_logit_softcapping cap = self.model.config.final_logit_softcapping
if cap: if cap:
logits = cap * torch.tanh(logits / cap) logits = cap * torch.tanh(logits / cap)
if self.model.config.suppress_tokens:
logits[..., self.model.config.suppress_tokens] = torch.finfo(logits.dtype).min
return logits return logits
def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): def init_kv_cache(self, batch, max_cache_len, device, execution_dtype):
@ -441,6 +519,28 @@ class Gemma4AudioMixin:
return None, None return None, None
class Gemma4UnifiedBase(Gemma4Base):
"""Encoder-free multimodal Gemma4 (gemma4_unified, e.g. 12B): raw image patches and audio frames projected directly into LM space."""
def _init_model(self, config, dtype, device, operations):
self.num_layers = config.num_hidden_layers
self.model = Gemma4Transformer(config, device=device, dtype=dtype, ops=operations)
self.dtype = dtype
self.vision_model = Gemma4UnifiedVisionEmbedder(config.vision_config, device=device, dtype=dtype, ops=operations)
self.multi_modal_projector = Gemma4RMSNormProjector(config.vision_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
self.audio_projector = Gemma4RMSNormProjector(config.audio_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
def preprocess_embed(self, embed, device):
if embed["type"] == "image":
pixels = embed.pop("data").movedim(-1, 1).to(device, dtype=self.dtype) # [B, H, W, C] -> [B, C, H, W], [0,1]
patches, positions = self.vision_model.patchify(pixels)
vision_out = self.vision_model(patches, positions)
return self.multi_modal_projector(vision_out), None
if embed["type"] == "audio":
audio = embed.pop("data").to(device, dtype=self.dtype) # [1, T, audio_samples_per_token]
return self.audio_projector(audio), None
return None, None
# Vision Encoder # Vision Encoder
def _compute_vision_2d_rope(head_dim, pixel_position_ids, theta=100.0, device=None): def _compute_vision_2d_rope(head_dim, pixel_position_ids, theta=100.0, device=None):
@ -713,6 +813,73 @@ class Gemma4MultiModalProjector(Gemma4RMSNormProjector):
super().__init__(config.vision_config["hidden_size"], config.hidden_size, dtype=dtype, device=device, ops=ops) super().__init__(config.vision_config["hidden_size"], config.hidden_size, dtype=dtype, device=device, ops=ops)
# Encoder-free vision (gemma4_unified): raw merged pixel patches projected directly into LM space.
def _patches_merge(patches, positions_xy, length):
patch_size = math.isqrt(patches.shape[-1] // 3)
k = math.isqrt(patches.shape[-2] // length)
batch = patches.shape[:-2]
max_x = positions_xy[..., 0].max(dim=-1, keepdim=True)[0] + 1
kidx = torch.div(positions_xy, k, rounding_mode="floor")
rem = torch.remainder(positions_xy, k)
order = rem[..., 0] + rem[..., 1] * k + k * k * kidx[..., 0] + k * max_x * kidx[..., 1]
perm = order.long().argsort(dim=-1)
merged = patches.gather(-2, perm.unsqueeze(-1).expand_as(patches))
merged = merged.reshape(*batch, length, k, k, patch_size, patch_size, 3)
merged = merged.permute(*range(len(batch)), -6, -5, -3, -4, -2, -1).reshape(*batch, length, (k * patch_size) ** 2 * 3)
pos = positions_xy.gather(-2, perm.unsqueeze(-1).expand_as(positions_xy))
pad = (positions_xy == -1).all(dim=-1, keepdim=True)
pos = torch.where(pad, positions_xy, pos).reshape(*batch, length, k * k, 2)
pos = torch.div(pos, k, rounding_mode="floor").min(dim=-2)[0]
return merged, pos
class Gemma4UnifiedVisionEmbedder(nn.Module):
"""Encoder-free patch embedder (LN -> Dense -> LN -> +2D posemb -> LN); projection to text space is the separate multi_modal_projector."""
def __init__(self, config, device=None, dtype=None, ops=None):
super().__init__()
self.patch_size = config["patch_size"]
self.pooling_kernel_size = config["pooling_kernel_size"]
patch_dim = config["model_patch_size"] ** 2 * 3
mm_embed_dim = config["mm_embed_dim"]
self.patch_ln1 = ops.LayerNorm(patch_dim, device=device, dtype=dtype)
self.patch_dense = ops.Linear(patch_dim, mm_embed_dim, device=device, dtype=dtype)
self.patch_ln2 = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
self.pos_embedding = nn.Parameter(torch.empty(config["mm_posemb_size"], 2, mm_embed_dim, device=device, dtype=dtype))
self.pos_norm = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
def patchify(self, pixels):
"""pixels: [B, C, H, W] in [0,1] -> merged patches [B, N, 6912], positions [B, N, 2]."""
ps, k = self.patch_size, self.pooling_kernel_size
out_patches, out_positions = [], []
for img in pixels:
ph, pw = img.shape[-2] // ps, img.shape[-1] // ps
teacher = img.reshape(img.shape[0], ph, ps, pw, ps).permute(1, 3, 2, 4, 0).reshape(ph * pw, -1)
grid = torch.meshgrid(torch.arange(pw, device=img.device), torch.arange(ph, device=img.device), indexing="xy")
tpos = torch.stack(grid, dim=-1).reshape(teacher.shape[0], 2)
n_model = teacher.shape[0] // (k * k)
mp, mpos = _patches_merge(teacher.unsqueeze(0), tpos.unsqueeze(0), n_model)
out_patches.append(mp.squeeze(0))
out_positions.append(mpos.squeeze(0))
return torch.stack(out_patches), torch.stack(out_positions)
def forward(self, pixel_values, image_position_ids):
x = self.patch_ln1(pixel_values)
x = self.patch_dense(x)
x = self.patch_ln2(x)
clamped = image_position_ids.clamp(min=0).long()
valid = (image_position_ids != -1).to(x.dtype).unsqueeze(-1)
axes = torch.arange(2, device=image_position_ids.device)
pos = comfy.model_management.cast_to_device(self.pos_embedding, x.device, x.dtype)
pos_embs = (pos[clamped, axes] * valid).sum(-2)
x = x + pos_embs
return self.pos_norm(x)
# Audio Encoder # Audio Encoder
class Gemma4AudioConvSubsampler(nn.Module): class Gemma4AudioConvSubsampler(nn.Module):
@ -990,6 +1157,30 @@ class Gemma4AudioProjector(Gemma4RMSNormProjector):
# Tokenizer and Wrappers # Tokenizer and Wrappers
def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, pooling_kernel_size):
target_px = max_patches * patch_size ** 2
factor = math.sqrt(target_px / (height * width))
side_mult = pooling_kernel_size * patch_size
target_height = math.floor(factor * height / side_mult) * side_mult
target_width = math.floor(factor * width / side_mult) * side_mult
if target_height == 0 and target_width == 0:
raise ValueError(f"Attempting to resize to a 0 x 0 image. Resized height should be divisible by {side_mult}.")
max_side_length = (max_patches // pooling_kernel_size ** 2) * side_mult
if target_height == 0:
target_height = side_mult
target_width = min(math.floor(width / height) * side_mult, max_side_length)
elif target_width == 0:
target_width = side_mult
target_height = min(math.floor(height / width) * side_mult, max_side_length)
if target_height * target_width > target_px:
raise ValueError(f"Resizing [{height}x{width}] to [{target_height}x{target_width}] exceeds the patch budget.")
return target_height, target_width
class Gemma4_Tokenizer(): class Gemma4_Tokenizer():
tokenizer_json_data = None tokenizer_json_data = None
@ -998,25 +1189,35 @@ class Gemma4_Tokenizer():
return {"tokenizer_json": self.tokenizer_json_data} return {"tokenizer_json": self.tokenizer_json_data}
return {} return {}
def _extract_mel_spectrogram(self, waveform, sample_rate): def _audio_token_count(self, num_samples):
"""Extract 128-bin log mel spectrogram. # Default (E2B/E4B): mel frames after two stride-2 conv subsamples.
Uses numpy for FFT/matmul/log to produce bit-identical results with reference code. _fl = 320 # int(round(16000 * 20.0 / 1000.0))
""" _hl = 160 # int(round(16000 * 10.0 / 1000.0))
# Mix to mono first, then resample to 16kHz _nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
_t = _nmel
for _ in range(2):
_t = (_t + 2 - 3) // 2 + 1
return min(_t, 750)
@staticmethod
def _resample_16k(waveform, sample_rate):
"""Mix to mono and resample to 16kHz. Kaiser params reproduce the reference (transformers
load_audio -> librosa/soxr_hq) to ~1e-12 MSE using only torchaudio."""
if waveform.dim() > 1 and waveform.shape[0] > 1: if waveform.dim() > 1 and waveform.shape[0] > 1:
waveform = waveform.mean(dim=0, keepdim=True) waveform = waveform.mean(dim=0, keepdim=True)
if waveform.dim() == 1: if waveform.dim() == 1:
waveform = waveform.unsqueeze(0) waveform = waveform.unsqueeze(0)
audio = waveform.squeeze(0).float().numpy() audio = waveform.float()
if sample_rate != 16000: if sample_rate != 16000:
# Use scipy's resample_poly with a high-quality FIR filter to get as close as possible to librosa's resampling (while still not full match) audio = AF.resample(audio, sample_rate, 16000, resampling_method="sinc_interp_kaiser",
from scipy.signal import resample_poly, firwin lowpass_filter_width=121, rolloff=0.9568384289091556, beta=21.01531462440614)
from math import gcd return audio.squeeze(0).contiguous()
g = gcd(sample_rate, 16000)
up, down = 16000 // g, sample_rate // g def _extract_audio_features(self, waveform, sample_rate):
L = max(up, down) """Default (E2B/E4B): 128-bin log mel spectrogram for the conformer audio encoder.
h = firwin(160 * L + 1, 0.96 / L, window=('kaiser', 6.5)) Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
audio = resample_poly(audio, up, down, window=h).astype(np.float32) """
audio = self._resample_16k(waveform, sample_rate).numpy()
n = len(audio) n = len(audio)
# Pad to multiple of 128, build sample-level mask # Pad to multiple of 128, build sample-level mask
@ -1064,8 +1265,8 @@ class Gemma4_Tokenizer():
if audio is not None: if audio is not None:
waveform = audio["waveform"].squeeze(0) if hasattr(audio, "__getitem__") else audio waveform = audio["waveform"].squeeze(0) if hasattr(audio, "__getitem__") else audio
sample_rate = audio.get("sample_rate", 16000) if hasattr(audio, "get") else 16000 sample_rate = audio.get("sample_rate", 16000) if hasattr(audio, "get") else 16000
mel, mel_mask = self._extract_mel_spectrogram(waveform, sample_rate) feat, feat_mask = self._extract_audio_features(waveform, sample_rate)
audio_features = [(mel.unsqueeze(0), mel_mask.unsqueeze(0))] # ([1, T, 128], [1, T]) audio_features = [(feat.unsqueeze(0), feat_mask.unsqueeze(0))] # ([1, T, D], [1, T])
# Process image/video frames # Process image/video frames
is_video = video is not None is_video = video is not None
@ -1090,13 +1291,8 @@ class Gemma4_Tokenizer():
pooling_k = 3 pooling_k = 3
max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280) max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280)
max_patches = max_soft_tokens * pooling_k * pooling_k max_patches = max_soft_tokens * pooling_k * pooling_k
target_px = max_patches * patch_size * patch_size target_h, target_w = _get_aspect_ratio_preserving_size(h, w, patch_size, max_patches, pooling_k)
factor = (target_px / (h * w)) ** 0.5
side_mult = pooling_k * patch_size
target_h = max(int(factor * h // side_mult) * side_mult, side_mult)
target_w = max(int(factor * w // side_mult) * side_mult, side_mult)
import torchvision.transforms.functional as TVF
for i in range(num_frames): for i in range(num_frames):
# rescaling to match reference code # rescaling to match reference code
s = (samples[i].clamp(0, 1) * 255).to(torch.uint8) # [C, H, W] uint8 s = (samples[i].clamp(0, 1) * 255).to(torch.uint8) # [C, H, W] uint8
@ -1115,7 +1311,7 @@ class Gemma4_Tokenizer():
llama_text = llama_template.format(text) llama_text = llama_template.format(text)
else: else:
# Build template from modalities present # Build template from modalities present
system = "<|turn>system\n<|think|><turn|>\n" if thinking else "" system = "<|turn>system\n<|think|>\n<turn|>\n" if thinking else ""
media = "" media = ""
if len(images) > 0: if len(images) > 0:
if is_video: if is_video:
@ -1135,15 +1331,11 @@ class Gemma4_Tokenizer():
if len(audio_features) > 0: if len(audio_features) > 0:
# Compute audio token count (always at 16kHz) # Compute audio token count (always at 16kHz)
num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1] num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1]
_fl = 320 # int(round(16000 * 20.0 / 1000.0)) n_audio_tokens = self._audio_token_count(num_samples)
_hl = 160 # int(round(16000 * 10.0 / 1000.0))
_nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
_t = _nmel
for _ in range(2):
_t = (_t + 2 - 3) // 2 + 1
n_audio_tokens = min(_t, 750)
media += "<|audio>" + "<|audio|>" * n_audio_tokens + "<audio|>" media += "<|audio>" + "<|audio|>" * n_audio_tokens + "<audio|>"
llama_text = f"{system}<|turn>user\n{media}{text}<turn|>\n<|turn>model\n" # Non-thinking mode primes an empty thought channel so the model answers directly.
model_open = "" if thinking else "<|channel>thought\n<channel|>"
llama_text = f"{system}<|turn>user\n{text}{media}<turn|>\n<|turn>model\n{model_open}"
text_tokens = super().tokenize_with_weights(llama_text, return_word_ids) text_tokens = super().tokenize_with_weights(llama_text, return_word_ids)
@ -1178,7 +1370,6 @@ class Gemma4_Tokenizer():
class _Gemma4Tokenizer: class _Gemma4Tokenizer:
"""Tokenizer using the tokenizers (Gemma4 doesn't come with sentencepiece model)""" """Tokenizer using the tokenizers (Gemma4 doesn't come with sentencepiece model)"""
def __init__(self, tokenizer_json_bytes=None, **kwargs): def __init__(self, tokenizer_json_bytes=None, **kwargs):
from tokenizers import Tokenizer
if isinstance(tokenizer_json_bytes, torch.Tensor): if isinstance(tokenizer_json_bytes, torch.Tensor):
tokenizer_json_bytes = bytes(tokenizer_json_bytes.tolist()) tokenizer_json_bytes = bytes(tokenizer_json_bytes.tolist())
self.tokenizer = Tokenizer.from_str(tokenizer_json_bytes.decode("utf-8")) self.tokenizer = Tokenizer.from_str(tokenizer_json_bytes.decode("utf-8"))
@ -1224,6 +1415,30 @@ class Gemma4Tokenizer(sd1_clip.SD1Tokenizer):
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma4", tokenizer=self.tokenizer_class) super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma4", tokenizer=self.tokenizer_class)
class Gemma4UnifiedSDTokenizer(Gemma4SDTokenizer):
"""Encoder-free (gemma4_unified) audio: raw 16kHz waveform frames instead of mel spectrogram."""
embedding_size = 3840
def _extract_audio_features(self, waveform, sample_rate):
audio = self._resample_16k(waveform, sample_rate)
spt = 640 # audio_samples_per_token (40ms at 16kHz)
pad = (-audio.shape[0]) % spt
if pad:
audio = torch.nn.functional.pad(audio, (0, pad))
num_tokens = audio.shape[0] // spt
feats = audio[:num_tokens * spt].reshape(num_tokens, spt)
feats = feats[:750] # audio_seq_length cap (matches reference truncation, ~30s)
mask = torch.ones(feats.shape[0], dtype=torch.bool)
return feats, mask
def _audio_token_count(self, num_samples):
return min((num_samples + 639) // 640, 750)
class Gemma4UnifiedTokenizer(Gemma4Tokenizer):
tokenizer_class = Gemma4UnifiedSDTokenizer
# Model wrappers # Model wrappers
class Gemma4Model(sd1_clip.SDClipModel): class Gemma4Model(sd1_clip.SDClipModel):
model_class = None model_class = None
@ -1256,7 +1471,7 @@ class Gemma4Model(sd1_clip.SDClipModel):
expanded_idx += 1 expanded_idx += 1
initial_token_ids = [ids] initial_token_ids = [ids]
input_ids = torch.tensor(initial_token_ids, device=self.execution_device) input_ids = torch.tensor(initial_token_ids, device=self.execution_device)
return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids) return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids, embeds_info=embeds_info)
def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None): def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None):
@ -1296,3 +1511,11 @@ def _make_variant(config_cls):
Gemma4_E4B = _make_variant(Gemma4Config) Gemma4_E4B = _make_variant(Gemma4Config)
Gemma4_E2B = _make_variant(Gemma4_E2B_Config) Gemma4_E2B = _make_variant(Gemma4_E2B_Config)
Gemma4_31B = _make_variant(Gemma4_31B_Config) Gemma4_31B = _make_variant(Gemma4_31B_Config)
# Gemma4 12B Unified: encoder-free multimodal, distinct base/tokenizer (not via _make_variant).
class Gemma4_12B(Gemma4UnifiedBase):
def __init__(self, config_dict, dtype, device, operations):
super().__init__()
self._init_model(Gemma4_12B_Config(**config_dict), dtype, device, operations)
Gemma4_12B.tokenizer = Gemma4UnifiedTokenizer

View File

@ -0,0 +1,97 @@
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_

View File

@ -876,7 +876,7 @@ class BaseGenerate:
torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), 0)) torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), 0))
return past_key_values return past_key_values
def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None): def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None, embeds_info=None):
device = embeds.device device = embeds.device
if stop_tokens is None: if stop_tokens is None:
@ -911,7 +911,7 @@ class BaseGenerate:
if step == 0 and deepstack_embeds is not None: if step == 0 and deepstack_embeds is not None:
extra["deepstack_embeds"] = deepstack_embeds extra["deepstack_embeds"] = deepstack_embeds
extra["visual_pos_masks"] = visual_pos_masks extra["visual_pos_masks"] = visual_pos_masks
x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra) x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra, embeds_info=(embeds_info if step == 0 else None))
logits = self.logits(x)[:, -1] logits = self.logits(x)[:, -1]
next_token = self.sample_token(logits, temperature, top_k, top_p, min_p, repetition_penalty, initial_tokens + generated_token_ids, generator, do_sample=do_sample, presence_penalty=presence_penalty) next_token = self.sample_token(logits, temperature, top_k, top_p, min_p, repetition_penalty, initial_tokens + generated_token_ids, generator, do_sample=do_sample, presence_penalty=presence_penalty)
token_id = next_token[0].item() token_id = next_token[0].item()

View File

@ -0,0 +1,94 @@
"""Mage-Flow text encoder: Qwen3-VL-4B, last hidden state (2560-dim).
Mage-Flow conditions on the final hidden state of Qwen3-VL-4B with the leading
system + user-opening template tokens stripped (reference start_idx 34 for t2i,
64 for edit). The t2i template is identical to Qwen-Image's; the edit template
uses the same system prompt as Qwen-Image-Edit with "Image N: " reference
prefixes and no <think> block.
"""
import numbers
import torch
import comfy.text_encoders.qwen3vl
from comfy import sd1_clip
MAGE_VISION_BLOCK = "<|vision_start|><|image_pad|><|vision_end|>"
MAGE_T2I_TEMPLATE = "<|im_start|>system\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"
MAGE_EDIT_TEMPLATE = "<|im_start|>system\nDescribe the key features of the input image (color, shape, size, texture, objects, background), then explain how the user's text instruction should alter or modify the image. Generate a new image that meets the user's requirements while maintaining consistency with the original input where appropriate.<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
class MageFlowTokenizer(comfy.text_encoders.qwen3vl.Qwen3VLTokenizer):
def __init__(self, embedding_directory=None, tokenizer_data={}):
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, model_type="qwen3vl_4b")
self.llama_template = MAGE_T2I_TEMPLATE
self.llama_template_images = MAGE_EDIT_TEMPLATE
def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=True, **kwargs):
image = kwargs.get("image", None)
if image is not None and len(images) == 0:
images = [image[i:i + 1] for i in range(image.shape[0])]
if llama_template is None:
if len(images) > 0:
# Training-time multi-reference body: "Image 1: <ph>Image 2: <ph>...{instruction}"
prefix = "".join("Image {}: {}".format(j + 1, MAGE_VISION_BLOCK) for j in range(len(images)))
llama_template = self.llama_template_images.replace("{}", prefix + "{}", 1)
else:
llama_template = self.llama_template
# thinking=True: Mage templates end at "<|im_start|>assistant\n" with no <think> block.
return super().tokenize_with_weights(text, return_word_ids=return_word_ids, llama_template=llama_template, images=images, prevent_empty_text=prevent_empty_text, thinking=thinking, **kwargs)
class MageFlowQwen3VLClipModel(comfy.text_encoders.qwen3vl.Qwen3VLClipModel):
def __init__(self, device="cpu", dtype=None, attention_mask=True, model_options={}, model_type="qwen3vl_4b"):
super().__init__(device=device, dtype=dtype, attention_mask=attention_mask, model_options=model_options, model_type=model_type)
# apply the final RMSNorm to the tapped last layer (HF last_hidden_state)
self.layer_norm_hidden_state = True
class MageFlowTEModel(sd1_clip.SD1ClipModel):
def __init__(self, device="cpu", dtype=None, model_options={}):
clip_model = lambda **kw: MageFlowQwen3VLClipModel(**kw, model_type="qwen3vl_4b") # noqa: E731
super().__init__(device=device, dtype=dtype, name="qwen3vl_4b", clip_model=clip_model, model_options=model_options)
def encode_token_weights(self, token_weight_pairs, template_end=-1):
# Strip the system + user-opening prefix (reference drop_idx: 34 t2i / 64 edit).
out, pooled, extra = super().encode_token_weights(token_weight_pairs)
tok_pairs = token_weight_pairs["qwen3vl_4b"][0]
count_im_start = 0
if template_end == -1:
for i, v in enumerate(tok_pairs):
elem = v[0]
if not torch.is_tensor(elem):
if isinstance(elem, numbers.Integral):
if elem == 151644 and count_im_start < 2: # <|im_start|>
template_end = i
count_im_start += 1
if out.shape[1] > (template_end + 3):
if tok_pairs[template_end + 1][0] == 872: # "user"
if tok_pairs[template_end + 2][0] == 198: # "\n"
template_end += 3
out = out[:, template_end:]
if "attention_mask" in extra:
extra["attention_mask"] = extra["attention_mask"][:, template_end:]
if extra["attention_mask"].sum() == torch.numel(extra["attention_mask"]):
extra.pop("attention_mask") # attention mask is useless if no masked elements
return out, pooled, extra
def te(dtype_llama=None, llama_quantization_metadata=None):
class MageFlowTEModel_(MageFlowTEModel):
def __init__(self, device="cpu", dtype=None, model_options={}):
if dtype_llama is not None:
dtype = dtype_llama
if llama_quantization_metadata is not None:
model_options = model_options.copy()
model_options["quantization_metadata"] = llama_quantization_metadata
super().__init__(device=device, dtype=dtype, model_options=model_options)
return MageFlowTEModel_

View File

@ -158,12 +158,12 @@ class Qwen3VLTokenizer(sd1_clip.SD1Tokenizer):
self.llama_template = "<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" self.llama_template = "<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
self.llama_template_images = "<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>{}<|im_end|>\n<|im_start|>assistant\n" self.llama_template_images = "<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>{}<|im_end|>\n<|im_start|>assistant\n"
def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=False, **kwargs): def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=False, skip_template=False, **kwargs):
image = kwargs.get("image", None) image = kwargs.get("image", None)
if image is not None and len(images) == 0: if image is not None and len(images) == 0:
images = [image[i:i + 1] for i in range(image.shape[0])] images = [image[i:i + 1] for i in range(image.shape[0])]
skip_template = text.startswith('<|im_start|>') skip_template = skip_template or text.startswith('<|im_start|>')
if prevent_empty_text and text == '': if prevent_empty_text and text == '':
text = ' ' text = ' '

View File

@ -15,6 +15,7 @@ def process_qwen2vl_images(
merge_size: int = 2, merge_size: int = 2,
image_mean: list = None, image_mean: list = None,
image_std: list = None, image_std: list = None,
interpolation: str = "bilinear",
): ):
if image_mean is None: if image_mean is None:
image_mean = [0.48145466, 0.4578275, 0.40821073] image_mean = [0.48145466, 0.4578275, 0.40821073]
@ -47,10 +48,9 @@ def process_qwen2vl_images(
img_resized = F.interpolate( img_resized = F.interpolate(
img.unsqueeze(0), img.unsqueeze(0),
size=(h_bar, w_bar), size=(h_bar, w_bar),
mode='bilinear', mode=interpolation,
align_corners=False align_corners=False
).squeeze(0) ).squeeze(0)
normalized = img_resized.clone() normalized = img_resized.clone()
for c in range(3): for c in range(3):
normalized[c] = (img_resized[c] - image_mean[c]) / image_std[c] normalized[c] = (img_resized[c] - image_mean[c]) / image_std[c]

View File

@ -1,5 +1,6 @@
from av.container import InputContainer from av.container import InputContainer
from av.subtitles.stream import SubtitleStream from av.subtitles.stream import SubtitleStream
from av.video.reformatter import ColorRange
from fractions import Fraction from fractions import Fraction
from typing import Optional from typing import Optional
from .._input import AudioInput, VideoInput from .._input import AudioInput, VideoInput
@ -9,6 +10,7 @@ import itertools
import json import json
import numpy as np import numpy as np
import math import math
import os
import torch import torch
from .._util import VideoContainer, VideoCodec, VideoComponents from .._util import VideoContainer, VideoCodec, VideoComponents
import logging import logging
@ -58,6 +60,57 @@ def video_stream_bit_depth(stream) -> int:
return max(component.bits for component in stream.format.components) 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 VideoFromFile(VideoInput):
""" """
Class representing video input from a file. Class representing video input from a file.
@ -192,13 +245,10 @@ class VideoFromFile(VideoInput):
return estimated_frames return estimated_frames
# 3. Last resort: decode frames and count them (streaming) # 3. Last resort: decode frames and count them (streaming)
if self.__start_time < 0: start_time, duration = self.get_active_trim_window()
start_time = max(self._get_raw_duration() + self.__start_time, 0)
else:
start_time = self.__start_time
frame_count = 1 frame_count = 1
start_pts = int(start_time / video_stream.time_base) start_pts = int(start_time / video_stream.time_base)
end_pts = int((start_time + self.__duration) / video_stream.time_base) end_pts = int((start_time + duration) / video_stream.time_base)
container.seek(start_pts, stream=video_stream) container.seek(start_pts, stream=video_stream)
frame_iterator = ( frame_iterator = (
container.decode(video_stream) container.decode(video_stream)
@ -253,17 +303,14 @@ class VideoFromFile(VideoInput):
def get_components_internal(self, container: InputContainer) -> VideoComponents: def get_components_internal(self, container: InputContainer) -> VideoComponents:
video_stream = self._get_first_video_stream(container) video_stream = self._get_first_video_stream(container)
if self.__start_time < 0: start_time, duration = self.get_active_trim_window()
start_time = max(self._get_raw_duration() + self.__start_time, 0)
else:
start_time = self.__start_time
# Get video frames # Get video frames
frames = [] frames = []
audio_frames = [] audio_frames = []
alphas = None alphas = None
start_pts = int(start_time / video_stream.time_base) start_pts = int(start_time / video_stream.time_base)
end_pts = int((start_time + self.__duration) / video_stream.time_base) end_pts = int((start_time + duration) / video_stream.time_base)
if start_pts != 0: if start_pts != 0:
container.seek(start_pts, stream=video_stream) container.seek(start_pts, stream=video_stream)
@ -281,18 +328,11 @@ class VideoFromFile(VideoInput):
video_done = False video_done = False
audio_done = True audio_done = True
# Use the last decodable audio stream. Streams FFmpeg has no decoder for have no codec context, audio_stream = last_decodable_audio_stream(container)
# 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: if audio_stream is not None:
streams += [audio_stream] streams += [audio_stream]
resampler = av.audio.resampler.AudioResampler(format='fltp') resampler = av.audio.resampler.AudioResampler(format='fltp')
audio_done = False 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): for packet in container.demux(*streams):
if video_done and audio_done: if video_done and audio_done:
@ -305,7 +345,7 @@ class VideoFromFile(VideoInput):
for frame in packet.decode(): for frame in packet.decode():
if frame.pts < start_pts: if frame.pts < start_pts:
continue continue
if self.__duration and frame.pts >= end_pts: if duration and frame.pts >= end_pts:
video_done = True video_done = True
break break
@ -372,7 +412,7 @@ class VideoFromFile(VideoInput):
map(resampler.resample, packet.decode()) map(resampler.resample, packet.decode())
) )
for frame in aframes: for frame in aframes:
if self.__duration and frame.time > start_time + self.__duration: if duration and frame.time > start_time + duration:
audio_done = True audio_done = True
break break
@ -394,8 +434,8 @@ class VideoFromFile(VideoInput):
if len(audio_frames) > 0: if len(audio_frames) > 0:
audio_data = np.concatenate(audio_frames, axis=1) # shape: (channels, total_samples) audio_data = np.concatenate(audio_frames, axis=1) # shape: (channels, total_samples)
if self.__duration: if duration:
audio_data = audio_data[..., :int(self.__duration * audio_stream.sample_rate)] audio_data = audio_data[..., :int(duration * audio_stream.sample_rate)]
audio_tensor = torch.from_numpy(audio_data).unsqueeze(0) # shape: (1, channels, total_samples) audio_tensor = torch.from_numpy(audio_data).unsqueeze(0) # shape: (1, channels, total_samples)
audio = AudioInput({ audio = AudioInput({
@ -441,28 +481,14 @@ class VideoFromFile(VideoInput):
if not reuse_streams: if not reuse_streams:
if bit_depth is None: if bit_depth is None:
bit_depth = source_bit_depth bit_depth = source_bit_depth
components = self.get_components_internal(container) return self._save_transcoded(container, path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth)
video = VideoFromComponents(components)
return video.save_to(
path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth,
)
streams = container.streams streams = container.streams
open_kwargs = get_open_write_kwargs(path, container_format, format) open_kwargs = get_open_write_kwargs(path, container_format, format)
with av.open(path, **open_kwargs) as output_container: with av.open(path, **open_kwargs) as output_container:
# Copy over the original metadata # Add metadata before writing any streams
for key, value in container.metadata.items(): write_output_metadata(container, output_container, metadata)
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. # Add streams to the new container. Streams with no codec context cannot be used as an output template.
stream_map = {} stream_map = {}
@ -480,6 +506,282 @@ class VideoFromFile(VideoInput):
packet.stream = stream_map[packet.stream] packet.stream = stream_map[packet.stream]
output_container.mux(packet) 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): def _get_first_video_stream(self, container: InputContainer):
if len(container.streams.video): if len(container.streams.video):
return container.streams.video[0] return container.streams.video[0]
@ -527,22 +829,12 @@ class VideoFromComponents(VideoInput):
bit_depth: int | None = None, bit_depth: int | None = None,
): ):
"""Save the video to a file path or BytesIO buffer.""" """Save the video to a file path or BytesIO buffer."""
if format != VideoContainer.AUTO and format != VideoContainer.MP4: open_kwargs = mp4_output_open_kwargs(path, format, codec)
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). # None means "use the depth this video was created with" (CreateVideo's choice).
if bit_depth is None: if bit_depth is None:
bit_depth = self.__bit_depth bit_depth = self.__bit_depth
is_10bit = bit_depth >= 10 is_10bit = bit_depth >= 10
extra_kwargs = {} with av.open(path, **open_kwargs) as output:
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 # Add metadata before writing any streams
if metadata is not None: if metadata is not None:
for key, value in metadata.items(): for key, value in metadata.items():

View File

@ -28,6 +28,9 @@ ANTHROPIC_IMAGE_MAX_PIXELS = 1568 * 1568
CLAUDE_MAX_IMAGES = 20 CLAUDE_MAX_IMAGES = 20
CLAUDE_MODELS: dict[str, str] = { CLAUDE_MODELS: dict[str, str] = {
"Opus 4.8": "claude-opus-4-8",
"Fable 5": "claude-fable-5",
"Sonnet 5": "claude-sonnet-5",
"Opus 4.7": "claude-opus-4-7", "Opus 4.7": "claude-opus-4-7",
"Opus 4.6": "claude-opus-4-6", "Opus 4.6": "claude-opus-4-6",
"Sonnet 4.6": "claude-sonnet-4-6", "Sonnet 4.6": "claude-sonnet-4-6",
@ -36,9 +39,12 @@ CLAUDE_MODELS: dict[str, str] = {
} }
_THINKING_UNSUPPORTED = {"Haiku 4.5"} _THINKING_UNSUPPORTED = {"Haiku 4.5"}
# Models that use the newer "adaptive" thinking mode (Opus 4.7 requires it; older models keep the explicit budget API). # Models that use the newer "adaptive" thinking mode (Opus 4.7+ require it; older models keep the explicit budget API).
# Anthropic decides the actual budget when adaptive is used, based on the `output_config.effort` hint. # Anthropic decides the actual budget when adaptive is used, based on the `output_config.effort` hint.
_ADAPTIVE_THINKING_MODELS = {"Opus 4.7", "Opus 4.6", "Sonnet 4.6"} _ADAPTIVE_THINKING_MODELS = {"Opus 4.8", "Sonnet 5", "Opus 4.7", "Opus 4.6", "Sonnet 4.6"}
_ALWAYS_THINKING_MODELS = {"Fable 5"}
_EXPLICIT_THINKING_OFF_MODELS = {"Sonnet 5"}
_NO_TEMPERATURE_MODELS = {"Opus 4.8", "Fable 5", "Sonnet 5"}
# Budget mode (Sonnet 4.5): effort -> reasoning budget in tokens. Must be < max_tokens. # Budget mode (Sonnet 4.5): effort -> reasoning budget in tokens. Must be < max_tokens.
# Sized so even the "high" budget fits comfortably under the default max_tokens=32768. # Sized so even the "high" budget fits comfortably under the default max_tokens=32768.
@ -60,20 +66,33 @@ def _claude_model_inputs(model_label: str):
tooltip="Maximum number of tokens to generate (includes reasoning tokens when enabled).", tooltip="Maximum number of tokens to generate (includes reasoning tokens when enabled).",
advanced=True, advanced=True,
), ),
IO.Float.Input(
"temperature",
default=1.0,
min=0.0,
max=1.0,
step=0.01,
tooltip=(
"Controls randomness. 0.0 is deterministic, 1.0 is most random. "
"Ignored for Opus 4.7 and any model when reasoning_effort is set."
),
advanced=True,
),
] ]
if model_label not in _THINKING_UNSUPPORTED: if model_label not in _NO_TEMPERATURE_MODELS:
inputs.append(
IO.Float.Input(
"temperature",
default=1.0,
min=0.0,
max=1.0,
step=0.01,
tooltip=(
"Controls randomness. 0.0 is deterministic, 1.0 is most random. "
"Ignored for Opus 4.7 and any model when reasoning_effort is set."
),
advanced=True,
)
)
if model_label in _ALWAYS_THINKING_MODELS:
inputs.append(
IO.Combo.Input(
"reasoning_effort",
options=[e for e in _REASONING_EFFORTS if e != "off"],
default="high",
tooltip="Extended thinking effort. Reasoning is always enabled for this model.",
advanced=True,
)
)
elif model_label not in _THINKING_UNSUPPORTED:
inputs.append( inputs.append(
IO.Combo.Input( IO.Combo.Input(
"reasoning_effort", "reasoning_effort",
@ -88,6 +107,12 @@ def _claude_model_inputs(model_label: str):
def _model_price_per_million(model: str) -> tuple[float, float] | None: def _model_price_per_million(model: str) -> tuple[float, float] | None:
"""Return (input_per_1M, output_per_1M) USD for a Claude model, or None if unknown.""" """Return (input_per_1M, output_per_1M) USD for a Claude model, or None if unknown."""
if "fable-5" in model:
return 14.30, 71.50
if "opus-4-8" in model:
return 7.15, 35.75
if "sonnet-5" in model:
return 2.86, 14.30
if "opus-4-7" in model or "opus-4-6" in model or "opus-4-5" in model: if "opus-4-7" in model or "opus-4-6" in model or "opus-4-5" in model:
return 5.0, 25.0 return 5.0, 25.0
if "sonnet-4" in model: if "sonnet-4" in model:
@ -213,7 +238,22 @@ class ClaudeNode(IO.ComfyNode):
expr=""" expr="""
( (
$m := widgets.model; $m := widgets.model;
$contains($m, "opus") ? { $contains($m, "fable") ? {
"type": "list_usd",
"usd": [0.0143, 0.0715],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "opus 4.8") ? {
"type": "list_usd",
"usd": [0.00715, 0.03575],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "sonnet 5") ? {
"type": "list_usd",
"usd": [0.00286, 0.0143],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "opus") ? {
"type": "list_usd", "type": "list_usd",
"usd": [0.005, 0.025], "usd": [0.005, 0.025],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
@ -247,18 +287,23 @@ class ClaudeNode(IO.ComfyNode):
model_label = model["model"] model_label = model["model"]
max_tokens = model.get("max_tokens", 32768) max_tokens = model.get("max_tokens", 32768)
reasoning_effort = model.get("reasoning_effort", "off") reasoning_effort = model.get("reasoning_effort", "off")
thinking_enabled = reasoning_effort not in ("off", None) and model_label not in _THINKING_UNSUPPORTED always_thinking = model_label in _ALWAYS_THINKING_MODELS
thinking_enabled = always_thinking or (
reasoning_effort not in ("off", None) and model_label not in _THINKING_UNSUPPORTED
)
# Anthropic requires temperature to be unset (defaults to 1.0) when thinking is enabled. # Anthropic requires temperature to be unset (defaults to 1.0) when thinking is enabled.
# Opus 4.7 also rejects user-supplied temperature. # Opus 4.7 also rejects user-supplied temperature.
if thinking_enabled or model_label == "Opus 4.7": if model_label in _NO_TEMPERATURE_MODELS or thinking_enabled or model_label == "Opus 4.7":
temperature = None temperature = None
else: else:
temperature = model.get("temperature", 1.0) temperature = model.get("temperature", 1.0)
thinking_cfg: AnthropicThinkingConfig | None = None thinking_cfg: AnthropicThinkingConfig | None = None
output_cfg: AnthropicOutputConfig | None = None output_cfg: AnthropicOutputConfig | None = None
if thinking_enabled: if always_thinking:
output_cfg = AnthropicOutputConfig(effort=reasoning_effort)
elif thinking_enabled:
if model_label in _ADAPTIVE_THINKING_MODELS: if model_label in _ADAPTIVE_THINKING_MODELS:
# Adaptive mode - Anthropic chooses the budget based on effort hint # Adaptive mode - Anthropic chooses the budget based on effort hint
thinking_cfg = AnthropicThinkingConfig(type="adaptive") thinking_cfg = AnthropicThinkingConfig(type="adaptive")
@ -268,6 +313,8 @@ class ClaudeNode(IO.ComfyNode):
budget = _REASONING_BUDGET[reasoning_effort] budget = _REASONING_BUDGET[reasoning_effort]
budget = min(budget, max(1024, max_tokens - 1024)) budget = min(budget, max(1024, max_tokens - 1024))
thinking_cfg = AnthropicThinkingConfig(type="enabled", budget_tokens=budget) thinking_cfg = AnthropicThinkingConfig(type="enabled", budget_tokens=budget)
elif model_label in _EXPLICIT_THINKING_OFF_MODELS:
thinking_cfg = AnthropicThinkingConfig(type="disabled")
image_tensors: list[Input.Image] = [t for t in (images or {}).values() if t is not None] image_tensors: list[Input.Image] = [t for t in (images or {}).values() if t is not None]
if sum(get_number_of_images(t) for t in image_tensors) > CLAUDE_MAX_IMAGES: if sum(get_number_of_images(t) for t in image_tensors) > CLAUDE_MAX_IMAGES:
@ -293,6 +340,11 @@ class ClaudeNode(IO.ComfyNode):
), ),
price_extractor=calculate_tokens_price, price_extractor=calculate_tokens_price,
) )
if response.stop_reason == "refusal":
raise ValueError(
"Claude declined to answer this request for safety reasons. "
"Rephrase the prompt or try a different model."
)
return IO.NodeOutput(_get_text_from_response(response) or "Empty response from Claude model.") return IO.NodeOutput(_get_text_from_response(response) or "Empty response from Claude model.")

View File

@ -2690,7 +2690,8 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
"with ByteDance Seed Audio 1.0. Describe the voice(s), emotion, ambience, background music " "with ByteDance Seed Audio 1.0. Describe the voice(s), emotion, ambience, background music "
"and sound effects in the prompt, and include the lines to speak. Optionally pick a built-in " "and sound effects in the prompt, and include the lines to speak. Optionally pick a built-in "
"preset voice, clone voices from up to 3 reference clips (tagged @Audio1-3 in the prompt), " "preset voice, clone voices from up to 3 reference clips (tagged @Audio1-3 in the prompt), "
"or derive a voice from a character image. Up to 2 minutes of audio per run." "or derive a voice from a character image. Up to 2 minutes of audio per run. "
"The multilingual model supports 20 languages and timestamp-based timing control."
), ),
inputs=[ inputs=[
IO.String.Input( IO.String.Input(
@ -2701,7 +2702,9 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
"Describe the voice(s), emotion, pacing, ambience, background music and sound " "Describe the voice(s), emotion, pacing, ambience, background music and sound "
"effects, and include the lines to speak (name characters inline for dialogue). " "effects, and include the lines to speak (name characters inline for dialogue). "
"In 'audio reference' mode, refer to connected clips by order as @Audio1, @Audio2, " "In 'audio reference' mode, refer to connected clips by order as @Audio1, @Audio2, "
"@Audio3. Maximum 3000 characters." "@Audio3. With the multilingual model, a quoted line can start with a timestamp "
'range that controls when and how long it is spoken, e.g. "[5.5s:8.0s] Wait for me!". '
"Write the prompt in the same language as the lines to speak. Maximum 3000 characters."
), ),
), ),
IO.DynamicCombo.Input( IO.DynamicCombo.Input(
@ -2796,6 +2799,19 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
tooltip="Seed controls whether the node should re-run; " tooltip="Seed controls whether the node should re-run; "
"results are non-deterministic regardless of seed.", "results are non-deterministic regardless of seed.",
), ),
IO.Combo.Input(
"model",
options=["seed-audio-1.0-multilingual", "seed-audio-1.0"],
default="seed-audio-1.0-multilingual",
optional=True,
tooltip=(
"seed-audio-1.0-multilingual: 20 languages (English, Chinese, Japanese, Korean, "
"Mexican & Castilian Spanish, Indonesian, German, Brazilian Portuguese, French, "
"Thai, Vietnamese, Malay, Filipino, Italian, Russian, Dutch, Polish, Turkish, "
'Swedish) plus per-sentence timing control via "[5.5s:8.0s] ..." timestamps. '
"seed-audio-1.0: English and Chinese only, no timing control."
),
),
], ],
outputs=[IO.Audio.Output()], outputs=[IO.Audio.Output()],
hidden=[ hidden=[
@ -2819,6 +2835,7 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
loudness_rate: int, loudness_rate: int,
pitch_rate: int, pitch_rate: int,
seed: int, seed: int,
model: str = "seed-audio-1.0-multilingual",
) -> IO.NodeOutput: ) -> IO.NodeOutput:
mode = reference_mode["reference_mode"] mode = reference_mode["reference_mode"]
audio_indices = connected_audio_indices(reference_mode) audio_indices = connected_audio_indices(reference_mode)
@ -2845,6 +2862,7 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
ApiEndpoint(path="/proxy/byteplus/api/v3/tts/create", method="POST"), ApiEndpoint(path="/proxy/byteplus/api/v3/tts/create", method="POST"),
response_model=SeedAudioResponse, response_model=SeedAudioResponse,
data=SeedAudioRequest( data=SeedAudioRequest(
model=model,
text_prompt=text_prompt, text_prompt=text_prompt,
references=references, references=references,
audio_config=SeedAudioConfig( audio_config=SeedAudioConfig(

View File

@ -60,6 +60,7 @@ GEMINI_INTERACTIONS_ENDPOINT = "/proxy/gemini-interactions"
GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB
GEMINI_URL_INPUT_BUDGET = 10 GEMINI_URL_INPUT_BUDGET = 10
GEMINI_MAX_INLINE_BYTES = 18 * 1024 * 1024 GEMINI_MAX_INLINE_BYTES = 18 * 1024 * 1024
GEMINI_INTERACTIONS_MAX_INLINE_BYTES = 90 * 1024 * 1024 # the Interactions API rejects requests over ~100MiB
GEMINI_IMAGE_SYS_PROMPT = ( GEMINI_IMAGE_SYS_PROMPT = (
"You are an expert image-generation engine. You must ALWAYS produce an image.\n" "You are an expert image-generation engine. You must ALWAYS produce an image.\n"
"Interpret all user input—regardless of " "Interpret all user input—regardless of "
@ -469,9 +470,10 @@ async def build_gemini_media_parts(
part, nbytes = _media_inline_part(kind, payload) part, nbytes = _media_inline_part(kind, payload)
inline_bytes += nbytes inline_bytes += nbytes
if inline_bytes > max_inline_bytes: if inline_bytes > max_inline_bytes:
detail = f" after the first {url_budget} inputs are uploaded as URLs" if url_budget else ""
raise ValueError( raise ValueError(
f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB after the first " f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB{detail}). "
f"{url_budget} inputs are uploaded as URLs). Reduce the number or size of attached media." "Reduce the number or size of attached media."
) )
parts.append(part) parts.append(part)
return parts return parts
@ -1738,7 +1740,14 @@ class GeminiVideoOmni(IO.ComfyNode):
parts: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = [] parts: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = []
if images or videos: if images or videos:
media_parts = await build_gemini_media_parts(cls, images, [], videos) # The Interactions API accepts video only inline or as a Files API URI, not as an HTTP URL.
media_parts = await build_gemini_media_parts(
cls, [], [], videos, url_budget=0, max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES
)
video_inline_bytes = sum(len(p.inlineData.data) for p in media_parts)
media_parts += await build_gemini_media_parts(
cls, images, [], [], max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES - video_inline_bytes
)
parts.extend(to_interaction_media_part(p) for p in media_parts) parts.extend(to_interaction_media_part(p) for p in media_parts)
parts.append(GeminiInteractionTextPart(text=prompt)) parts.append(GeminiInteractionTextPart(text=prompt))
interaction = await sync_op( interaction = await sync_op(

View File

@ -45,27 +45,40 @@ class _ModelSpec:
MODELS: list[_ModelSpec] = [ MODELS: list[_ModelSpec] = [
_ModelSpec("anthropic/claude-opus-4.7", "frontier_reasoning", 0.000005, 0.000025, max_images=20), _ModelSpec("anthropic/claude-opus-5", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20),
_ModelSpec("openai/gpt-5.5-pro", "frontier_reasoning", 0.00003, 0.00018, max_images=20), _ModelSpec("anthropic/claude-opus-4.8", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20),
_ModelSpec("openai/gpt-5.5", "frontier_reasoning", 0.000005, 0.00003, max_images=20), _ModelSpec("anthropic/claude-opus-4.7", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20),
_ModelSpec("google/gemini-3.5-flash", "reasoning", 0.0000015, 0.000009, max_images=20, max_videos=4), _ModelSpec("anthropic/claude-fable-5", "frontier_reasoning", 0.0000143, 0.0000715, max_images=20),
_ModelSpec("x-ai/grok-4.20", "reasoning", 0.00000125, 0.0000025, max_images=20), _ModelSpec("anthropic/claude-sonnet-5", "frontier_reasoning", 0.00000286, 0.0000143, max_images=20),
_ModelSpec("x-ai/grok-4.3", "reasoning", 0.00000125, 0.0000025, max_images=20), _ModelSpec("anthropic/claude-haiku-4.5", "frontier_reasoning", 0.00000143, 0.00000715, max_images=20),
_ModelSpec("deepseek/deepseek-v4-pro", "reasoning", 0.000000435, 0.00000087), _ModelSpec("openai/gpt-5.6-sol-pro", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20),
_ModelSpec("deepseek/deepseek-v4-flash", "reasoning", 0.000000112, 0.000000224), _ModelSpec("openai/gpt-5.6-sol", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20),
_ModelSpec("deepseek/deepseek-v3.2", "reasoning", 0.000000252, 0.000000378), _ModelSpec("openai/gpt-5.6-terra-pro", "frontier_reasoning", 0.000003575, 0.00002145, max_images=20),
_ModelSpec("qwen/qwen3.6-max-preview", "reasoning", 0.00000104, 0.00000624), _ModelSpec("openai/gpt-5.6-terra", "frontier_reasoning", 0.000003575, 0.00002145, max_images=20),
_ModelSpec("qwen/qwen3.6-plus", "reasoning", 0.000000325, 0.00000195, max_images=10, max_videos=4), _ModelSpec("openai/gpt-5.6-luna-pro", "frontier_reasoning", 0.00000143, 0.00000858, max_images=20),
_ModelSpec("qwen/qwen3.6-flash", "reasoning", 0.0000001875, 0.000001125, max_images=10, max_videos=4), _ModelSpec("openai/gpt-5.6-luna", "frontier_reasoning", 0.00000143, 0.00000858, max_images=20),
_ModelSpec("mistralai/mistral-large-2512", "standard", 0.0000005, 0.0000015, max_images=8), _ModelSpec("openai/gpt-5.5-pro", "frontier_reasoning", 0.0000429, 0.0002574, max_images=20),
_ModelSpec("mistralai/mistral-medium-3-5", "reasoning", 0.0000015, 0.0000075, max_images=8), _ModelSpec("openai/gpt-5.5", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20),
_ModelSpec("z-ai/glm-4.6", "reasoning", 0.00000043, 0.00000174), _ModelSpec("google/gemini-3.5-flash", "reasoning", 0.000002145, 0.00001287, max_images=20, max_videos=4),
_ModelSpec("z-ai/glm-5", "reasoning", 0.0000006, 0.00000192), _ModelSpec("x-ai/grok-4.5", "reasoning", 0.00000286, 0.00000858, max_images=20),
_ModelSpec("moonshotai/kimi-k2.6", "reasoning", 0.00000073, 0.00000349, max_images=10), _ModelSpec("x-ai/grok-4.20", "reasoning", 0.0000017875, 0.000003575, max_images=20),
_ModelSpec("moonshotai/kimi-k2-thinking", "reasoning", 0.0000006, 0.0000025), _ModelSpec("x-ai/grok-4.3", "reasoning", 0.0000017875, 0.000003575, max_images=20),
_ModelSpec("perplexity/sonar-pro", "perplexity", 0.000003, 0.000015), _ModelSpec("deepseek/deepseek-v4-pro", "reasoning", 0.00000062205, 0.0000012441),
_ModelSpec("perplexity/sonar-reasoning-pro", "perplexity_reasoning", 0.000002, 0.000008), _ModelSpec("deepseek/deepseek-v4-flash", "reasoning", 0.00000016016, 0.00000032032),
_ModelSpec("perplexity/sonar-deep-research", "perplexity_reasoning", 0.000002, 0.000008), _ModelSpec("deepseek/deepseek-v3.2", "reasoning", 0.00000036036, 0.00000054054),
_ModelSpec("qwen/qwen3.6-max-preview", "reasoning", 0.0000014872, 0.0000089232),
_ModelSpec("qwen/qwen3.6-plus", "reasoning", 0.00000046475, 0.0000027885, max_images=10, max_videos=4),
_ModelSpec("qwen/qwen3.6-flash", "reasoning", 0.000000268125, 0.00000160875, max_images=10, max_videos=4),
_ModelSpec("mistralai/mistral-large-2512", "standard", 0.000000715, 0.000002145, max_images=8),
_ModelSpec("mistralai/mistral-medium-3-5", "reasoning", 0.000002145, 0.000010725, max_images=8),
_ModelSpec("z-ai/glm-4.6", "reasoning", 0.0000006149, 0.0000024882),
_ModelSpec("z-ai/glm-5", "reasoning", 0.000000858, 0.0000027456),
_ModelSpec("moonshotai/kimi-k3", "reasoning", 0.00000429, 0.00002145, max_images=10),
_ModelSpec("moonshotai/kimi-k2.6", "reasoning", 0.0000010439, 0.0000049907, max_images=10),
_ModelSpec("moonshotai/kimi-k2-thinking", "reasoning", 0.000000858, 0.000003575),
_ModelSpec("perplexity/sonar-pro", "perplexity", 0.00000429, 0.00002145),
_ModelSpec("perplexity/sonar-reasoning-pro", "perplexity_reasoning", 0.00000286, 0.00001144),
_ModelSpec("perplexity/sonar-deep-research", "perplexity_reasoning", 0.00000286, 0.00001144),
] ]
_MODELS_BY_SLUG: dict[str, _ModelSpec] = {m.slug: m for m in MODELS} _MODELS_BY_SLUG: dict[str, _ModelSpec] = {m.slug: m for m in MODELS}
@ -148,7 +161,7 @@ def _build_model_options() -> list[IO.DynamicCombo.Option]:
def _calculate_price(response: OpenRouterChatResponse) -> float | None: def _calculate_price(response: OpenRouterChatResponse) -> float | None:
if response.usage and response.usage.cost is not None: if response.usage and response.usage.cost is not None:
return float(response.usage.cost) return float(response.usage.cost) * 1.43
return None return None
@ -269,8 +282,8 @@ class OpenRouterLLMNode(IO.ComfyNode):
essentials_category="Text Generation", essentials_category="Text Generation",
description=( description=(
"Generate text responses through OpenRouter. Routes to a curated set of popular " "Generate text responses through OpenRouter. Routes to a curated set of popular "
"models from xAI, DeepSeek, Qwen, Mistral, Z.AI (GLM), Moonshot (Kimi), and " "models from Anthropic (Claude), OpenAI (GPT), Google (Gemini), xAI (Grok), "
"Perplexity Sonar." "DeepSeek, Qwen, Mistral, Z.AI (GLM), Moonshot (Kimi), and Perplexity Sonar."
), ),
inputs=[ inputs=[
IO.String.Input( IO.String.Input(

View File

@ -194,6 +194,7 @@ class RunwayImageToVideoNodeGen3a(IO.ComfyNode):
depends_on=IO.PriceBadgeDepends(widgets=["duration"]), depends_on=IO.PriceBadgeDepends(widgets=["duration"]),
expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""", expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""",
), ),
is_deprecated=True,
) )
@classmethod @classmethod
@ -390,6 +391,7 @@ class RunwayFirstLastFrameNode(IO.ComfyNode):
depends_on=IO.PriceBadgeDepends(widgets=["duration"]), depends_on=IO.PriceBadgeDepends(widgets=["duration"]),
expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""", expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""",
), ),
is_deprecated=True,
) )
@classmethod @classmethod

View File

@ -15,6 +15,7 @@ from comfy.comfy_api_env import normalize_comfy_api_base
from comfy.deploy_environment import get_deploy_environment from comfy.deploy_environment import get_deploy_environment
from comfy.model_management import processing_interrupted from comfy.model_management import processing_interrupted
from comfy_api.latest import IO from comfy_api.latest import IO
from comfy_execution.utils import get_executing_context
from comfyui_version import __version__ as comfyui_version from comfyui_version import __version__ as comfyui_version
from .common_exceptions import ProcessingInterrupted from .common_exceptions import ProcessingInterrupted
@ -57,12 +58,16 @@ def get_comfy_api_headers(node_cls: type[IO.ComfyNode]) -> dict[str, str]:
relative/cloud URLs resolved against ``default_base_url()``; because the result relative/cloud URLs resolved against ``default_base_url()``; because the result
includes auth, callers must not attach it to arbitrary absolute/presigned URLs. includes auth, callers must not attach it to arbitrary absolute/presigned URLs.
""" """
return { headers = {
**get_auth_header(node_cls), **get_auth_header(node_cls),
"Comfy-Env": get_deploy_environment(), "Comfy-Env": get_deploy_environment(),
"Comfy-Usage-Source": get_usage_source(node_cls), "Comfy-Usage-Source": get_usage_source(node_cls),
"Comfy-Core-Version": comfyui_version, "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: def default_base_url() -> str:

View File

@ -2,6 +2,7 @@ import logging
import os import os
import json import json
import av
import numpy as np import numpy as np
import torch import torch
from PIL import Image from PIL import Image
@ -9,7 +10,7 @@ from typing_extensions import override
import folder_paths import folder_paths
import node_helpers import node_helpers
from comfy_api.latest import ComfyExtension, io from comfy_api.latest import ComfyExtension, io, Input, InputImpl, Types
def load_and_process_images(image_files, input_dir): def load_and_process_images(image_files, input_dir):
@ -42,6 +43,38 @@ def load_and_process_images(image_files, input_dir):
return output_images return output_images
VALID_VIDEO_EXTENSIONS = [".mp4", ".avi", ".mov", ".webm", ".mkv", ".flv"]
def _decode_selected_frames(video: Input.Video, indices: list[int]) -> Input.Video:
"""Decode only the requested frame indices from a video.
Opens the underlying container once, decodes frames in presentation order,
keeps only the ones whose index is in ``indices``, and returns the result
wrapped in a VideoFromComponents so it still satisfies the VideoInput
contract for downstream nodes.
"""
indices_sorted = sorted(set(indices))
max_idx = indices_sorted[-1]
source = video.get_stream_source()
frames_by_idx: dict[int, torch.Tensor] = {}
with av.open(source, mode="r") as container:
stream = container.streams.video[0]
wanted = set(indices_sorted)
for frame_idx, frame in enumerate(container.decode(stream)):
if frame_idx in wanted:
img = frame.to_ndarray(format="rgb24")
frames_by_idx[frame_idx] = torch.from_numpy(img.copy()).float() / 255.0
if frame_idx >= max_idx:
break
stacked = torch.stack([frames_by_idx[i] for i in indices])
return InputImpl.VideoFromComponents(
Types.VideoComponents(images=stacked, frame_rate=video.get_frame_rate())
)
class LoadImageDataSetFromFolderNode(io.ComfyNode): class LoadImageDataSetFromFolderNode(io.ComfyNode):
@classmethod @classmethod
def define_schema(cls): def define_schema(cls):
@ -157,6 +190,116 @@ class LoadImageTextDataSetFromFolderNode(io.ComfyNode):
return io.NodeOutput(output_tensor, captions) return io.NodeOutput(output_tensor, captions)
class LoadVideoDataSetFromFolderNode(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LoadVideoDataSetFromFolder",
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
display_name="Load Video (from Folder)",
category="video",
description="Load a dataset of videos from a specified folder and return a list of videos. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
is_experimental=True,
inputs=[
io.Combo.Input(
"folder",
options=folder_paths.get_input_subfolders(),
tooltip="The folder containing video files.",
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Lazy video references; frames are decoded only when needed downstream.",
),
],
)
@classmethod
def execute(cls, folder):
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
video_files = sorted([
f for f in os.listdir(sub_input_dir)
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
])
if not video_files:
raise ValueError(f"No video files found in {sub_input_dir}")
videos = [InputImpl.VideoFromFile(os.path.join(sub_input_dir, f)) for f in video_files]
logging.info(f"Loaded {len(videos)} lazy video references from {sub_input_dir}")
return io.NodeOutput(videos)
class LoadVideoTextDataSetFromFolderNode(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LoadVideoTextDataSetFromFolder",
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
display_name="Load Video-Text (from Folder)",
category="video",
description="Load a dataset of pairs of videos and text captions from a specified folder and return them as a list. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
is_experimental=True,
inputs=[
io.Combo.Input(
"folder",
options=folder_paths.get_input_subfolders(),
tooltip="The folder containing video files and .txt captions.",
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Lazy video references; frames are decoded only when needed downstream.",
),
io.String.Output(
display_name="texts",
is_output_list=True,
tooltip="List of text captions.",
),
],
)
@classmethod
def execute(cls, folder):
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
video_files = []
for item in sorted(os.listdir(sub_input_dir)):
path = os.path.join(sub_input_dir, item)
if any(item.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS):
video_files.append(path)
elif os.path.isdir(path):
# Support kohya-ss/sd-scripts folder structure: {repeat}_{desc}/
repeat = 1
if item.split("_")[0].isdigit():
repeat = int(item.split("_")[0])
video_files.extend([
os.path.join(path, f)
for f in sorted(os.listdir(path))
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
] * repeat)
if not video_files:
raise ValueError(f"No video files found in {sub_input_dir}")
captions = []
for vf in video_files:
caption_path = os.path.splitext(vf)[0] + ".txt"
if os.path.exists(caption_path):
with open(caption_path, "r", encoding="utf-8") as f:
captions.append(f.read().strip())
else:
captions.append("")
videos = [InputImpl.VideoFromFile(vf) for vf in video_files]
logging.info(f"Loaded {len(videos)} lazy video references with captions from {sub_input_dir}")
return io.NodeOutput(videos, captions)
def save_images_to_folder(image_list, output_dir, prefix="image", overwrite=True): def save_images_to_folder(image_list, output_dir, prefix="image", overwrite=True):
"""Utility function to save a list of image tensors to disk. """Utility function to save a list of image tensors to disk.
@ -470,7 +613,15 @@ class ImageProcessingNode(io.ComfyNode):
@classmethod @classmethod
def execute(cls, images, **kwargs): def execute(cls, images, **kwargs):
"""Execute the node. Routes to _process or _group_process based on mode.""" """Execute the node. Routes to _process or _group_process based on mode.
For individual processing (_process), automatically handles multi-frame
inputs (video tensors [T, H, W, C]) by applying _process per-frame and
concatenating the results. This allows all spatial transform nodes to
work with video without modification. Nodes that natively handle batched
tensors (e.g. pure tensor math) can set per_frame_process = False to
skip the per-frame loop.
"""
is_group = cls._detect_processing_mode() is_group = cls._detect_processing_mode()
if is_group: if is_group:
@ -489,7 +640,16 @@ class ImageProcessingNode(io.ComfyNode):
result = cls._group_process(images, **params) result = cls._group_process(images, **params)
else: else:
# Individual processing: images is single item, call _process # Individual processing: images is single item, call _process
result = cls._process(images, **params) # Auto-loop over frames for multi-frame inputs (video [T, H, W, C])
# so that PIL-based spatial transforms work per-frame automatically.
if images.shape[0] > 1 and getattr(cls, 'per_frame_process', True):
results = []
for i in range(images.shape[0]):
frame_result = cls._process(images[i:i + 1], **params)
results.append(frame_result)
result = torch.cat(results, dim=0)
else:
result = cls._process(images, **params)
return io.NodeOutput(result) return io.NodeOutput(result)
@ -803,6 +963,7 @@ class NormalizeImagesNode(ImageProcessingNode):
display_name = "Normalize Image Colors" display_name = "Normalize Image Colors"
category = "image/color" category = "image/color"
description = "Normalize images using mean and standard deviation." description = "Normalize images using mean and standard deviation."
per_frame_process = False # Pure tensor math, handles any batch size
extra_inputs = [ extra_inputs = [
io.Float.Input( io.Float.Input(
"mean", "mean",
@ -833,6 +994,7 @@ class AdjustBrightnessNode(ImageProcessingNode):
display_name = "Adjust Brightness" display_name = "Adjust Brightness"
category="image/adjustments" category="image/adjustments"
description = "Adjust the brightness of an image." description = "Adjust the brightness of an image."
per_frame_process = False # Pure tensor math, handles any batch size
extra_inputs = [ extra_inputs = [
io.Float.Input( io.Float.Input(
"factor", "factor",
@ -854,6 +1016,7 @@ class AdjustContrastNode(ImageProcessingNode):
display_name = "Adjust Contrast" display_name = "Adjust Contrast"
category="image/adjustments" category="image/adjustments"
description = "Adjust the contrast of an image." description = "Adjust the contrast of an image."
per_frame_process = False # Pure tensor math, handles any batch size
extra_inputs = [ extra_inputs = [
io.Float.Input( io.Float.Input(
"factor", "factor",
@ -935,6 +1098,261 @@ class ShuffleImageTextDatasetNode(io.ComfyNode):
return io.NodeOutput(shuffled_images, shuffled_texts) return io.NodeOutput(shuffled_images, shuffled_texts)
# ========== Video Processing Nodes ==========
class VideoFrameSampleNode(io.ComfyNode):
"""Sample a fixed number of frames from a video using various strategies.
For contiguous strategies ("head"/"tail") the result is a fully lazy
VideoInput (no frames decoded). For non-contiguous strategies
("uniform"/"random") only the selected indices are decoded.
"""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="VideoFrameSample",
search_aliases=["sample frames", "extract frames"],
display_name="Sample Video Frame",
category="video",
description="Sample a fixed number of frames from a video using various strategies.",
is_experimental=True,
inputs=[
io.Video.Input("video", tooltip="Input video."),
io.Int.Input(
"num_frames",
default=16,
min=1,
max=9999,
tooltip="Number of frames to sample.",
),
io.Combo.Input(
"strategy",
options=["uniform", "head", "tail", "random"],
default="uniform",
tooltip="uniform: evenly spaced, head: first N, tail: last N, random: random sorted.",
),
io.Int.Input(
"seed",
default=0,
min=0,
max=0xFFFFFFFFFFFFFFFF,
tooltip="Random seed (only used with 'random' strategy).",
),
],
outputs=[
io.Video.Output(display_name="video", tooltip="Sampled video."),
],
)
@classmethod
def execute(cls, video, num_frames, strategy, seed):
total_frames = video.get_frame_count()
num_frames = min(num_frames, total_frames)
fps = float(video.get_frame_rate())
if strategy == "head":
return io.NodeOutput(
video.as_trimmed(0.0, num_frames / fps, strict_duration=False)
)
if strategy == "tail":
start_t = (total_frames - num_frames) / fps
return io.NodeOutput(
video.as_trimmed(start_t, num_frames / fps, strict_duration=False)
)
if strategy == "uniform":
if num_frames == 1:
indices = [total_frames // 2]
else:
indices = [round(i * (total_frames - 1) / (num_frames - 1)) for i in range(num_frames)]
elif strategy == "random":
rng = np.random.RandomState(seed % (2**32 - 1))
indices = sorted(rng.choice(total_frames, size=num_frames, replace=False).tolist())
else:
raise ValueError(f"Unknown strategy: {strategy}")
return io.NodeOutput(_decode_selected_frames(video, indices))
class VideoTemporalCropNode(io.ComfyNode):
"""Crop a continuous range of frames from a video (fully lazy)."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="VideoTemporalCrop",
search_aliases=["crop", "crop video", "temporal crop", "truncate video"],
display_name="Crop Video (Temporal)",
category="video/transform",
description="Crop a continuous range of frames from a video.",
is_experimental=True,
inputs=[
io.Video.Input("video", tooltip="Input video."),
io.Int.Input(
"start_frame",
default=0,
min=0,
max=99999,
tooltip="Starting frame index.",
),
io.Int.Input(
"length",
default=16,
min=1,
max=99999,
tooltip="Number of frames to keep.",
),
],
outputs=[
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
],
)
@classmethod
def execute(cls, video, start_frame, length):
total_frames = video.get_frame_count()
fps = float(video.get_frame_rate())
start_frame = min(start_frame, max(total_frames - 1, 0))
length = min(length, total_frames - start_frame)
return io.NodeOutput(
video.as_trimmed(start_frame / fps, length / fps, strict_duration=False)
)
class VideoRandomTemporalCropNode(io.ComfyNode):
"""Randomly crop a continuous range of frames from a video (fully lazy)."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="VideoRandomTemporalCrop",
search_aliases=["crop", "crop video", "temporal crop", "truncate video", "random crop"],
display_name="Crop Video (Temporal Random)",
category="video/transform",
description="Randomly crop a continuous range of frames from a video.",
is_experimental=True,
inputs=[
io.Video.Input("video", tooltip="Input video."),
io.Int.Input(
"length",
default=16,
min=1,
max=99999,
tooltip="Number of frames to keep.",
),
io.Int.Input(
"seed",
default=0,
min=0,
max=0xFFFFFFFFFFFFFFFF,
tooltip="Random seed.",
),
],
outputs=[
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
],
)
@classmethod
def execute(cls, video, length, seed):
total_frames = video.get_frame_count()
fps = float(video.get_frame_rate())
length = min(length, total_frames)
max_start = total_frames - length
rng = np.random.RandomState(seed % (2**32 - 1))
start = rng.randint(0, max_start + 1) if max_start > 0 else 0
return io.NodeOutput(
video.as_trimmed(start / fps, length / fps, strict_duration=False)
)
class ShuffleVideoDatasetNode(io.ComfyNode):
"""Randomly shuffle the order of videos in the dataset."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ShuffleVideoDataset",
search_aliases=["shuffle", "randomize", "mix"],
display_name="Shuffle Videos List",
category="video/batch",
description="Randomly shuffle the order of videos in a list.",
is_experimental=True,
is_input_list=True,
inputs=[
io.Video.Input("videos", tooltip="List of videos to shuffle."),
io.Int.Input(
"seed", default=0, min=0, max=0xFFFFFFFFFFFFFFFF, tooltip="Random seed."
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Shuffled videos",
),
],
)
@classmethod
def execute(cls, videos, seed):
seed = seed[0] if isinstance(seed, list) else seed
np.random.seed(seed % (2**32 - 1))
indices = np.random.permutation(len(videos))
return io.NodeOutput([videos[i] for i in indices])
class ShuffleVideoTextDatasetNode(io.ComfyNode):
"""Shuffle videos and their captions together, preserving pairs."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ShuffleVideoTextDataset",
search_aliases=["shuffle", "randomize", "mix"],
display_name="Shuffle Pairs of Video-Text",
category="dataset/video",
description="Randomly shuffle the order of pairs of video-text in a list.",
is_experimental=True,
is_input_list=True,
inputs=[
io.Video.Input("videos", tooltip="List of videos to shuffle."),
io.String.Input("texts", tooltip="List of texts to shuffle."),
io.Int.Input(
"seed",
default=0,
min=0,
max=0xFFFFFFFFFFFFFFFF,
tooltip="Random seed.",
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Shuffled videos",
),
io.String.Output(
display_name="texts",
is_output_list=True,
tooltip="Shuffled texts",
),
],
)
@classmethod
def execute(cls, videos, texts, seed):
seed = seed[0] if isinstance(seed, list) else seed
np.random.seed(seed % (2**32 - 1))
indices = np.random.permutation(len(videos))
return io.NodeOutput(
[videos[i] for i in indices],
[texts[i] for i in indices],
)
# ========== Text Transform Nodes ========== # ========== Text Transform Nodes ==========
@ -1608,7 +2026,10 @@ class DatasetExtension(ComfyExtension):
LoadImageTextDataSetFromFolderNode, LoadImageTextDataSetFromFolderNode,
SaveImageDataSetToFolderNode, SaveImageDataSetToFolderNode,
SaveImageTextDataSetToFolderNode, SaveImageTextDataSetToFolderNode,
# Image transform nodes # Video data loading nodes
LoadVideoDataSetFromFolderNode,
LoadVideoTextDataSetFromFolderNode,
# Image transform nodes (auto-handle video via per-frame processing)
ResizeImagesByShorterEdgeNode, ResizeImagesByShorterEdgeNode,
ResizeImagesByLongerEdgeNode, ResizeImagesByLongerEdgeNode,
CenterCropImagesNode, CenterCropImagesNode,
@ -1618,6 +2039,12 @@ class DatasetExtension(ComfyExtension):
AdjustContrastNode, AdjustContrastNode,
ShuffleDatasetNode, ShuffleDatasetNode,
ShuffleImageTextDatasetNode, ShuffleImageTextDatasetNode,
# Video processing nodes (lazy VideoInput in/out)
VideoFrameSampleNode,
VideoTemporalCropNode,
VideoRandomTemporalCropNode,
ShuffleVideoDatasetNode,
ShuffleVideoTextDatasetNode,
# Text transform nodes # Text transform nodes
TextToLowercaseNode, TextToLowercaseNode,
TextToUppercaseNode, TextToUppercaseNode,

View File

@ -10,7 +10,7 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
Apply frequency-dependent scaling to an image tensor using Fourier transforms. Apply frequency-dependent scaling to an image tensor using Fourier transforms.
Parameters: Parameters:
x: Input tensor of shape (B, C, H, W) x: Input tensor of shape (..., H, W)
scale_low: Scaling factor for low-frequency components (default: 1.0) scale_low: Scaling factor for low-frequency components (default: 1.0)
scale_high: Scaling factor for high-frequency components (default: 1.5) scale_high: Scaling factor for high-frequency components (default: 1.5)
freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20) freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20)
@ -31,8 +31,8 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
# Initialize mask with high-frequency scaling factor # Initialize mask with high-frequency scaling factor
mask = torch.ones(x_freq.shape, device=device) * scale_high mask = torch.ones(x_freq.shape, device=device) * scale_high
m = mask m = mask
for d in range(len(x_freq.shape) - 2): for d in range(2):
dim = d + 2 dim = len(x_freq.shape) - 2 + d
cc = x_freq.shape[dim] // 2 cc = x_freq.shape[dim] // 2
f_c = min(freq_cutoff, cc) f_c = min(freq_cutoff, cc)
m = m.narrow(dim, cc - f_c, f_c * 2) m = m.narrow(dim, cc - f_c, f_c * 2)

View File

@ -2,6 +2,8 @@ import nodes
import node_helpers import node_helpers
import torch import torch
import comfy.model_management import comfy.model_management
import comfy.model_patcher
import comfy.ops
from typing_extensions import override from typing_extensions import override
from comfy_api.latest import ComfyExtension, io from comfy_api.latest import ComfyExtension, io
from comfy.ldm.hunyuan_video.upsampler import HunyuanVideo15SRModel from comfy.ldm.hunyuan_video.upsampler import HunyuanVideo15SRModel
@ -217,8 +219,11 @@ class LatentUpscaleModelLoader(io.ComfyNode):
model.load_sd(sd) model.load_sd(sd)
elif "post_upsample_res_blocks.0.conv2.bias" in sd: elif "post_upsample_res_blocks.0.conv2.bias" in sd:
config = json.loads(metadata["config"]) config = json.loads(metadata["config"])
model = LatentUpsampler.from_config(config).to(dtype=comfy.model_management.vae_dtype(allowed_dtypes=[torch.bfloat16, torch.float32])) model = LatentUpsampler.from_config(config, operations=comfy.ops.disable_weight_init).to(dtype=comfy.model_management.vae_dtype(allowed_dtypes=[torch.bfloat16, torch.float32]))
model.load_state_dict(sd) comfy.model_management.archive_model_dtypes(model)
model_patcher = comfy.model_patcher.CoreModelPatcher(model, load_device=comfy.model_management.get_torch_device(), offload_device=comfy.model_management.unet_offload_device())
model.load_state_dict(sd, assign=model_patcher.is_dynamic())
model = model_patcher
return io.NodeOutput(model) return io.NodeOutput(model)

View File

@ -0,0 +1,102 @@
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()

View File

@ -50,8 +50,8 @@ class GetICLoRAParameters(io.ComfyNode):
factor = 1 factor = 1
if metadata: if metadata:
try: try:
factor = max(1, round(float(metadata.get("reference_downscale_factor", 1)))) factor = max(1, round(float(next(v for k, v in metadata.items() if k.endswith("reference_downscale_factor")))))
except (TypeError, ValueError): except (StopIteration, TypeError, ValueError):
factor = 1 factor = 1
parameters = {"reference_downscale_factor": factor} parameters = {"reference_downscale_factor": factor}
return io.NodeOutput(parameters) return io.NodeOutput(parameters)

View File

@ -38,26 +38,20 @@ class LTXVLatentUpsampler(IO.ComfyNode):
Returns: Returns:
tuple: Tuple containing the upsampled latent tuple: Tuple containing the upsampled latent
""" """
device = model_management.get_torch_device() device = upscale_model.load_device
memory_required = model_management.module_size(upscale_model) model = upscale_model.model
model_dtype = upscale_model.model_dtype()
model_dtype = next(upscale_model.parameters()).dtype
latents = samples["samples"] latents = samples["samples"]
input_dtype = latents.dtype input_dtype = latents.dtype
memory_required += math.prod(latents.shape) * 3000.0 # TODO: more accurate memory_required = math.prod(latents.shape) * 3000.0 # TODO: more accurate
model_management.free_memory(memory_required, device) model_management.load_models_gpu([upscale_model], memory_required=memory_required)
try: latents = latents.to(dtype=model_dtype, device=device)
upscale_model.to(device) # TODO: use the comfy model management system.
latents = latents.to(dtype=model_dtype, device=device) """Upsample latents without tiling."""
latents = vae.first_stage_model.per_channel_statistics.un_normalize(latents)
"""Upsample latents without tiling.""" upsampled_latents = model(latents)
latents = vae.first_stage_model.per_channel_statistics.un_normalize(latents)
upsampled_latents = upscale_model(latents)
finally:
upscale_model.cpu()
upsampled_latents = vae.first_stage_model.per_channel_statistics.normalize( upsampled_latents = vae.first_stage_model.per_channel_statistics.normalize(
upsampled_latents upsampled_latents

103
comfy_extras/nodes_mage.py Normal file
View File

@ -0,0 +1,103 @@
from typing_extensions import override
import comfy.utils
import node_helpers
import torch
import comfy.model_management
from comfy_api.latest import ComfyExtension, io
class TextEncodeMageFlowEdit(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="TextEncodeMageFlowEdit",
category="model/conditioning/mage",
description="Encode an edit instruction with one or more reference images for Mage-Flow-Edit. Reference latents are resized to the output resolution (width/height, or the first image's size when 0). Use the latent output for sampling so the sizes always match.",
inputs=[
io.Clip.Input("clip"),
io.String.Input("prompt", multiline=True, dynamic_prompts=True),
io.String.Input("negative_prompt", multiline=True, dynamic_prompts=True, advanced=True),
io.Vae.Input("vae", optional=True),
io.Autogrow.Input(
"images",
template=io.Autogrow.TemplateNames(
io.Image.Input("image"),
names=[f"image_{i}" for i in range(1, 17)],
min=0,
),
tooltip="Reference image(s) to edit. All references are resized to the output resolution before encoding.",
),
io.Int.Input("width", default=0, min=0, max=8192, step=16, tooltip="Output width. 0 = use the first reference image's size."),
io.Int.Input("height", default=0, min=0, max=8192, step=16, tooltip="Output height. 0 = use the first reference image's size."),
io.Int.Input("batch_size", default=1, min=1, max=4096),
],
outputs=[
io.Conditioning.Output(display_name="positive"),
io.Conditioning.Output(display_name="negative"),
io.Latent.Output(display_name="latent"),
],
)
@classmethod
def execute(cls, clip, prompt, negative_prompt="", vae=None, images: io.Autogrow.Type = None, width=0, height=0, batch_size=1) -> io.NodeOutput:
ref_latents = []
images = images or {}
images = [images[name] for name in sorted(images, key=lambda n: int(n.rsplit("_", 1)[-1])) if images[name] is not None]
images_vl = []
# Output resolution: explicit width/height, else the primary reference's own size, floored to /16.
# Each dimension falls back independently so a 0 on one axis keeps an explicit value on the other.
if width == 0 or height == 0:
if len(images) > 0:
ref_h, ref_w = images[0].shape[1], images[0].shape[2]
else:
ref_h, ref_w = 1024, 1024
height = height or ref_h
width = width or ref_w
width = max(16, (width // 16) * 16)
height = max(16, (height // 16) * 16)
for image in images:
samples = image.movedim(-1, 1)
# VL conditioning copy: cap the long edge at 384 (training preprocessing).
long_edge = max(samples.shape[3], samples.shape[2])
if long_edge > 384:
scale_by = 384 / long_edge
s = comfy.utils.common_upscale(samples, max(1, round(samples.shape[3] * scale_by)), max(1, round(samples.shape[2] * scale_by)), "bicubic", "disabled")
images_vl.append(s.movedim(1, -1))
else:
images_vl.append(image)
if vae is not None:
# All references are resized to the output resolution before encoding, because Mage's RoPE aligns reference and target content by position
if samples.shape[3] != width or samples.shape[2] != height:
s = comfy.utils.common_upscale(samples, width, height, "bicubic", "disabled")
else:
s = samples
ref_latents.append(vae.encode(s.movedim(1, -1)[:, :, :, :3]))
# Negative branch keeps the same reference images (VL tokens + ref latents), only the instruction differs.
positive = clip.encode_from_tokens_scheduled(clip.tokenize(prompt, images=images_vl))
negative = clip.encode_from_tokens_scheduled(clip.tokenize(negative_prompt if negative_prompt else " ", images=images_vl))
if len(ref_latents) > 0:
positive = node_helpers.conditioning_set_values(positive, {"reference_latents": ref_latents}, append=True)
negative = node_helpers.conditioning_set_values(negative, {"reference_latents": ref_latents}, append=True)
latent = torch.zeros([batch_size, 128, height // 16, width // 16], device=comfy.model_management.intermediate_device())
return io.NodeOutput(positive, negative, {"samples": latent})
class MageExtension(ComfyExtension):
@override
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [
TextEncodeMageFlowEdit,
]
async def comfy_entrypoint() -> MageExtension:
return MageExtension()

View File

@ -8,6 +8,8 @@ import comfy.ldm.common_dit
import comfy.latent_formats import comfy.latent_formats
import comfy.ldm.lumina.controlnet import comfy.ldm.lumina.controlnet
import comfy.ldm.supir.supir_modules import comfy.ldm.supir.supir_modules
import comfy.ldm.anima.lllite
import comfy.ldm.wan.uni3c
from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel
from comfy_api.latest import io from comfy_api.latest import io
from comfy.ldm.supir.supir_patch import SUPIRPatch from comfy.ldm.supir.supir_patch import SUPIRPatch
@ -236,10 +238,12 @@ class ModelPatchLoader:
def load_model_patch(self, name): def load_model_patch(self, name):
model_patch_path = folder_paths.get_full_path_or_raise("model_patches", name) model_patch_path = folder_paths.get_full_path_or_raise("model_patches", name)
sd = comfy.utils.load_torch_file(model_patch_path, safe_load=True) sd, metadata = comfy.utils.load_torch_file(model_patch_path, safe_load=True, return_metadata=True)
dtype = comfy.utils.weight_dtype(sd) dtype = comfy.utils.weight_dtype(sd)
if 'controlnet_blocks.0.y_rms.weight' in sd: if 'lllite_conditioning1.conv1.weight' in sd:
model = comfy.ldm.anima.lllite.AnimaLLLite(sd, metadata, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast)
elif 'controlnet_blocks.0.y_rms.weight' in sd:
additional_in_dim = sd["img_in.weight"].shape[1] - 64 additional_in_dim = sd["img_in.weight"].shape[1] - 64
model = QwenImageBlockWiseControlNet(additional_in_dim=additional_in_dim, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast) model = QwenImageBlockWiseControlNet(additional_in_dim=additional_in_dim, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast)
elif 'feature_embedder.mid_layer_norm.bias' in sd: elif 'feature_embedder.mid_layer_norm.bias' in sd:
@ -261,6 +265,37 @@ class ModelPatchLoader:
if torch.count_nonzero(ref_weight) == 0: if torch.count_nonzero(ref_weight) == 0:
config['broken'] = True config['broken'] = True
model = comfy.ldm.lumina.controlnet.ZImage_Control(device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast, **config) model = comfy.ldm.lumina.controlnet.ZImage_Control(device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast, **config)
elif 'controlnet_patch_embedding.weight' in sd: # Uni3C controlnet for Wan
attn_key_replace = {".self_attn.to_q.": ".self_attn.q.",
".self_attn.to_k.": ".self_attn.k.",
".self_attn.to_v.": ".self_attn.v.",
".self_attn.to_out.0.": ".self_attn.o."}
converted_sd = {}
for k, w in sd.items():
for r, rr in attn_key_replace.items():
k = k.replace(r, rr)
converted_sd[k] = w
sd = converted_sd
num_layers = sum(1 for k in sd if k.startswith("proj_out.") and k.endswith(".weight"))
conv_out_dim = sd["controlnet_patch_embedding.weight"].shape[0]
if "proj_in.weight" in sd:
dim = sd["proj_in.weight"].shape[0]
else:
dim = conv_out_dim
model = comfy.ldm.wan.uni3c.WanUni3CControlnet(
in_channels=sd["controlnet_patch_embedding.weight"].shape[1],
conv_out_dim=conv_out_dim,
dim=dim,
ffn_dim=sd["controlnet_blocks.0.ffn.0.bias"].shape[0],
num_layers=num_layers,
time_embed_dim=sd["controlnet_blocks.0.norm1.linear.weight"].shape[1],
out_proj_dim=sd["proj_out.0.weight"].shape[0],
add_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[1],
mid_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[0],
device=comfy.model_management.unet_offload_device(),
dtype=dtype,
operations=comfy.ops.manual_cast)
elif "audio_proj.proj1.weight" in sd: elif "audio_proj.proj1.weight" in sd:
model = MultiTalkModelPatch( model = MultiTalkModelPatch(
audio_window=5, context_tokens=32, vae_scale=4, audio_window=5, context_tokens=32, vae_scale=4,
@ -296,6 +331,50 @@ class ModelPatchLoader:
return (model_patcher,) return (model_patcher,)
class AnimaLLLiteApply:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"model_patch": ("MODEL_PATCH",),
"image": ("IMAGE",),
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
},
"optional": {"mask": ("MASK",),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "apply_patch"
EXPERIMENTAL = True
CATEGORY = "model_patches/anima"
def apply_patch(self, model, model_patch, image, strength, start_percent, end_percent, mask=None):
image = image[..., :3]
if model_patch.model.cond_in_channels == 4 and mask is None:
mask = torch.zeros_like(image[..., 0])
elif model_patch.model.cond_in_channels != 4:
mask = None
model_sampling = model.get_model_object("model_sampling")
sigma_start = float(model_sampling.percent_to_sigma(start_percent))
sigma_end = float(model_sampling.percent_to_sigma(end_percent))
patch = comfy.ldm.anima.lllite.AnimaLLLitePatch(model_patch, image, mask, strength, sigma_start, sigma_end)
model_patched = model.clone()
model_patched.set_model_post_input_patch(patch)
model_patched.set_model_attn1_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch(
patch,
{"q": "self_attn_q_proj", "k": "self_attn_k_proj", "v": "self_attn_v_proj"},
))
model_patched.set_model_attn2_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch(
patch,
{"q": "cross_attn_q_proj"},
))
model_patched.set_model_patch(comfy.ldm.anima.lllite.AnimaLLLiteMLPPatch(patch), "mlp_patch")
return (model_patched,)
class DiffSynthCnetPatch: class DiffSynthCnetPatch:
def __init__(self, model_patch, vae, image, strength, mask=None): def __init__(self, model_patch, vae, image, strength, mask=None):
self.model_patch = model_patch self.model_patch = model_patch
@ -514,6 +593,150 @@ class ZImageFunControlnet(QwenImageDiffsynthControlnet):
CATEGORY = "model/patch/z-image" CATEGORY = "model/patch/z-image"
class WanUni3CCnetPatch:
def __init__(self, model_patch, render_video, vae, latent_format, strength, sigma_start, sigma_end):
self.model_patch = model_patch
self.render_video = render_video
self.vae = vae
self.latent_format = latent_format
self.strength = strength
self.sigma_start = sigma_start
self.sigma_end = sigma_end
self.prepared_render = None
self.temp_data = None
def encode_render_video(self, target_latent_shape):
t_len, h_len, w_len = target_latent_shape
temporal_compression = self.vae.temporal_compression_decode() or 1
spatial_compression = self.vae.spacial_compression_encode()
target_frames = (t_len - 1) * temporal_compression + 1
target_height = h_len * spatial_compression
target_width = w_len * spatial_compression
frames = self.render_video
if frames.shape[0] > target_frames:
frames = frames[:target_frames]
elif frames.shape[0] < target_frames:
last_frame = frames[-1:].expand(target_frames - frames.shape[0], -1, -1, -1)
frames = torch.cat([frames, last_frame], dim=0)
if frames.shape[1] != target_height or frames.shape[2] != target_width:
frames = comfy.utils.common_upscale(frames.movedim(-1, 1), target_width, target_height, "bilinear", "center").movedim(1, -1)
loaded_models = comfy.model_management.loaded_models(only_currently_used=True)
render_latent = self.vae.encode(frames)
comfy.model_management.load_models_gpu(loaded_models)
return self.latent_format.process_in(render_latent)
def build_controlnet_input(self, x, dtype, samples_per_cond):
# first 20 channels of the model input: noise latent + I2V mask (zero padded for T2V)
hidden = x[:samples_per_cond, :20].to(dtype)
if hidden.shape[1] < 20:
pad_shape = list(hidden.shape)
pad_shape[1] = 20 - hidden.shape[1]
hidden = torch.cat([hidden, torch.zeros(pad_shape, dtype=hidden.dtype, device=hidden.device)], dim=1)
render = self.prepared_render
if render is None or render.shape[2:] != hidden.shape[2:]:
render = self.encode_render_video(hidden.shape[2:])
render = render.to(device=hidden.device, dtype=dtype)
self.prepared_render = render
if render.shape[0] != hidden.shape[0]:
render = render.expand(hidden.shape[0], -1, -1, -1, -1)
return torch.cat([hidden, render], dim=1)
def __call__(self, kwargs):
img = kwargs.get("img")
block_index = kwargs.get("block_index")
transformer_options = kwargs.get("transformer_options", {})
if block_index == 0:
self.temp_data = None
active = True
sigmas = transformer_options.get("sigmas", None)
if sigmas is not None:
sigma = sigmas[0].item()
if sigma > self.sigma_start or sigma < self.sigma_end:
active = False
if active:
x = kwargs.get("x")
# cond and uncond chunks share latents, so we can reuse residuals
num_conds = len(transformer_options.get("cond_or_uncond", [0]))
samples_per_cond = x.shape[0]
if num_conds > 0 and x.shape[0] % num_conds == 0:
samples_per_cond = x.shape[0] // num_conds
temb = kwargs.get("vec")[:samples_per_cond]
if temb.ndim == 3:
temb = temb[:, 0]
model = self.model_patch.model
controlnet_input = self.build_controlnet_input(x, img.dtype, samples_per_cond)
hidden, freqs = model.process_input(controlnet_input)
self.temp_data = (hidden, temb.to(img.dtype), freqs)
num_layers = self.model_patch.model.num_layers
if self.temp_data is not None and block_index < num_layers:
hidden, temb, freqs = self.temp_data
hidden, residual = self.model_patch.model.forward_block(block_index, hidden, temb, freqs)
residual = residual.to(img.dtype) * self.strength
if residual.shape[0] != img.shape[0]:
residual = residual.repeat(img.shape[0] // residual.shape[0], 1, 1)
img_offset = kwargs.get("img_offset", 0)
img[:, img_offset:img_offset + residual.shape[1]] += residual
if block_index >= num_layers - 1:
self.temp_data = None
else:
self.temp_data = (hidden, temb, freqs)
return kwargs
def to(self, device_or_dtype):
if isinstance(device_or_dtype, torch.device):
if self.prepared_render is not None:
self.prepared_render = self.prepared_render.to(device_or_dtype)
self.temp_data = None
return self
def models(self):
return [self.model_patch]
class WanUni3CControlnetApply:
@classmethod
def INPUT_TYPES(s):
return {"required": { "model": ("MODEL",),
"model_patch": ("MODEL_PATCH",),
"vae": ("VAE",),
"render_video": ("IMAGE", {"tooltip": "The guidance video rendered from the camera trajectory, most commonly warped point cloud renders of the input image."}),
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "apply_patch"
EXPERIMENTAL = True
CATEGORY = "model/patch/wan"
def apply_patch(self, model, model_patch, vae, render_video, strength, start_percent, end_percent):
if not isinstance(model_patch.model, comfy.ldm.wan.uni3c.WanUni3CControlnet):
raise ValueError("The connected model patch is not a Uni3C ControlNet.")
cnet_dim = model_patch.model.controlnet_blocks[0].norm1.linear.in_features
model_dim = getattr(model.get_model_object("diffusion_model"), "dim", None)
if model_dim is None:
raise ValueError("The Uni3C ControlNet only works with Wan models.")
if model_dim != cnet_dim:
raise ValueError("This Uni3C ControlNet expects a Wan model with dim {}, the loaded model has dim {}.".format(cnet_dim, model_dim))
model_patched = model.clone()
model_sampling = model.get_model_object("model_sampling")
sigma_start = model_sampling.percent_to_sigma(start_percent)
sigma_end = model_sampling.percent_to_sigma(end_percent)
latent_format = model.get_model_object("latent_format")
patch = WanUni3CCnetPatch(model_patch, render_video[:, :, :, :3], vae, latent_format, strength, sigma_start, sigma_end)
model_patched.set_model_double_block_patch(patch)
return (model_patched,)
class UsoStyleProjectorPatch: class UsoStyleProjectorPatch:
def __init__(self, model_patch, encoded_image): def __init__(self, model_patch, encoded_image):
self.model_patch = model_patch self.model_patch = model_patch
@ -672,14 +895,18 @@ NODE_CLASS_MAPPINGS = {
"ModelPatchLoader": ModelPatchLoader, "ModelPatchLoader": ModelPatchLoader,
"QwenImageDiffsynthControlnet": QwenImageDiffsynthControlnet, "QwenImageDiffsynthControlnet": QwenImageDiffsynthControlnet,
"ZImageFunControlnet": ZImageFunControlnet, "ZImageFunControlnet": ZImageFunControlnet,
"WanUni3CControlnetApply": WanUni3CControlnetApply,
"USOStyleReference": USOStyleReference, "USOStyleReference": USOStyleReference,
"SUPIRApply": SUPIRApply, "SUPIRApply": SUPIRApply,
"AnimaLLLiteApply": AnimaLLLiteApply,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"ModelPatchLoader": "Load Model Patch", "ModelPatchLoader": "Load Model Patch",
"QwenImageDiffsynthControlnet": "Apply Qwen Image DiffSynth ControlNet", "QwenImageDiffsynthControlnet": "Apply Qwen Image DiffSynth ControlNet",
"ZImageFunControlnet": "Apply Z-Image Fun ControlNet", "ZImageFunControlnet": "Apply Z-Image Fun ControlNet",
"WanUni3CControlnetApply": "Apply Wan Uni3C ControlNet",
"USOStyleReference": "Apply USO Style Reference", "USOStyleReference": "Apply USO Style Reference",
"SUPIRApply": "Apply SUPIR Patch", "SUPIRApply": "Apply SUPIR Patch",
"AnimaLLLiteApply": "Apply Anima LLLite",
} }

View File

@ -920,10 +920,11 @@ def _run_training_loop(
""" """
sigmas = torch.tensor(range(num_images)) sigmas = torch.tensor(range(num_images))
noise = comfy_extras.nodes_custom_sampler.Noise_RandomNoise(seed) noise = comfy_extras.nodes_custom_sampler.Noise_RandomNoise(seed)
ndim = latents[0].ndim
if bucket_mode: if bucket_mode:
# Use first bucket's first latent as dummy for guider # Use first bucket's first latent as dummy for guider
dummy_latent = latents[0][:1].repeat(num_images, 1, 1, 1) dummy_latent = latents[0][:1].repeat(num_images, *[1]*(ndim-1))
guider.sample( guider.sample(
noise.generate_noise({"samples": dummy_latent}), noise.generate_noise({"samples": dummy_latent}),
dummy_latent, dummy_latent,
@ -933,7 +934,7 @@ def _run_training_loop(
) )
elif multi_res: elif multi_res:
# use first latent as dummy latent if multi_res # use first latent as dummy latent if multi_res
latents = latents[0].repeat(num_images, 1, 1, 1) latents = latents[0].repeat(num_images, *[1]*(ndim-1))
guider.sample( guider.sample(
noise.generate_noise({"samples": latents}), noise.generate_noise({"samples": latents}),
latents, latents,

View File

@ -7,6 +7,7 @@ import folder_paths
from typing_extensions import override from typing_extensions import override
from comfy_api.latest import ComfyExtension, io from comfy_api.latest import ComfyExtension, io
import comfy.model_management import comfy.model_management
import comfy.model_patcher
try: try:
from spandrel_extra_arches import EXTRA_REGISTRY from spandrel_extra_arches import EXTRA_REGISTRY
@ -42,6 +43,7 @@ class UpscaleModelLoader(io.ComfyNode):
if not isinstance(out, ImageModelDescriptor): if not isinstance(out, ImageModelDescriptor):
raise Exception("Upscale model must be a single-image model.") raise Exception("Upscale model must be a single-image model.")
out.patcher = comfy.model_patcher.CoreModelPatcher(out.model, load_device=model_management.get_torch_device(), offload_device=model_management.unet_offload_device())
return io.NodeOutput(out) return io.NodeOutput(out)
load_model = execute # TODO: remove load_model = execute # TODO: remove
@ -66,14 +68,12 @@ class ImageUpscaleWithModel(io.ComfyNode):
@classmethod @classmethod
def execute(cls, upscale_model, image) -> io.NodeOutput: def execute(cls, upscale_model, image) -> io.NodeOutput:
device = model_management.get_torch_device() device = upscale_model.patcher.load_device
memory_required = model_management.module_size(upscale_model.model) memory_required = (512 * 512 * 3) * image.element_size() * max(upscale_model.scale, 1.0) * 384.0 #The 384.0 is an estimate of how much some of these models take, TODO: make it more accurate
memory_required += (512 * 512 * 3) * image.element_size() * max(upscale_model.scale, 1.0) * 384.0 #The 384.0 is an estimate of how much some of these models take, TODO: make it more accurate
memory_required += image.nelement() * image.element_size() memory_required += image.nelement() * image.element_size()
model_management.free_memory(memory_required, device) model_management.load_models_gpu([upscale_model.patcher], memory_required=memory_required)
upscale_model.to(device)
in_img = image.movedim(-1,-3).to(device) in_img = image.movedim(-1,-3).to(device)
tile = 512 tile = 512
@ -82,20 +82,17 @@ class ImageUpscaleWithModel(io.ComfyNode):
output_device = comfy.model_management.intermediate_device() output_device = comfy.model_management.intermediate_device()
oom = True oom = True
try: while oom:
while oom: try:
try: steps = in_img.shape[0] * comfy.utils.get_tiled_scale_steps(in_img.shape[3], in_img.shape[2], tile_x=tile, tile_y=tile, overlap=overlap)
steps = in_img.shape[0] * comfy.utils.get_tiled_scale_steps(in_img.shape[3], in_img.shape[2], tile_x=tile, tile_y=tile, overlap=overlap) pbar = comfy.utils.ProgressBar(steps)
pbar = comfy.utils.ProgressBar(steps) s = comfy.utils.tiled_scale(in_img, lambda a: upscale_model(a.float()), tile_x=tile, tile_y=tile, overlap=overlap, upscale_amount=upscale_model.scale, pbar=pbar, output_device=output_device)
s = comfy.utils.tiled_scale(in_img, lambda a: upscale_model(a.float()), tile_x=tile, tile_y=tile, overlap=overlap, upscale_amount=upscale_model.scale, pbar=pbar, output_device=output_device) oom = False
oom = False except Exception as e:
except Exception as e: model_management.raise_non_oom(e)
model_management.raise_non_oom(e) tile //= 2
tile //= 2 if tile < 128:
if tile < 128: raise e
raise e
finally:
upscale_model.to("cpu")
s = torch.clamp(s.movedim(-3,-1), min=0, max=1.0).to(comfy.model_management.intermediate_dtype()) s = torch.clamp(s.movedim(-3,-1), min=0, max=1.0).to(comfy.model_management.intermediate_dtype())
return io.NodeOutput(s) return io.NodeOutput(s)

View File

@ -1,3 +1,3 @@
# This file is automatically generated by the build process when version is # This file is automatically generated by the build process when version is
# updated in pyproject.toml. # updated in pyproject.toml.
__version__ = "0.28.2" __version__ = "0.28.0"

View File

@ -319,7 +319,7 @@ def prompt_worker(q, server_instance):
cache_ram_inactive = 0 cache_ram_inactive = 0
if not args.cache_classic and not args.cache_none and args.cache_lru <= 0: if not args.cache_classic and not args.cache_none and args.cache_lru <= 0:
cache_ram = min(10.0, max(2.0, comfy.model_management.total_ram * 0.10 / 1024.0)) cache_ram = min(10.0, max(2.0, comfy.model_management.total_ram * 0.10 / 1024.0))
cache_ram_inactive = min(96.0, comfy.model_management.total_ram / 1024.0) cache_ram_inactive = min(128.0, comfy.model_management.total_ram / 1024.0)
if len(args.cache_ram) > 0: if len(args.cache_ram) > 0:
cache_ram = args.cache_ram[0] cache_ram = args.cache_ram[0]
if len(args.cache_ram) > 1: if len(args.cache_ram) > 1:

View File

@ -992,7 +992,7 @@ class CLIPLoader:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": { "clip_name": (folder_paths.get_filename_list("text_encoders"), ), 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"], ), "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", "mage"], ),
}, },
"optional": { "optional": {
"device": (["default", "cpu"], {"advanced": True}), "device": (["default", "cpu"], {"advanced": True}),
@ -1002,7 +1002,7 @@ class CLIPLoader:
CATEGORY = "model/loaders" 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\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\njoyimage: qwen3-vl 8B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm"
def load_clip(self, clip_name, type="stable_diffusion", device="default"): def load_clip(self, clip_name, type="stable_diffusion", device="default"):
clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION) clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION)
@ -2462,6 +2462,8 @@ async def init_builtin_extra_nodes():
"nodes_seedvr.py", "nodes_seedvr.py",
"nodes_context_windows.py", "nodes_context_windows.py",
"nodes_qwen.py", "nodes_qwen.py",
"nodes_mage.py",
"nodes_joyimage.py",
"nodes_boogu.py", "nodes_boogu.py",
"nodes_chroma_radiance.py", "nodes_chroma_radiance.py",
"nodes_pid.py", "nodes_pid.py",

View File

@ -530,6 +530,10 @@ components:
description: Job creation timestamp (Unix timestamp in milliseconds) description: Job creation timestamp (Unix timestamp in milliseconds)
format: int64 format: int64
type: integer type: integer
execution_end_time:
description: Workflow execution completion timestamp (Unix milliseconds, only present for terminal states)
format: int64
type: integer
execution_error: execution_error:
allOf: allOf:
- $ref: '#/components/schemas/ExecutionError' - $ref: '#/components/schemas/ExecutionError'
@ -538,6 +542,10 @@ components:
additionalProperties: true additionalProperties: true
description: Node-level execution metadata (only for terminal states) description: Node-level execution metadata (only for terminal states)
type: object type: object
execution_start_time:
description: Workflow execution start timestamp (Unix milliseconds, only present once execution has started)
format: int64
type: integer
execution_status: execution_status:
additionalProperties: true additionalProperties: true
description: ComfyUI execution status and timeline (only for terminal states) description: ComfyUI execution status and timeline (only for terminal states)
@ -570,6 +578,12 @@ components:
description: Last update timestamp (Unix timestamp in milliseconds) description: Last update timestamp (Unix timestamp in milliseconds)
format: int64 format: int64
type: integer type: integer
user_id:
description: |
ID of the user that owns this job (see the `workspace_id`
description above for why this is always the caller's own id
on a successful response).
type: string
workflow: workflow:
additionalProperties: true additionalProperties: true
description: | description: |
@ -583,6 +597,18 @@ components:
workflow_id: workflow_id:
description: UUID identifying the workflow graph definition description: UUID identifying the workflow graph definition
type: string type: string
workspace_id:
description: |
ID of the workspace that owns this job. A successful (200)
response from this operation is only ever returned for the
caller's own job (see this operation's ownership-scoped
query), so this is always the caller's own workspace —
consumers that also need to correlate this job to its
live-progress broadcast channel (workspace+user scoped; see
the internal common/gateways/broadcast package) can use this
value directly rather than resolving their own identity a
second way.
type: string
required: required:
- id - id
- status - status
@ -1565,7 +1591,13 @@ paths:
schema: schema:
default: true default: true
type: boolean type: boolean
- description: Filter assets by exact content hash. - description: |
Filter assets by content hash, in the canonical `blake3:<hex>`
form. Matches regardless of which of this asset store's two
internal hash storage formats the matching row was written
under (the canonical form used by from-hash-created references,
or the raw `<hex>.<ext>`/bare `<hex>` storage key used by direct
uploads) — both represent the same content hash.
in: query in: query
name: hash name: hash
schema: schema:
@ -2464,6 +2496,23 @@ paths:
schema: schema:
additionalProperties: true additionalProperties: true
properties: properties:
free_tier_balance:
description: Free-tier job allowance for an authenticated non-paid (FREE-tier) user in the rollout. Absent for paid users and unauthenticated requests. Synthesized from config before a grant row exists so a brand-new user still sees their full allowance.
properties:
allowance:
description: Total free jobs granted for the current period
type: integer
remaining:
description: Free jobs remaining (allowance - used, floored at 0)
type: integer
used:
description: Free jobs consumed so far
type: integer
required:
- allowance
- used
- remaining
type: object
max_upload_size: max_upload_size:
description: Maximum upload size in bytes description: Maximum upload size in bytes
type: integer type: integer

View File

@ -1,6 +1,6 @@
[project] [project]
name = "ComfyUI" name = "ComfyUI"
version = "0.28.2" version = "0.28.0"
readme = "README.md" readme = "README.md"
license = { file = "LICENSE" } license = { file = "LICENSE" }
requires-python = ">=3.10" requires-python = ">=3.10"

View File

@ -1,5 +1,5 @@
comfyui-frontend-package==1.45.21 comfyui-frontend-package==1.47.10
comfyui-workflow-templates==0.11.12 comfyui-workflow-templates==0.11.17
comfyui-embedded-docs==0.5.8 comfyui-embedded-docs==0.5.8
torch torch
torchsde torchsde
@ -22,7 +22,7 @@ alembic
SQLAlchemy>=2.0.0 SQLAlchemy>=2.0.0
filelock filelock
av>=16.0.0 av>=16.0.0
comfy-kitchen==0.2.20 comfy-kitchen==0.2.22
comfy-aimdo==0.4.10 comfy-aimdo==0.4.10
requests requests
simpleeval>=1.0.0 simpleeval>=1.0.0

View File

@ -2,11 +2,12 @@ import pytest
import torch import torch
import tempfile import tempfile
import os import os
import sys
import av import av
import io import io
from fractions import Fraction from fractions import Fraction
from comfy_api.input_impl.video_types import VideoFromFile, VideoFromComponents from comfy_api.input_impl.video_types import VideoFromFile, VideoFromComponents
from comfy_api.util.video_types import VideoComponents from comfy_api.util.video_types import VideoComponents, VideoContainer, VideoCodec
from comfy_api.input.basic_types import AudioInput from comfy_api.input.basic_types import AudioInput
from av.error import InvalidDataError from av.error import InvalidDataError
@ -237,3 +238,526 @@ def test_duration_consistency(video_components):
manual_duration = float(components.images.shape[0] / components.frame_rate) manual_duration = float(components.images.shape[0] / components.frame_rate)
assert duration == pytest.approx(manual_duration) 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)

View File

@ -112,6 +112,17 @@ def _make_pid_v1_5_sd(latent_proj_channels=16):
return sd 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): def _add_model_diffusion_prefix(sd):
return {f"model.diffusion_model.{k}": v for k, v in sd.items()} return {f"model.diffusion_model.{k}": v for k, v in sd.items()}
@ -258,6 +269,26 @@ class TestModelDetection:
assert processed["pixel_blocks.0.adaLN_modulation_msa.bias"].shape == (12288,) assert processed["pixel_blocks.0.adaLN_modulation_msa.bias"].shape == (12288,)
assert processed["pixel_blocks.0.adaLN_modulation_mlp.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): def test_unet_config_and_required_keys_combination_is_unique(self):
"""Each model in the registry must have a unique combination of """Each model in the registry must have a unique combination of
``unet_config`` and ``required_keys``. If two models share the same ``unet_config`` and ``required_keys``. If two models share the same