From cd84f47efe4aedd4e2b97370c21a07db9e50bc96 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 7 Aug 2026 19:14:04 -0700 Subject: [PATCH 01/76] Make it easier to debug nested tensors. (#15383) --- comfy/nested_tensor.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/comfy/nested_tensor.py b/comfy/nested_tensor.py index 08c7133f8..43835b15f 100644 --- a/comfy/nested_tensor.py +++ b/comfy/nested_tensor.py @@ -83,6 +83,9 @@ class NestedTensor: def layout(self): return self.tensors[0].layout + def __repr__(self): + return f"{type(self).__name__}({self.tensors!r})" + def cat_nested(tensors, *args, **kwargs): cated_tensors = [] From 5599a05fea715cb2aff11f30f5b06e16d0dfa0c4 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Sat, 8 Aug 2026 10:27:21 +0800 Subject: [PATCH 02/76] chore: update workflow templates to v0.11.37 (#15415) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 771504af4..94cd1c5eb 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.48.7 -comfyui-workflow-templates==0.11.34 +comfyui-workflow-templates==0.11.37 comfyui-embedded-docs==0.5.9 torch torchsde From dd79c643a95402136a75a28f6187d843bcf457ed Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 7 Aug 2026 23:12:21 -0700 Subject: [PATCH 03/76] Minimum officially supported pytorch is now 2.7 (#15413) --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 8830b62ce..1721cefa5 100644 --- a/README.md +++ b/README.md @@ -194,7 +194,7 @@ Python 3.14 works but some custom nodes may have issues. The free threaded varia Python 3.13 is very well supported. If you have trouble with some custom node dependencies on 3.13 you can try 3.12 -torch 2.5 is minimally supported but using a newer version is extremely recommended. Some features and optimizations might only work on newer versions. We generally recommend using the latest major version of pytorch with the latest cuda version unless it is less than 2 weeks old. If your pytorch is more than 6 months old, please update it. +torch 2.7 is minimally supported but using a newer version is extremely recommended. Using a cu130 or above version of pytorch is required on Nvidia 20 series and above. Some features and optimizations might only work on newer versions. We generally recommend using the latest major version of pytorch with the latest cuda version unless it is less than 2 weeks old. If your pytorch is more than 6 months old, please update it. ### Instructions: From 00d02f2854892ee5b9808bc2f6348b972017886a Mon Sep 17 00:00:00 2001 From: Terry Jia Date: Sat, 8 Aug 2026 13:50:09 -0400 Subject: [PATCH 04/76] fix: make Create Layered Image discoverable and its flags self-explanatory (#15429) --- comfy_extras/nodes_compositor.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/comfy_extras/nodes_compositor.py b/comfy_extras/nodes_compositor.py index 10cea3adc..f66a9feab 100644 --- a/comfy_extras/nodes_compositor.py +++ b/comfy_extras/nodes_compositor.py @@ -486,7 +486,10 @@ class ImageCompositor(io.ComfyNode): node_id="ImageCompositor", display_name="Create Layered Image", category="image", + search_aliases=["compositor", "composite", "layer", "layers", "layer editor", "psd"], is_experimental=True, + # both flags on purpose: terminal compositor graphs must execute (the + # editor needs a run to open), and cache hits must replay the layer UI is_output_node=True, has_intermediate_output=True, inputs=[ @@ -605,7 +608,7 @@ class AddLayer(io.ComfyNode): options=list(_LAYER_MODES), default="normal", optional=True, - tooltip="Initial blend mode.", + tooltip="Initial blend mode, applied against the layers below. On the bottom layer over the default transparent background, non-normal modes produce transparency.", ), io.Float.Input( "rotation", From 9eaba63e1a9f2b27701cf0a0694aeed777da42f5 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 8 Aug 2026 14:51:05 -0700 Subject: [PATCH 05/76] Fix upscale models breaking on non dynamic vram low vram. (#15437) --- comfy_extras/nodes_upscale_model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy_extras/nodes_upscale_model.py b/comfy_extras/nodes_upscale_model.py index a4d692955..d1e0401b9 100644 --- a/comfy_extras/nodes_upscale_model.py +++ b/comfy_extras/nodes_upscale_model.py @@ -72,7 +72,7 @@ class ImageUpscaleWithModel(io.ComfyNode): 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() - model_management.load_models_gpu([upscale_model.patcher], memory_required=memory_required) + model_management.load_models_gpu([upscale_model.patcher], memory_required=memory_required, force_full_load=True) in_img = image.movedim(-1,-3).to(device) From a683fa6e577f3f73ab8a0b4d7434173ccaec9c12 Mon Sep 17 00:00:00 2001 From: "claude[bot]" <209825114+claude[bot]@users.noreply.github.com> Date: Sat, 8 Aug 2026 18:06:03 -0700 Subject: [PATCH 06/76] Add previewable_outputs_count to /api/jobs (backend half of Media Assets badge fix) (#15148) --- comfy_execution/jobs.py | 30 +++++++++ tests/execution/test_jobs.py | 123 +++++++++++++++++++++++++++++++++++ 2 files changed, 153 insertions(+) diff --git a/comfy_execution/jobs.py b/comfy_execution/jobs.py index 34c06363b..60f9b8f90 100644 --- a/comfy_execution/jobs.py +++ b/comfy_execution/jobs.py @@ -197,6 +197,7 @@ def normalize_queue_item(item: tuple, status: str) -> dict: 'priority': priority, 'create_time': create_time, 'outputs_count': 0, + 'previewable_outputs_count': 0, 'workflow_id': workflow_id, }) @@ -215,6 +216,7 @@ def normalize_history_item(prompt_id: str, history_item: dict, include_outputs: outputs = history_item.get('outputs', {}) outputs_count, preview_output = get_outputs_summary(outputs) + previewable_outputs_count = count_previewable_outputs(outputs) execution_error = None execution_start_time = None @@ -251,6 +253,7 @@ def normalize_history_item(prompt_id: str, history_item: dict, include_outputs: 'execution_end_time': execution_end_time, 'execution_error': execution_error, 'outputs_count': outputs_count, + 'previewable_outputs_count': previewable_outputs_count, 'preview_output': preview_output, 'workflow_id': workflow_id, }) @@ -345,6 +348,33 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]: return count, preview_output or fallback_preview or text_file_fallback or text_fallback +def count_previewable_outputs(outputs: dict) -> int: + """ + Count only outputs that would actually render in the expanded asset view, + i.e. items is_previewable() accepts (image/video/audio/3D/text). Kept + separate from get_outputs_summary()'s outputs_count, which counts every + output item regardless of media type, so a job with a non-previewable + saved file alongside real media (e.g. SaveLatent's .latent output next to + a SaveImage output) doesn't inflate the Media Assets badge beyond what + the expanded view shows. + """ + count = 0 + for node_outputs in outputs.values(): + if not isinstance(node_outputs, dict): + continue + for media_type, items in node_outputs.items(): + if media_type == 'animated' or not isinstance(items, list): + continue + for item in items: + if not isinstance(item, dict): + item = normalize_output_item(item) + if item is None: + continue + if is_previewable(media_type, item): + count += 1 + return count + + def apply_sorting(jobs: list[dict], sort_by: str, sort_order: str) -> list[dict]: """Sort jobs list by specified field and order.""" reverse = (sort_order == 'desc') diff --git a/tests/execution/test_jobs.py b/tests/execution/test_jobs.py index cef2b41cb..ffa069435 100644 --- a/tests/execution/test_jobs.py +++ b/tests/execution/test_jobs.py @@ -10,6 +10,7 @@ from comfy_execution.jobs import ( normalize_output_item, normalize_outputs, get_outputs_summary, + count_previewable_outputs, apply_sorting, has_3d_extension, validate_job_id, @@ -361,6 +362,79 @@ class TestGetOutputsSummary: assert preview['mediaType'] == 'files' +class TestCountPreviewableOutputs: + """Unit tests for count_previewable_outputs() + + Kept separate from get_outputs_summary()'s outputs_count: the Media Assets + badge should reflect only what the expanded asset view actually renders + (previewable outputs), while outputs_count keeps counting every output + item for other consumers. + """ + + def test_empty_outputs(self): + assert count_previewable_outputs({}) == 0 + + def test_previewable_outputs_all_counted(self): + """When every output is previewable, the two counts should match.""" + outputs = { + 'node1': {'images': [{'filename': 'a.png', 'type': 'output'}]}, + 'node2': {'images': [{'filename': 'b.png', 'type': 'output'}]}, + } + outputs_count, _ = get_outputs_summary(outputs) + assert count_previewable_outputs(outputs) == outputs_count == 2 + + def test_save_latent_counted_but_not_previewable(self): + """SaveLatent (nodes.py) emits a real saved file under the 'latents' + media type: {'latents': [{'filename': '..._00001_.latent', + 'subfolder': '', 'type': 'output'}]}. It has no previewable media + type, format, or extension, so it inflates outputs_count without + ever rendering in the expanded asset view.""" + outputs = { + 'node1': { + 'images': [{'filename': 'ComfyUI_00001_.png', 'subfolder': '', 'type': 'output'}] + }, + 'node2': { + 'latents': [{'filename': 'ComfyUI_00001_.latent', 'subfolder': '', 'type': 'output'}] + }, + } + outputs_count, _ = get_outputs_summary(outputs) + assert outputs_count == 2 + assert count_previewable_outputs(outputs) == 1 + + def test_save_text_file_output_is_previewable_by_extension(self): + """SaveText (comfy_extras/nodes_text.py) emits its saved file under a + 'files' media type via ui.SavedResult: {'files': [{'filename': + '..._00001.txt', 'subfolder': ..., 'type': 'output'}]}. The .txt + extension makes it previewable even though 'files' itself isn't a + previewable media type.""" + outputs = { + 'node1': { + 'files': [{'filename': 'ComfyUI_00001.txt', 'subfolder': '', 'type': 'output'}] + } + } + assert count_previewable_outputs(outputs) == 1 + + def test_preview_any_text_tuple_not_counted(self): + """PreviewAny (comfy_extras/nodes_preview_any.py) emits only + {'text': (value,)} with no saved file. Since the value is a tuple, + not a list, it is excluded from both outputs_count and + previewable_outputs_count — matching get_outputs_summary().""" + outputs = { + 'node1': {'text': ('some previewed value',)} + } + outputs_count, _ = get_outputs_summary(outputs) + assert outputs_count == 0 + assert count_previewable_outputs(outputs) == 0 + + def test_string_3d_filename_previewable(self): + """String 3D filenames (e.g. Preview3D) normalize into a previewable + item just like they do for outputs_count.""" + outputs = { + 'node1': {'result': ['preview3d_abc123.glb', None]} + } + assert count_previewable_outputs(outputs) == 1 + + class TestHas3DExtension: """Unit tests for has_3d_extension()""" @@ -447,6 +521,7 @@ class TestNormalizeQueueItem: assert 'execution_error' not in job assert 'preview_output' not in job assert job['outputs_count'] == 0 + assert job['previewable_outputs_count'] == 0 assert job['workflow_id'] == 'workflow-abc' @@ -635,6 +710,54 @@ class TestNormalizeHistoryItem: {'filename': 'photo.png', 'type': 'output', 'subfolder': ''}, ] + def test_previewable_outputs_count_excludes_non_previewable_outputs(self): + """Regression test for the Media Assets badge overcount: a job with an + image (SaveImage) and a SaveLatent output should report previewable_ + outputs_count == 1 while outputs_count == 2, so the frontend badge + (once switched to previewable_outputs_count) matches what the + expanded asset view actually renders.""" + history_item = { + 'prompt': ( + 5, + 'prompt-mixed', + {'nodes': {}}, + {'create_time': 1234567890}, + ['node1', 'node2'], + ), + 'status': {'status_str': 'success', 'completed': True, 'messages': []}, + 'outputs': { + 'node1': { + 'images': [{'filename': 'ComfyUI_00001_.png', 'subfolder': '', 'type': 'output'}] + }, + 'node2': { + 'latents': [{'filename': 'ComfyUI_00001_.latent', 'subfolder': '', 'type': 'output'}] + }, + }, + } + job = normalize_history_item('prompt-mixed', history_item) + + assert job['outputs_count'] == 2 + assert job['previewable_outputs_count'] == 1 + + def test_previewable_outputs_count_zero_pruned_by_prune_dict(self): + """A job with no outputs at all should still report both counts as 0, + not omit the field (prune_dict only strips None, not 0).""" + history_item = { + 'prompt': ( + 5, + 'prompt-empty', + {'nodes': {}}, + {'create_time': 1234567890}, + ['node1'], + ), + 'status': {'status_str': 'success', 'completed': True, 'messages': []}, + 'outputs': {}, + } + job = normalize_history_item('prompt-empty', history_item) + + assert job['outputs_count'] == 0 + assert job['previewable_outputs_count'] == 0 + class TestNormalizeOutputItem: """Unit tests for normalize_output_item()""" From 40e46c711025947f126cccaa1a692a14937a3096 Mon Sep 17 00:00:00 2001 From: chaObserv <154517000+chaObserv@users.noreply.github.com> Date: Sun, 9 Aug 2026 10:06:12 +0800 Subject: [PATCH 07/76] Extend ER-SDE noise scaler by scaling h(t) (#15428) --- comfy_extras/nodes_custom_sampler.py | 36 +++++++++++++++++++--------- 1 file changed, 25 insertions(+), 11 deletions(-) diff --git a/comfy_extras/nodes_custom_sampler.py b/comfy_extras/nodes_custom_sampler.py index e81b6328b..d5aa730d2 100644 --- a/comfy_extras/nodes_custom_sampler.py +++ b/comfy_extras/nodes_custom_sampler.py @@ -591,7 +591,7 @@ class SamplerER_SDE(io.ComfyNode): inputs=[ io.Combo.Input("solver_type", options=["ER-SDE", "Reverse-time SDE", "ODE"]), io.Int.Input("max_stage", default=3, min=1, max=3, advanced=True), - io.Float.Input("eta", default=1.0, min=0.0, max=100.0, step=0.01, round=False, tooltip="Stochastic strength of reverse-time SDE.\nWhen eta=0, it reduces to deterministic ODE. This setting doesn't apply to ER-SDE solver type.", advanced=True), + io.Float.Input("eta", default=1.0, min=0.0, max=10.0, step=0.01, round=False, tooltip="Stochastic strength of SDEs.\nWhen eta=0, they reduce to deterministic ODE.\nLarge eta may cause invalid outputs. If this occurs, try decreasing this value.", advanced=True), io.Float.Input("s_noise", default=1.0, min=0.0, max=100.0, step=0.01, round=False, advanced=True), ], outputs=[io.Sampler.Output()] @@ -599,21 +599,35 @@ class SamplerER_SDE(io.ComfyNode): @classmethod def execute(cls, solver_type, max_stage, eta, s_noise) -> io.NodeOutput: - if solver_type == "ODE" or (solver_type == "Reverse-time SDE" and eta == 0): - eta = 0 - s_noise = 0 + # Extend existing noise scalers phi(x) with eta-controlled noise scalers: + # psi(x) = x**(1-eta) * phi(x)**eta + # where eta is constant and directly scales the h^2(t) contribution. - def reverse_time_sde_noise_scaler(x): + def er_sde_noise_scaler(x: torch.Tensor) -> torch.Tensor: + return x * ((x ** 0.3).exp() + 10.0) ** eta + + def reverse_time_sde_noise_scaler(x: torch.Tensor) -> torch.Tensor: return x ** (eta + 1) - if solver_type == "ER-SDE": - # Use the default one in sample_er_sde() - noise_scaler = None - else: - noise_scaler = reverse_time_sde_noise_scaler + def ode_noise_scaler(x: torch.Tensor) -> torch.Tensor: + return x + + solver_scalers = { + "ER-SDE": er_sde_noise_scaler, + "Reverse-time SDE": reverse_time_sde_noise_scaler, + "ODE": ode_noise_scaler, + } + + if solver_type == "ODE" or eta == 0: + s_noise = 0.0 + solver_type = "ODE" + noise_scaler = solver_scalers[solver_type] sampler_name = "er_sde" - sampler = comfy.samplers.ksampler(sampler_name, {"s_noise": s_noise, "noise_scaler": noise_scaler, "max_stage": max_stage}) + sampler = comfy.samplers.ksampler( + sampler_name, + {"s_noise": s_noise, "noise_scaler": noise_scaler, "max_stage": max_stage}, + ) return io.NodeOutput(sampler) get_sampler = execute From cbbc9dab1f03d0d9a6caa8a8be7d77a7e37e1e44 Mon Sep 17 00:00:00 2001 From: blepping <157360029+blepping@users.noreply.github.com> Date: Sat, 8 Aug 2026 20:38:34 -0600 Subject: [PATCH 08/76] Make a context manager for cast_bias_weight and use it. (#14750) --- comfy/background_removal/birefnet.py | 23 ++- comfy/controlnet.py | 11 +- comfy/ops.py | 241 ++++++++++++--------------- comfy/text_encoders/llama.py | 14 +- 4 files changed, 124 insertions(+), 165 deletions(-) diff --git a/comfy/background_removal/birefnet.py b/comfy/background_removal/birefnet.py index 78a80246e..ba3f710d4 100644 --- a/comfy/background_removal/birefnet.py +++ b/comfy/background_removal/birefnet.py @@ -433,19 +433,16 @@ class DeformableConv2d(nn.Module): def forward(self, x): offset = self.offset_conv(x) modulator = 2. * torch.sigmoid(self.modulator_conv(x)) - weight, bias, offload_info = comfy.ops.cast_bias_weight(self.regular_conv, x, offloadable=True) - - x = deform_conv2d( - input=x, - offset=offset, - weight=weight, - bias=None, - padding=self.padding, - mask=modulator, - stride=self.stride, - ) - comfy.ops.uncast_bias_weight(self.regular_conv, weight, bias, offload_info) - return x + with comfy.ops.CastBiasWeightContext(self.regular_conv, x, offloadable=True) as (weight, _bias): + return deform_conv2d( + input=x, + offset=offset, + weight=weight, + bias=None, + padding=self.padding, + mask=modulator, + stride=self.stride, + ) class BasicDecBlk(nn.Module): def __init__(self, in_channels=64, out_channels=64, inter_channels=64, device=None, dtype=None, operations=None): diff --git a/comfy/controlnet.py b/comfy/controlnet.py index 6dbbaa959..7e35fe027 100644 --- a/comfy/controlnet.py +++ b/comfy/controlnet.py @@ -381,13 +381,10 @@ class ControlLoraOps: self.bias = None def forward(self, input): - weight, bias, offload_stream = comfy.ops.cast_bias_weight(self, input, offloadable=True) - if self.up is not None: - x = torch.nn.functional.linear(input, weight + (torch.mm(self.up.flatten(start_dim=1), self.down.flatten(start_dim=1))).reshape(self.weight.shape).type(input.dtype), bias) - else: - x = torch.nn.functional.linear(input, weight, bias) - comfy.ops.uncast_bias_weight(self, weight, bias, offload_stream) - return x + with comfy.ops.CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + if self.up is None: + return torch.nn.functional.linear(input, weight, bias) + return torch.nn.functional.linear(input, weight + (torch.mm(self.up.flatten(start_dim=1), self.down.flatten(start_dim=1))).reshape(self.weight.shape).type(input.dtype), bias) class Conv2d(torch.nn.Module, comfy.ops.CastWeightBiasOp): def __init__( diff --git a/comfy/ops.py b/comfy/ops.py index 14599997b..9ec44cfa2 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -452,6 +452,26 @@ def uncast_bias_weight(s, weight, bias, offload_stream): device = bias_a.device os.wait_stream(comfy.model_management.current_stream(device)) +class CastBiasWeightContext: + # When initialized with no arguments or the first is None, the context + # will return the tuple (None, None). + def __init__(self, *args, **kwargs): + self.slf = args[0] if len(args) else None + self.state = (None, None) if self.slf is None else cast_bias_weight(*args, **kwargs) + + def __enter__(self): + result = self.state + if len(result) < 3 or result[2] is None: + # Not offloaded, immediately drop references. + self.state = self.slf = None + return result[:2] + + def __exit__(self, *_args) -> None: + if self.slf is None: + return + slf, state = self.slf, self.state + self.state = self.slf = None + uncast_bias_weight(slf, *state) class CastWeightBiasOp: comfy_cast_weights = False @@ -538,10 +558,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.linear(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.linear(input, weight, bias) def forward(self, *args, **kwargs): run_every_op() @@ -555,10 +573,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = self._conv_forward(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return self._conv_forward(input, weight, bias) def forward(self, *args, **kwargs): run_every_op() @@ -572,10 +588,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = self._conv_forward(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return self._conv_forward(input, weight, bias) def forward(self, *args, **kwargs): run_every_op() @@ -600,10 +614,8 @@ class disable_weight_init: return super()._conv_forward(input, weight, bias, *args, **kwargs) def forward_comfy_cast_weights(self, input, autopad=None): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = self._conv_forward(input, weight, bias, autopad=autopad) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return self._conv_forward(input, weight, bias, autopad=autopad) def forward(self, *args, **kwargs): run_every_op() @@ -617,10 +629,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.group_norm(input, self.num_groups, weight, bias, self.eps) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.group_norm(input, self.num_groups, weight, bias, self.eps) def forward(self, *args, **kwargs): run_every_op() @@ -634,12 +644,10 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - running_mean = self.running_mean.to(device=input.device, dtype=weight.dtype) if self.running_mean is not None else None - running_var = self.running_var.to(device=input.device, dtype=weight.dtype) if self.running_var is not None else None - x = torch.nn.functional.batch_norm(input, running_mean, running_var, weight, bias, self.training, self.momentum, self.eps) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + running_mean = self.running_mean.to(device=input.device, dtype=weight.dtype) if self.running_mean is not None else None + running_var = self.running_var.to(device=input.device, dtype=weight.dtype) if self.running_var is not None else None + return torch.nn.functional.batch_norm(input, running_mean, running_var, weight, bias, self.training, self.momentum, self.eps) def forward(self, *args, **kwargs): run_every_op() @@ -653,15 +661,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - if self.weight is not None: - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - else: - weight = None - bias = None - offload_stream = None - x = torch.nn.functional.layer_norm(input, self.normalized_shape, weight, bias, self.eps) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self if self.weight is not None else None, input, offloadable=True) as (weight, bias): + return torch.nn.functional.layer_norm(input, self.normalized_shape, weight, bias, self.eps) def forward(self, *args, **kwargs): run_every_op() @@ -676,15 +677,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - if self.weight is not None: - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - else: - weight = None - bias = None - offload_stream = None - x = torch.nn.functional.rms_norm(input, self.normalized_shape, weight, self.eps) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self if self.weight is not None else None, input, offloadable=True) as (weight, bias): + return torch.nn.functional.rms_norm(input, self.normalized_shape, weight, self.eps) def forward(self, *args, **kwargs): run_every_op() @@ -703,12 +697,10 @@ class disable_weight_init: input, output_size, self.stride, self.padding, self.kernel_size, num_spatial_dims, self.dilation) - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.conv_transpose2d( - input, weight, bias, self.stride, self.padding, - output_padding, self.groups, self.dilation) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.conv_transpose2d( + input, weight, bias, self.stride, self.padding, + output_padding, self.groups, self.dilation) def forward(self, *args, **kwargs): run_every_op() @@ -727,12 +719,10 @@ class disable_weight_init: input, output_size, self.stride, self.padding, self.kernel_size, num_spatial_dims, self.dilation) - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.conv_transpose1d( - input, weight, bias, self.stride, self.padding, - output_padding, self.groups, self.dilation) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.conv_transpose1d( + input, weight, bias, self.stride, self.padding, + output_padding, self.groups, self.dilation) def forward(self, *args, **kwargs): run_every_op() @@ -795,10 +785,8 @@ class disable_weight_init: output_dtype = out_dtype if self.weight.dtype == torch.float16 or self.weight.dtype == torch.bfloat16: out_dtype = None - weight, bias, offload_stream = cast_bias_weight(self, device=input.device, dtype=out_dtype, offloadable=True) - x = torch.nn.functional.embedding(input, weight, self.padding_idx, self.max_norm, self.norm_type, self.scale_grad_by_freq, self.sparse).to(dtype=output_dtype) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, device=input.device, dtype=out_dtype, offloadable=True) as (weight, bias): + return torch.nn.functional.embedding(input, weight, self.padding_idx, self.max_norm, self.norm_type, self.scale_grad_by_freq, self.sparse).to(dtype=output_dtype) def forward(self, *args, **kwargs): @@ -874,7 +862,6 @@ def fp8_linear(self, input): if input.ndim != 2: return None lora_compute_dtype=comfy.model_management.lora_compute_dtype(input.device) - w, bias, offload_stream = cast_bias_weight(self, input, dtype=dtype, bias_dtype=input_dtype, offloadable=True, compute_dtype=lora_compute_dtype, want_requant=True) scale_weight = torch.ones((), device=input.device, dtype=torch.float32) scale_input = torch.ones((), device=input.device, dtype=torch.float32) @@ -883,15 +870,16 @@ def fp8_linear(self, input): layout_params_input = TensorCoreFP8Layout.Params(scale=scale_input, orig_dtype=input_dtype, orig_shape=tuple(input_fp8.shape)) quantized_input = QuantizedTensor(input_fp8, "TensorCoreFP8Layout", layout_params_input) - # Wrap weight in QuantizedTensor - this enables unified dispatch - # Call F.linear - __torch_dispatch__ routes to fp8_linear handler in quant_ops.py! - layout_params_weight = TensorCoreFP8Layout.Params(scale=scale_weight, orig_dtype=input_dtype, orig_shape=tuple(w.shape)) - quantized_weight = QuantizedTensor(w, "TensorCoreFP8Layout", layout_params_weight) - o = torch.nn.functional.linear(quantized_input, quantized_weight, bias) + with CastBiasWeightContext(self, input, dtype=dtype, bias_dtype=input_dtype, offloadable=True, compute_dtype=lora_compute_dtype, want_requant=True) as (w, bias): + # Wrap weight in QuantizedTensor - this enables unified dispatch + # Call F.linear - __torch_dispatch__ routes to fp8_linear handler in quant_ops.py! + w_shape = tuple(w.shape) + layout_params_weight = TensorCoreFP8Layout.Params(scale=scale_weight, orig_dtype=input_dtype, orig_shape=w_shape) + quantized_weight = QuantizedTensor(w, "TensorCoreFP8Layout", layout_params_weight) + o = torch.nn.functional.linear(quantized_input, quantized_weight, bias) - uncast_bias_weight(self, w, bias, offload_stream) if tensor_3d: - o = o.reshape((input_shape[0], input_shape[1], w.shape[0])) + o = o.reshape((input_shape[0], input_shape[1], w_shape[0])) return o @@ -911,10 +899,8 @@ class fp8_ops(manual_cast): except Exception as e: logging.info("Exception during fp8 op: {}".format(e)) - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.linear(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.linear(input, weight, bias) CUBLAS_IS_AVAILABLE = False try: @@ -930,10 +916,8 @@ if CUBLAS_IS_AVAILABLE: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = cublas_half_matmul(input, weight, bias, self._epilogue_str, self.has_bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return cublas_half_matmul(input, weight, bias, self._epilogue_str, self.has_bias) def forward(self, *args, **kwargs): run_every_op() @@ -1344,29 +1328,28 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec want_requant=False, weight_only_quant=False, ): - if weight_only_quant: - weight, bias, offload_stream = cast_bias_weight( - self, - input=None, - dtype=self.weight.dtype, - device=input.device, - bias_dtype=input.dtype, - offloadable=True, - compute_dtype=compute_dtype, - want_requant=True, - ) - weight = weight.to(dtype=input.dtype) - else: - weight, bias, offload_stream = cast_bias_weight( + if not weight_only_quant: + with CastBiasWeightContext( self, input, offloadable=True, compute_dtype=compute_dtype, want_requant=want_requant, - ) - x = self._forward(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + ) as (weight, bias): + return self._forward(input, weight, bias) + + with CastBiasWeightContext( + self, + input=None, + dtype=self.weight.dtype, + device=input.device, + bias_dtype=input.dtype, + offloadable=True, + compute_dtype=compute_dtype, + want_requant=True, + ) as (weight, bias): + weight = weight.to(dtype=input.dtype) + return self._forward(input, weight, bias) def forward(self, input, *args, **kwargs): run_every_op() @@ -1391,25 +1374,20 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec # Training path: quantized forward with compute_dtype backward via autograd function if (input.requires_grad and _use_quantized and quantize_input): - - weight, bias, offload_stream = cast_bias_weight( + with CastBiasWeightContext( self, input, offloadable=True, compute_dtype=compute_dtype, want_requant=True - ) + ) as (weight, bias): + scale = getattr(self, 'input_scale', None) + if scale is not None: + scale = comfy.model_management.cast_to_device(scale, input.device, None) - scale = getattr(self, 'input_scale', None) - if scale is not None: - scale = comfy.model_management.cast_to_device(scale, input.device, None) - - output = QuantLinearFunc.apply( - input, weight, bias, self.layout_type, scale, compute_dtype - ) - - uncast_bias_weight(self, weight, bias, offload_stream) - return output + return QuantLinearFunc.apply( + input, weight, bias, self.layout_type, scale, compute_dtype + ) # Inference path (unchanged) if _use_quantized and quantize_input: @@ -1520,13 +1498,11 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec """Cast the whole bank once; expert_linear inside reuses the cast. Not re-entrant — do not nest calls on the same instance. """ - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - self._resident_bank = (weight, bias) - try: - yield self - finally: - self._resident_bank = None - uncast_bias_weight(self, weight, bias, offload_stream) + with CastBiasWeightContext(self, input, offloadable=True) as self._resident_bank: + try: + yield self + finally: + self._resident_bank = None def expert_linear(self, input: torch.Tensor, i: int) -> torch.Tensor: """Linear against expert i's weight (with optional bias).""" @@ -1534,11 +1510,8 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec if resident is not None: weight, bias = resident return self._expert_linear_impl(input, weight, bias, i) - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - try: + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): return self._expert_linear_impl(input, weight, bias, i) - finally: - uncast_bias_weight(self, weight, bias, offload_stream) def _expert_linear_impl(self, input, weight, bias, i): if isinstance(weight, QuantizedTensor): @@ -1641,25 +1614,23 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec # Optimized path: lookup in fp8/int8, dequantize only the selected rows. if isinstance(weight, QuantizedTensor) and len(self.weight_function) == 0: - qdata, _, offload_stream = cast_bias_weight(self, device=input.device, dtype=weight.dtype, offloadable=True) - if isinstance(qdata, QuantizedTensor): - params = qdata._params - scale = params.scale - qdata = qdata._qdata - else: - params = weight._params - scale = None + with CastBiasWeightContext(self, device=input.device, dtype=weight.dtype, offloadable=True) as (qdata, _bias): + if isinstance(qdata, QuantizedTensor): + params = qdata._params + scale = params.scale + qdata = qdata._qdata + else: + params = weight._params + scale = None - # int8: per-row scale possible ConvRot, so let the layout do the gather - if self.quant_format == "int8_tensorwise": - x = get_layout_class(self.layout_type).dequantize_embedding(qdata, params, input) - uncast_bias_weight(self, qdata, None, offload_stream) - return x if out_dtype is None else x.to(dtype=out_dtype) + # int8: per-row scale possible ConvRot, so let the layout do the gather + if self.quant_format == "int8_tensorwise": + x = get_layout_class(self.layout_type).dequantize_embedding(qdata, params, input) + return x if out_dtype is None else x.to(dtype=out_dtype) - x = torch.nn.functional.embedding( - input, qdata, self.padding_idx, self.max_norm, - self.norm_type, self.scale_grad_by_freq, self.sparse) - uncast_bias_weight(self, qdata, None, offload_stream) + x = torch.nn.functional.embedding( + input, qdata, self.padding_idx, self.max_norm, + self.norm_type, self.scale_grad_by_freq, self.sparse) target_dtype = out_dtype if out_dtype is not None else weight._params.orig_dtype x = x.to(dtype=target_dtype) if scale is not None and scale != 1.0: diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index f5c5597ef..371ec1bbc 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -868,16 +868,10 @@ class BaseGenerate: else: module = self.model.embed_tokens - offload_stream = None - if module.comfy_cast_weights: - weight, _, offload_stream = comfy.ops.cast_bias_weight(module, input, offloadable=True) - else: - weight = self.model.embed_tokens.weight.to(x) - - x = torch.nn.functional.linear(input, weight, None) - - comfy.ops.uncast_bias_weight(module, weight, None, offload_stream) - return x + if not module.comfy_cast_weights: + return torch.nn.functional.linear(input, self.model.embed_tokens.weight.to(x), None) + with comfy.ops.CastBiasWeightContext(module, input, offloadable=True) as (weight, _bias): + return torch.nn.functional.linear(input, weight, None) def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): model_config = self.model.config From 2a68ce33b4c9ea6ee4283e618a74560cefb32694 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Sun, 9 Aug 2026 21:24:48 +0300 Subject: [PATCH 09/76] Optimize MiniMax-H3 VAE (#15446) --- comfy/ldm/minimax/vae.py | 104 ++++++++++++++++++++++----------------- comfy/sd.py | 9 ++++ 2 files changed, 68 insertions(+), 45 deletions(-) diff --git a/comfy/ldm/minimax/vae.py b/comfy/ldm/minimax/vae.py index 65d06f3e9..e1c146ec4 100644 --- a/comfy/ldm/minimax/vae.py +++ b/comfy/ldm/minimax/vae.py @@ -6,6 +6,7 @@ import torch import torch.nn as nn import torch.nn.functional as F +import comfy.model_management import comfy.ops import comfy.quant_ops import comfy.rmsnorm @@ -321,6 +322,8 @@ class ViT3DDecoder(nn.Module): # Full VAE class MiniMaxH3VideoVAE(nn.Module): + comfy_has_chunked_io = True + def __init__( self, in_channels=3, @@ -389,6 +392,23 @@ class MiniMaxH3VideoVAE(nn.Module): def _decode_pixels(self, z): return self.decoder(self.post_quant_conv(z)) + def _normalize_pixels(self, x): + return x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x)) + + def _finalize_pixels(self, part): + # raw decoder output -> float32 pixels in [0, 1] (the VAE wrapper's process_output is identity) + part = part * self.pixel_std.to(device=part.device, dtype=torch.float32) + return part.add_(self.pixel_mean.to(device=part.device, dtype=torch.float32)).clamp_(0.0, 1.0) + + def decode_output_shape(self, input_shape): + b, c, t, h, w = input_shape + if t == 1: + frames = 1 + else: + pad_tokens, num_chunks = self._decode_temporal_chunks(t) + frames = self._decode_temporal_frame_plan(t + pad_tokens, num_chunks, pad_tokens) + return (b, self.decoder.out_channels, frames, h * self.vae_ratio, w * self.vae_ratio) + def _adaptive_encode(self, x): if self.tiling: return self.tiled_encode(x) @@ -521,18 +541,15 @@ class MiniMaxH3VideoVAE(nn.Module): # temporal chunking - def encode_temporal(self, x): - if x.shape[2] % self.clip_length != 0: - pad_size = (-x.shape[2]) % self.clip_length - pad_frames = x[:, :, -1:].repeat(1, 1, pad_size, 1, 1) - x = torch.cat([x, pad_frames], dim=2) - - num_chunks = x.shape[2] // self.clip_length - + def encode_temporal(self, x, device): + # chunked input io: x may live on the CPU, clips move to the device as they encode z_list = [] - for i in range(num_chunks): - clip_x = x[:, :, i * self.clip_length:(i + 1) * self.clip_length, :, :] - z_list.append(self._adaptive_encode(clip_x)) + for i in range(math.ceil(x.shape[2] / self.clip_length)): + clip_x = x[:, :, i * self.clip_length:(i + 1) * self.clip_length, :, :].to(device) + if clip_x.shape[2] < self.clip_length: + pad_frames = clip_x[:, :, -1:].repeat(1, 1, self.clip_length - clip_x.shape[2], 1, 1) + clip_x = torch.cat([clip_x, pad_frames], dim=2) + z_list.append(self._adaptive_encode(self._normalize_pixels(clip_x))) z = torch.cat(z_list, dim=2) if self.token_drop > 0: @@ -577,43 +594,42 @@ class MiniMaxH3VideoVAE(nn.Module): total_frames += final_overlap_frames return total_frames - self._decode_temporal_pad_frames(z_len, pad_tokens) - def decode_temporal(self, z): - chunk_dec = self.tokens_chunk_size * self.vae_ratio_t - split_count = int(self.token_drop > 0) + 1 - - pseudo_total_tokens = z.shape[2] + self.token_drop - - pad_tokens = 0 - remainder = pseudo_total_tokens % self.tokens_chunk_size - if remainder != 0: - pad_tokens = self.tokens_chunk_size - remainder - pseudo_total_tokens += pad_tokens + def _decode_temporal_chunks(self, z_len): + pseudo_total_tokens = z_len + self.token_drop + pad_tokens = (-pseudo_total_tokens) % self.tokens_chunk_size + pseudo_total_tokens += pad_tokens num_chunks = pseudo_total_tokens // self.tokens_chunk_size - int(self.token_drop > 0) if num_chunks < 1: # too few tokens for one chunk (e.g. T_lat == 2): pad one extra chunk pad_tokens += self.tokens_chunk_size num_chunks += 1 + return pad_tokens, num_chunks + def decode_temporal(self, z, output_buffer=None): + chunk_dec = self.tokens_chunk_size * self.vae_ratio_t + split_count = int(self.token_drop > 0) + 1 + + if output_buffer is None: + # finalized chunks stream out of VRAM so the full video never sits on the GPU + output_buffer = torch.empty(self.decode_output_shape(z.shape), dtype=torch.float32, + device=comfy.model_management.intermediate_device()) + + pad_tokens, num_chunks = self._decode_temporal_chunks(z.shape[2]) if pad_tokens > 0: pad_z = z[:, :, -1:, :, :].repeat(1, 1, pad_tokens, 1, 1) z = torch.cat([z, pad_z], dim=2) - output_frames = self._decode_temporal_frame_plan(z.shape[2], num_chunks, pad_tokens) - - dec = None + dec = output_buffer dec_overlap = None write_pos = 0 def write_part(part): - nonlocal dec, write_pos + nonlocal write_pos part_frames = part.shape[2] if part_frames <= 0: return - if dec is None: - out_shape = list(part.shape) - out_shape[2] = output_frames - dec = torch.empty(out_shape, dtype=part.dtype, device=part.device) + part = self._finalize_pixels(part) copy_frames = min(part_frames, max(0, dec.shape[2] - write_pos)) if copy_frames > 0: dec[:, :, write_pos:write_pos + copy_frames, :, :].copy_( @@ -653,18 +669,18 @@ class MiniMaxH3VideoVAE(nn.Module): return dec - def encode(self, x): + def encode(self, x, device=None): # x: [B, 3, T, H, W] in [-1, 1] -> normalized latents [B, 24, T_lat, H/16, W/16] if x.ndim == 4: x = x.unsqueeze(2) - - x = x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x)) + if device is None: + device = x.device if x.shape[2] == 1: - moments = self._adaptive_encode(x) + moments = self._adaptive_encode(self._normalize_pixels(x.to(device))) moments = moments[:, :, -1:, :, :] else: - moments = self.encode_temporal(x) + moments = self.encode_temporal(x, device) mean = torch.chunk(moments.float(), 2, dim=1)[0] @@ -679,18 +695,16 @@ class MiniMaxH3VideoVAE(nn.Module): def decode_tiled(self, z, **kwargs): return self.decode(z) - def decode(self, z): - # z: [B, 24, T_lat, H_lat, W_lat] normalized latents -> pixels [B, 3, T, H, W] in [-1, 1] + def decode(self, z, output_buffer=None): + # z: [B, 24, T_lat, H_lat, W_lat] normalized latents -> float32 pixels [B, 3, T, H, W] in [0, 1] latents_mean = self.latents_mean.view(1, -1, 1, 1, 1).to(z) latents_std = self.latents_std.view(1, -1, 1, 1, 1).to(z) z = z * latents_std + latents_mean if z.shape[2] == 1: - dec = self._adaptive_decode(z) - dec = dec[:, :, -1:, :, :] - else: - dec = self.decode_temporal(z) - - dec = dec.float() - dec.mul_(self.pixel_std.to(dec)).add_(self.pixel_mean.to(dec)).clamp_(0.0, 1.0).mul_(2.0).sub_(1.0) - return dec + dec = self._finalize_pixels(self._adaptive_decode(z)[:, :, -1:, :, :]) + if output_buffer is None: + return dec + output_buffer.copy_(dec) + return output_buffer + return self.decode_temporal(z, output_buffer) diff --git a/comfy/sd.py b/comfy/sd.py index 9ccd561bc..5fed4ca9a 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -955,13 +955,21 @@ class VAE: self.working_dtypes = [torch.float16, torch.float32] # the model tiles internally (256px spatial, 17-frame temporal chunks) self.handles_tiling = True + # decode finalizes straight to [0, 1] while streaming chunks out + self.process_output = lambda image: image + # one decoded temporal chunk (with overlap) is all that ever sits in VRAM + chunk_frames = (self.first_stage_model.tokens_chunk_size + self.first_stage_model.token_overlap) * self.first_stage_model.vae_ratio_t + def estimate_encode_memory(frames, height, width, dtype): fixed = 110_000_000 if frames == 1 else 1_300_000_000 elements_per_pixel = 7 if frames == 1 else 9.5 + # only one clip of the input video is ever resident on the GPU + frames = min(frames, self.first_stage_model.clip_length) return (elements_per_pixel * frames * height * width + fixed) * model_management.dtype_size(dtype) * 1.03 def estimate_decode_memory(frames, height, width, dtype): fixed = 110_000_000 if frames <= 22 else 270_000_000 + frames = min(frames, chunk_frames + 2) return (9.5 * frames * height * width + fixed) * model_management.dtype_size(dtype) * 1.03 self.memory_used_encode = lambda shape, dtype: estimate_encode_memory(shape[2], shape[3], shape[4], dtype) @@ -1198,6 +1206,7 @@ class VAE: do_tile = True if do_tile: + pixel_samples = None comfy.model_management.soft_empty_cache() dims = samples_in.ndim - 2 if dims == 1 or self.extra_1d_channel is not None: From 7d11ec31cb700d881fdf2d73731ecde0093b9540 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Mon, 10 Aug 2026 10:13:24 +0300 Subject: [PATCH 10/76] [Partner Nodes] feat(Qwen): add Qwen-Image 3.0 image generation and editing nodes (#15327) Signed-off-by: Alexander Piskun --- comfy_api_nodes/apis/qwen.py | 46 ++++ comfy_api_nodes/nodes_qwen.py | 442 ++++++++++++++++++++++++++++++++++ 2 files changed, 488 insertions(+) create mode 100644 comfy_api_nodes/apis/qwen.py create mode 100644 comfy_api_nodes/nodes_qwen.py diff --git a/comfy_api_nodes/apis/qwen.py b/comfy_api_nodes/apis/qwen.py new file mode 100644 index 000000000..90b68dee8 --- /dev/null +++ b/comfy_api_nodes/apis/qwen.py @@ -0,0 +1,46 @@ +from pydantic import BaseModel, Field + + +class QwenImageContentItem(BaseModel): + image: str | None = Field(None) + text: str | None = Field(None) + + +class QwenImageMessage(BaseModel): + role: str = Field("user") + content: list[QwenImageContentItem] = Field(...) + + +class QwenImageInputField(BaseModel): + messages: list[QwenImageMessage] = Field(...) + + +class QwenImageParametersField(BaseModel): + size: str | None = Field(None, description="Output resolution as 'width*height'; omit for the model default.") + n: int = Field(1, ge=1, le=6) + seed: int = Field(..., ge=0, le=2147483647) + prompt_extend: bool = Field(True) + watermark: bool = Field(False) + negative_prompt: str | None = Field(None) + + +class QwenImageGenerationRequest(BaseModel): + model: str = Field(...) + input: QwenImageInputField = Field(...) + parameters: QwenImageParametersField = Field(...) + + +class QwenImageChoice(BaseModel): + finish_reason: str | None = Field(None) + message: QwenImageMessage | None = Field(None) + + +class QwenImageOutputField(BaseModel): + choices: list[QwenImageChoice] = Field(default_factory=list) + + +class QwenImageGenerationResponse(BaseModel): + output: QwenImageOutputField | None = Field(None) + request_id: str = Field(...) + code: str | None = Field(None, description="Error code for the failed request.") + message: str | None = Field(None, description="Details about the failed request.") diff --git a/comfy_api_nodes/nodes_qwen.py b/comfy_api_nodes/nodes_qwen.py new file mode 100644 index 000000000..3b6c5023c --- /dev/null +++ b/comfy_api_nodes/nodes_qwen.py @@ -0,0 +1,442 @@ +import math +import re + +import torch +from typing_extensions import override + +from comfy_api.latest import IO, ComfyExtension +from comfy_api_nodes.apis.qwen import ( + QwenImageContentItem, + QwenImageGenerationRequest, + QwenImageGenerationResponse, + QwenImageInputField, + QwenImageMessage, + QwenImageParametersField, +) +from comfy_api_nodes.util import ( + ApiEndpoint, + download_url_to_image_tensor, + sync_op, + tensor_to_base64_string, + validate_string, +) + +GENERATION_PATH = "/proxy/qwen/api/v1/services/aigc/multimodal-generation/generation" +QWEN_IMAGE_MODELS = ["qwen-image-3.0-pro", "qwen-image-3.0"] +MIN_AREA = 262144 # 512*512 +MAX_AREA = 6553600 # 2560*2560 +MAX_ASPECT = 8 # the API allows aspect ratios from 1:8 to 8:1 +MAX_INPUT_BYTES = 10 * 1024 * 1024 # the API rejects decoded input images over 10MB + +_IMAGE_REF_RE = re.compile(r"@image(?P\d*)(?!\w)", re.IGNORECASE | re.ASCII) + + +def _resolve_image_refs(prompt: str, total_images: int) -> str: + """Rewrite @Image1-style references (shared partner-node syntax, 1-based; an unnumbered + @image means the first image) into the plain 'Image N' wording the model resolves + natively. A tag counts only at a word boundary or right after a previous tag, so + adjacent tags like '@Image1@Image2' all resolve while addresses like user@image1.com + pass through untouched.""" + parts = [] + pos = 0 + prev_end = -1 + for match in _IMAGE_REF_RE.finditer(prompt): + start = match.start() + if start > 0 and start != prev_end and (prompt[start - 1].isalnum() or prompt[start - 1] == "_"): + continue + idx = int(match.group("idx") or 1) + if not 1 <= idx <= total_images: + raise ValueError( + f"The prompt references @Image{idx}, but only {total_images} reference images " + f"are connected (a batched input counts once per image)." + ) + parts.append(prompt[pos:start]) + parts.append(f"Image {idx}") + pos = match.end() + prev_end = match.end() + parts.append(prompt[pos:]) + return "".join(parts) + + +def _validate_size(width: int, height: int) -> None: + if not MIN_AREA <= width * height <= MAX_AREA: + raise ValueError( + f"Image area must be between {MIN_AREA} (512x512) and {MAX_AREA} (2560x2560) pixels; " + f"got {width}x{height} = {width * height}." + ) + if width > MAX_ASPECT * height or height > MAX_ASPECT * width: + raise ValueError(f"Aspect ratio must be between 1:8 and 8:1; got {width}x{height}.") + + +def _fit_to_size(width: int, height: int) -> tuple[int, int]: + """Scale dimensions into the supported pixel area and 1:8..8:1 aspect range, preserving + the aspect ratio where possible.""" + if width > MAX_ASPECT * height: + height = math.ceil(width / MAX_ASPECT) + elif height > MAX_ASPECT * width: + width = math.ceil(height / MAX_ASPECT) + area = width * height + if area < MIN_AREA: + scale = math.sqrt(MIN_AREA / area) + width, height = math.ceil(width * scale), math.ceil(height * scale) + elif area > MAX_AREA: + scale = math.sqrt(MAX_AREA / area) + width, height = math.floor(width * scale), math.floor(height * scale) + # rounding can push the ratio a hair past the limit; trimming only ever shrinks the area + return min(width, MAX_ASPECT * height), min(height, MAX_ASPECT * width) + + +def _image_data_uri(image: torch.Tensor) -> str: + """PNG data URI of an RGB view of the image, downscaled to <=2048x2048; falls back to + JPEG when the PNG exceeds the API's decoded-size cap (e.g. noisy, incompressible images).""" + image = image[..., :3] + b64 = tensor_to_base64_string(image, total_pixels=2048 * 2048) + if len(b64) * 3 > MAX_INPUT_BYTES * 4: + return "data:image/jpeg;base64," + tensor_to_base64_string( + image, total_pixels=2048 * 2048, mime_type="image/jpeg" + ) + return "data:image/png;base64," + b64 + + +async def _download_result_images(response: QwenImageGenerationResponse) -> torch.Tensor: + if not response.output: + raise Exception(f"An unknown error occurred: {response.code} - {response.message}") + urls = [ + item.image + for choice in response.output.choices + if choice.message + for item in choice.message.content + if item.image + ] + if not urls: + raise Exception(f"The response contains no images: {response.code} - {response.message}") + return torch.cat([await download_url_to_image_tensor(url) for url in urls]) + + +def _size_inputs() -> list[IO.Int.Input]: + return [ + IO.Int.Input( + "width", + default=1024, + min=256, + max=2560, + step=16, + tooltip="The total pixel area must be between 512x512 and 2560x2560; " + "any aspect ratio within that area works.", + ), + IO.Int.Input( + "height", + default=1024, + min=256, + max=2560, + step=16, + tooltip="The total pixel area must be between 512x512 and 2560x2560; " + "any aspect ratio within that area works.", + ), + ] + + +def _t2i_model_option(model_id: str) -> IO.DynamicCombo.Option: + return IO.DynamicCombo.Option( + model_id, + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the image. Supports English and Chinese.", + ), + IO.String.Input( + "negative_prompt", + multiline=True, + default="", + tooltip="Negative prompt describing what to avoid.", + ), + *_size_inputs(), + ], + ) + + +def _edit_model_option(model_id: str) -> IO.DynamicCombo.Option: + return IO.DynamicCombo.Option( + model_id, + [ + IO.Autogrow.Input( + "images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("image"), + names=["image_1", "image_2", "image_3"], + min=1, + ), + tooltip="1-3 reference images. Refer to them in the prompt as @Image1, @Image2, " + "@Image3, numbered in input order; a batched input counts once per image.", + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Editing instructions. Supports English and Chinese, " + "and @Image1-style references to the input images.", + ), + IO.String.Input( + "negative_prompt", + multiline=True, + default="", + tooltip="Negative prompt describing what to avoid.", + ), + ], + ) + + +class QwenImageTextToImageApi(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="QwenImageTextToImageApi", + display_name="Qwen Image 3 Text to Image", + category="partner/image/Qwen", + description="Generates images from a text prompt using the Qwen-Image 3.0 models.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[_t2i_model_option(model_id) for model_id in QWEN_IMAGE_MODELS], + tooltip="Model to use.", + ), + IO.Int.Input( + "n", + default=1, + min=1, + max=6, + display_mode=IO.NumberDisplay.number, + tooltip="Number of images to generate, returned as a batch.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed to use for generation.", + ), + IO.Boolean.Input( + "prompt_extend", + default=True, + tooltip="Whether to enhance the prompt with AI assistance.", + advanced=True, + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AI-generated watermark to the result.", + advanced=True, + ), + ], + outputs=[ + IO.Image.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model", "model.width", "model.height", "n"]), + expr=""" + ( + $isPro := widgets.model = "qwen-image-3.0-pro"; + $area := $lookup(widgets, "model.width") * $lookup(widgets, "model.height"); + $rate := $isPro ? ($area > 2250000 ? 0.10725 : 0.0572) : 0.0429; + {"type":"usd","usd": $rate * widgets.n} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + n: int = 1, + seed: int = 42, + prompt_extend: bool = True, + watermark: bool = False, + ): + validate_string(model["prompt"], strip_whitespace=False, min_length=1) + width, height = model["width"], model["height"] + _validate_size(width, height) + response = await sync_op( + cls, + ApiEndpoint(path=GENERATION_PATH, method="POST"), + response_model=QwenImageGenerationResponse, + data=QwenImageGenerationRequest( + model=model["model"], + input=QwenImageInputField( + messages=[QwenImageMessage(content=[QwenImageContentItem(text=model["prompt"])])], + ), + parameters=QwenImageParametersField( + size=f"{width}*{height}", + n=n, + seed=seed, + prompt_extend=prompt_extend, + watermark=watermark, + negative_prompt=model["negative_prompt"] or None, + ), + ), + ) + return IO.NodeOutput(await _download_result_images(response)) + + +class QwenImageEditApi(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="QwenImageEditApi", + display_name="Qwen Image 3 Edit", + category="partner/image/Qwen", + description="Edits or combines up to 3 reference images guided by a text prompt " + "using the Qwen-Image 3.0 models.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[_edit_model_option(model_id) for model_id in QWEN_IMAGE_MODELS], + tooltip="Model to use.", + ), + IO.DynamicCombo.Input( + "size", + options=[ + IO.DynamicCombo.Option("match input", []), + IO.DynamicCombo.Option("auto", []), + IO.DynamicCombo.Option("custom", _size_inputs()), + ], + tooltip="Output resolution. 'match input' reuses the first reference image's size, " + "'auto' lets the model pick a size with the same aspect ratio, " + "'custom' sets an explicit width and height.", + ), + IO.Int.Input( + "n", + default=1, + min=1, + max=6, + display_mode=IO.NumberDisplay.number, + tooltip="Number of images to generate, returned as a batch.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed to use for generation.", + ), + IO.Boolean.Input( + "prompt_extend", + default=True, + tooltip="Whether to enhance the prompt with AI assistance.", + advanced=True, + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AI-generated watermark to the result.", + advanced=True, + ), + ], + outputs=[ + IO.Image.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends( + widgets=["model", "size", "size.width", "size.height", "n"], + input_groups=["model.images"], + ), + expr=""" + ( + $isPro := widgets.model = "qwen-image-3.0-pro"; + $mode := widgets.size; + $count := $max([$lookup(inputGroups, "model.images"), 1]); + $inputCost := 0.00429 * $count; + $area := $mode = "custom" + ? $lookup(widgets, "size.width") * $lookup(widgets, "size.height") : 0; + $customRate := $area > 2250000 ? 0.10725 : 0.0572; + $isPro and $mode != "custom" + ? {"type":"range_usd", + "min_usd": 0.0572 * widgets.n + $inputCost, + "max_usd": 0.10725 * widgets.n + $inputCost} + : {"type":"usd", + "usd": ($isPro ? $customRate : 0.0429) * widgets.n + $inputCost} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + size: dict, + n: int = 1, + seed: int = 42, + prompt_extend: bool = True, + watermark: bool = False, + ): + validate_string(model["prompt"], strip_whitespace=False, min_length=1) + reference_images = [image for key in model["images"] for image in model["images"][key]] + if len(reference_images) > 3: + raise ValueError( + f"A maximum of 3 reference images is supported; got {len(reference_images)} " + f"(a batched input counts once per image)." + ) + prompt = _resolve_image_refs(model["prompt"], len(reference_images)) + if size["size"] == "custom": + _validate_size(size["width"], size["height"]) + size_str = f"{size['width']}*{size['height']}" + elif size["size"] == "match input": + height, width = reference_images[0].shape[0], reference_images[0].shape[1] + width, height = _fit_to_size(width, height) + size_str = f"{width}*{height}" + else: # auto: the API picks a size preserving the input aspect ratio (1.9-4.2 MP) + size_str = None + content = [QwenImageContentItem(image=_image_data_uri(image)) for image in reference_images] + content.append(QwenImageContentItem(text=prompt)) + response = await sync_op( + cls, + ApiEndpoint(path=GENERATION_PATH, method="POST"), + response_model=QwenImageGenerationResponse, + data=QwenImageGenerationRequest( + model=model["model"], + input=QwenImageInputField(messages=[QwenImageMessage(content=content)]), + parameters=QwenImageParametersField( + size=size_str, + n=n, + seed=seed, + prompt_extend=prompt_extend, + watermark=watermark, + negative_prompt=model["negative_prompt"] or None, + ), + ), + ) + return IO.NodeOutput(await _download_result_images(response)) + + +class QwenApiExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[IO.ComfyNode]]: + return [ + QwenImageTextToImageApi, + QwenImageEditApi, + ] + + +async def comfy_entrypoint() -> QwenApiExtension: + return QwenApiExtension() From 34744cd29eacea9bbdec17e628a81c2ce0737d16 Mon Sep 17 00:00:00 2001 From: Simon Pinfold Date: Mon, 10 Aug 2026 14:05:21 -0700 Subject: [PATCH 11/76] Add tags_all / tags_any / tags_none tag filters to the assets list API (#15332) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Implement tags_all/tags_any/tags_none on the assets list API (BE-6600) Adds the three canonically-named tag filter params to GET /api/assets and GET /api/assets/tags/refine: - tags_all: asset carries every tag (replaces include_tags) - tags_any: asset carries at least one tag (new) - tags_none: asset carries no tag (replaces exclude_tags) Clauses intersect; tags_none always wins. include_tags/exclude_tags remain as permanent deprecated aliases and behave exactly as before when used on their own. Invalid combinations return 400 INVALID_TAG_FILTER, but only when the request uses at least one new-name parameter (non-empty after normalisation): - mixed spellings of one slot (include_tags with tags_all, exclude_tags with tags_none) - the same tag in the effective all-list and none-list (query can never match) Old-names-only requests gain no new error paths: include_tags=a&exclude_tags=a still returns an empty 200. tags_any/tags_none overlap stays valid (dead term, not a dead query). * Address review findings: positional-compat, deprecation metadata, test matrix - Move any_tags to the end of the four touched signatures: inserting it mid-signature silently misbound pre-existing positional callers (e.g. a caller passing name_contains positionally would have it consumed as any_tags). - Mark include_tags/exclude_tags Field(deprecated=True) on both list schemas so generated schema metadata matches the contract, not just a comment (schemas_out.py already uses this form for Asset.name). - Add tests: legal cross-slot old/new combinations, repeated query-key concatenation (pins Core behavior; outside the cross-platform contract), tags_any two-page cursor consistency (total/has_more/ no-overlap), refine-route mixed-spelling rejection + legacy-conflict preservation, and schema deprecation metadata. * Pin tag-value opacity: case-sensitive matching, byte-exact conflict check The prod tag survey (~/comfy/prod-model-tag-shape.md) found live case-distinct tag pairs (SEEDVR2/seedvr2) that resolve differently, so the contract now states tag values are opaque byte-strings. Pin that: case-distinct tags filter separately, and a case-distinct all/none pair is not an INVALID_TAG_FILTER conflict. * Document tags_all/tags_any/tags_none in openapi.yaml, deprecate aliases Add the three tag-filter parameters to both listAssets and getAssetTagHistogram parameter blocks and mark include_tags/exclude_tags deprecated: true, keeping the spec in step with the runtime schemas so generated clients can discover the new filters while the aliases stay present for existing consumers. * Move schemas_in import to module scope in test_list_filter Review feedback: no import cycle requires the local import. * Silence per-request DeprecationWarning in the tag-filter remap shim Reading the deprecated include_tags/exclude_tags fields by attribute fires pydantic's DeprecationWarning on every list/refine request even for callers using only the new names. The warning is aimed at API clients, not the server's own remap; read via model_dump instead. * Cap tag-filter lists at 100 entries, all spellings Review finding: unbounded tag lists fan out into one correlated EXISTS per tag on both page and count statements. Cap each list at 100 normalized entries with 400 INVALID_TAG_FILTER naming the parameter. Applies to the legacy spellings as well — a deliberate, decided exception to the old-names-behave-identically rule, since a cap only on new names would leave the same fan-out reachable through the aliases. * Strip process narration from comments Comments carried decision dates, contract cross-references, and review context. Keep only the constraints the code cannot show, one line each. --- app/assets/api/routes.py | 103 ++++- app/assets/api/schemas_in.py | 26 +- .../database/queries/asset_reference.py | 6 +- app/assets/database/queries/common.py | 13 +- app/assets/database/queries/tags.py | 4 +- app/assets/services/asset_management.py | 3 + app/assets/services/tagging.py | 3 + openapi.yaml | 66 ++- tests-unit/assets_test/test_list_filter.py | 419 ++++++++++++++++++ 9 files changed, 624 insertions(+), 19 deletions(-) diff --git a/app/assets/api/routes.py b/app/assets/api/routes.py index e25b8a57f..d43485861 100644 --- a/app/assets/api/routes.py +++ b/app/assets/api/routes.py @@ -18,7 +18,7 @@ from app.assets.api.schemas_in import ( AssetValidationError, UploadError, ) -from app.assets.helpers import validate_blake3_hash +from app.assets.helpers import normalize_tags, validate_blake3_hash from app.assets.api.upload import ( delete_temp_file_if_exists, parse_multipart_upload, @@ -117,6 +117,87 @@ def _build_validation_error_response(code: str, ve: ValidationError) -> web.Resp return _build_error_response(400, code, "Validation failed.", {"errors": errors}) +class InvalidTagFilterError(Exception): + """Invalid combination of tag-filter query parameters.""" + + def __init__(self, message: str, details: dict): + super().__init__(message) + self.details = details + + +# Caps the per-tag EXISTS fan-out; deliberately covers the legacy spellings too. +MAX_TAG_FILTER_TAGS = 100 + + +def _resolve_tag_filters( + q: schemas_in.ListAssetsQuery | schemas_in.TagsRefineQuery, +) -> tuple[list[str], list[str], list[str]]: + """Resolve legacy (include/exclude) and new (all/any/none) tag-filter + spellings into effective (all, any, none) lists. + + Combination validation applies only when the request uses at least one + new-name parameter (non-empty after normalisation); requests using only + the legacy names keep their historical behaviour, including degenerate + combinations like include_tags=a&exclude_tags=a. + """ + # model_dump, not attribute access: deprecated fields warn on every attribute read. + legacy = q.model_dump(include={"include_tags", "exclude_tags"}) + include_tags = normalize_tags(legacy["include_tags"]) + exclude_tags = normalize_tags(legacy["exclude_tags"]) + tags_all = normalize_tags(q.tags_all) + tags_any = normalize_tags(q.tags_any) + tags_none = normalize_tags(q.tags_none) + + for param_name, values in ( + ("include_tags", include_tags), + ("exclude_tags", exclude_tags), + ("tags_all", tags_all), + ("tags_any", tags_any), + ("tags_none", tags_none), + ): + if len(values) > MAX_TAG_FILTER_TAGS: + raise InvalidTagFilterError( + f"'{param_name}' lists {len(values)} tags; the maximum is " + f"{MAX_TAG_FILTER_TAGS}.", + { + "parameter": param_name, + "count": len(values), + "max": MAX_TAG_FILTER_TAGS, + }, + ) + + if not (tags_all or tags_any or tags_none): + return include_tags, [], exclude_tags + + if include_tags and tags_all: + raise InvalidTagFilterError( + "Cannot combine 'include_tags' with 'tags_all'; use 'tags_all'.", + {"parameters": ["include_tags", "tags_all"]}, + ) + if exclude_tags and tags_none: + raise InvalidTagFilterError( + "Cannot combine 'exclude_tags' with 'tags_none'; use 'tags_none'.", + {"parameters": ["exclude_tags", "tags_none"]}, + ) + + all_param, all_list = ( + ("tags_all", tags_all) if tags_all else ("include_tags", include_tags) + ) + none_param, none_list = ( + ("tags_none", tags_none) if tags_none else ("exclude_tags", exclude_tags) + ) + + conflicting = sorted(set(all_list) & set(none_list)) + if conflicting: + raise InvalidTagFilterError( + f"Query can never match: {', '.join(repr(t) for t in conflicting)} " + f"required by '{all_param}' but rejected by '{none_param}'.", + {"conflicting_tags": conflicting, "parameters": [all_param, none_param]}, + ) + + return all_list, tags_any, none_list + + def _validate_sort_field(requested: str | None) -> str: if not requested: return "created_at" @@ -217,6 +298,11 @@ async def list_assets_route(request: web.Request) -> web.Response: except ValidationError as ve: return _build_validation_error_response("INVALID_QUERY", ve) + try: + tags_all, tags_any, tags_none = _resolve_tag_filters(q) + except InvalidTagFilterError as e: + return _build_error_response(400, "INVALID_TAG_FILTER", str(e), e.details) + sort = _validate_sort_field(q.sort) order_candidate = (q.order or "desc").lower() order = order_candidate if order_candidate in {"asc", "desc"} else "desc" @@ -224,8 +310,9 @@ async def list_assets_route(request: web.Request) -> web.Response: try: result = list_assets_page( owner_id=USER_MANAGER.get_request_user_id(request), - include_tags=q.include_tags, - exclude_tags=q.exclude_tags, + include_tags=tags_all, + exclude_tags=tags_none, + any_tags=tags_any, name_contains=q.name_contains, metadata_filter=q.metadata_filter, limit=q.limit, @@ -715,10 +802,16 @@ async def get_tags_refine(request: web.Request) -> web.Response: except ValidationError as ve: return _build_validation_error_response("INVALID_QUERY", ve) + try: + tags_all, tags_any, tags_none = _resolve_tag_filters(q) + except InvalidTagFilterError as e: + return _build_error_response(400, "INVALID_TAG_FILTER", str(e), e.details) + tag_counts = list_tag_histogram( owner_id=USER_MANAGER.get_request_user_id(request), - include_tags=q.include_tags, - exclude_tags=q.exclude_tags, + include_tags=tags_all, + exclude_tags=tags_none, + any_tags=tags_any, name_contains=q.name_contains, metadata_filter=q.metadata_filter, limit=q.limit, diff --git a/app/assets/api/schemas_in.py b/app/assets/api/schemas_in.py index 38a942b7b..862700a24 100644 --- a/app/assets/api/schemas_in.py +++ b/app/assets/api/schemas_in.py @@ -50,8 +50,12 @@ class ParsedUpload: class ListAssetsQuery(BaseModel): - include_tags: list[str] = Field(default_factory=list) - exclude_tags: list[str] = Field(default_factory=list) + # Deprecated spellings: include_tags ≡ tags_all, exclude_tags ≡ tags_none. + include_tags: list[str] = Field(default_factory=list, deprecated=True) + exclude_tags: list[str] = Field(default_factory=list, deprecated=True) + tags_all: list[str] = Field(default_factory=list) + tags_any: list[str] = Field(default_factory=list) + tags_none: list[str] = Field(default_factory=list) name_contains: str | None = None # Accept either a JSON string (query param) or a dict @@ -70,7 +74,10 @@ class ListAssetsQuery(BaseModel): ) order: Literal["asc", "desc"] = "desc" - @field_validator("include_tags", "exclude_tags", mode="before") + @field_validator( + "include_tags", "exclude_tags", "tags_all", "tags_any", "tags_none", + mode="before", + ) @classmethod def _split_csv_tags(cls, v): # Accept "a,b,c" or ["a","b"] (we are liberal in what we accept) @@ -154,13 +161,20 @@ class CreateFromHashBody(BaseModel): class TagsRefineQuery(BaseModel): - include_tags: list[str] = Field(default_factory=list) - exclude_tags: list[str] = Field(default_factory=list) + # Deprecated spellings: include_tags ≡ tags_all, exclude_tags ≡ tags_none. + include_tags: list[str] = Field(default_factory=list, deprecated=True) + exclude_tags: list[str] = Field(default_factory=list, deprecated=True) + tags_all: list[str] = Field(default_factory=list) + tags_any: list[str] = Field(default_factory=list) + tags_none: list[str] = Field(default_factory=list) name_contains: str | None = None metadata_filter: dict[str, Any] | None = None limit: conint(ge=1, le=1000) = 100 - @field_validator("include_tags", "exclude_tags", mode="before") + @field_validator( + "include_tags", "exclude_tags", "tags_all", "tags_any", "tags_none", + mode="before", + ) @classmethod def _split_csv_tags(cls, v): if v is None: diff --git a/app/assets/database/queries/asset_reference.py b/app/assets/database/queries/asset_reference.py index 967b0e43a..126a0c9a4 100644 --- a/app/assets/database/queries/asset_reference.py +++ b/app/assets/database/queries/asset_reference.py @@ -268,6 +268,8 @@ def list_references_page( order: str | None = None, after_cursor_value: object | None = None, after_cursor_id: str | None = None, + # Appended last so pre-existing positional callers keep binding correctly. + any_tags: Sequence[str] | None = None, ) -> tuple[list[AssetReference], dict[str, list[str]], int]: """List references with pagination, filtering, and sorting. @@ -293,7 +295,7 @@ def list_references_page( escaped, esc = escape_sql_like_string(name_contains) base = base.where(AssetReference.name.ilike(f"%{escaped}%", escape=esc)) - base = apply_tag_filters(base, include_tags, exclude_tags) + base = apply_tag_filters(base, include_tags, exclude_tags, any_tags) base = apply_metadata_filter(base, metadata_filter) sort = (sort or "created_at").lower() @@ -345,7 +347,7 @@ def list_references_page( count_stmt = count_stmt.where( AssetReference.name.ilike(f"%{escaped}%", escape=esc) ) - count_stmt = apply_tag_filters(count_stmt, include_tags, exclude_tags) + count_stmt = apply_tag_filters(count_stmt, include_tags, exclude_tags, any_tags) count_stmt = apply_metadata_filter(count_stmt, metadata_filter) total = int(session.execute(count_stmt).scalar_one() or 0) diff --git a/app/assets/database/queries/common.py b/app/assets/database/queries/common.py index 89bb49327..7b0c211a0 100644 --- a/app/assets/database/queries/common.py +++ b/app/assets/database/queries/common.py @@ -60,10 +60,13 @@ def apply_tag_filters( stmt: sa.sql.Select, include_tags: Sequence[str] | None = None, exclude_tags: Sequence[str] | None = None, + any_tags: Sequence[str] | None = None, ) -> sa.sql.Select: - """include_tags: every tag must be present; exclude_tags: none may be present.""" + """include_tags: every tag must be present; any_tags: at least one must be + present; exclude_tags: none may be present.""" include_tags = normalize_tags(include_tags) exclude_tags = normalize_tags(exclude_tags) + any_tags = normalize_tags(any_tags) if include_tags: for tag_name in include_tags: @@ -74,6 +77,14 @@ def apply_tag_filters( ) ) + if any_tags: + stmt = stmt.where( + exists().where( + (AssetReferenceTag.asset_reference_id == AssetReference.id) + & (AssetReferenceTag.tag_name.in_(any_tags)) + ) + ) + if exclude_tags: stmt = stmt.where( ~exists().where( diff --git a/app/assets/database/queries/tags.py b/app/assets/database/queries/tags.py index 148f34801..e5f70e3df 100644 --- a/app/assets/database/queries/tags.py +++ b/app/assets/database/queries/tags.py @@ -340,6 +340,8 @@ def list_tag_counts_for_filtered_assets( name_contains: str | None = None, metadata_filter: dict | None = None, limit: int = 100, + # Appended last so pre-existing positional callers keep binding correctly. + any_tags: Sequence[str] | None = None, ) -> dict[str, int]: """Return tag counts for assets matching the given filters. @@ -359,7 +361,7 @@ def list_tag_counts_for_filtered_assets( escaped, esc = escape_sql_like_string(name_contains) ref_sq = ref_sq.where(AssetReference.name.ilike(f"%{escaped}%", escape=esc)) - ref_sq = apply_tag_filters(ref_sq, include_tags, exclude_tags) + ref_sq = apply_tag_filters(ref_sq, include_tags, exclude_tags, any_tags) ref_sq = apply_metadata_filter(ref_sq, metadata_filter) ref_sq = ref_sq.subquery() diff --git a/app/assets/services/asset_management.py b/app/assets/services/asset_management.py index a4c8b5a75..efdfd31a8 100644 --- a/app/assets/services/asset_management.py +++ b/app/assets/services/asset_management.py @@ -279,6 +279,8 @@ def list_assets_page( sort: str = "created_at", order: str = "desc", after: str | None = None, + # Appended last so pre-existing positional callers keep binding correctly. + any_tags: Sequence[str] | None = None, ) -> ListAssetsResult: """List assets with optional cursor pagination. @@ -317,6 +319,7 @@ def list_assets_page( owner_id=owner_id, include_tags=include_tags, exclude_tags=exclude_tags, + any_tags=any_tags, name_contains=name_contains, metadata_filter=metadata_filter, limit=fetch_limit, diff --git a/app/assets/services/tagging.py b/app/assets/services/tagging.py index 5fa39d26a..69c8cf39c 100644 --- a/app/assets/services/tagging.py +++ b/app/assets/services/tagging.py @@ -85,6 +85,8 @@ def list_tag_histogram( name_contains: str | None = None, metadata_filter: dict | None = None, limit: int = 100, + # Appended last so pre-existing positional callers keep binding correctly. + any_tags: Sequence[str] | None = None, ) -> dict[str, int]: with create_session() as session: return list_tag_counts_for_filtered_assets( @@ -92,6 +94,7 @@ def list_tag_histogram( owner_id=owner_id, include_tags=include_tags, exclude_tags=exclude_tags, + any_tags=any_tags, name_contains=name_contains, metadata_filter=metadata_filter, limit=limit, diff --git a/openapi.yaml b/openapi.yaml index a50312226..59c659dd3 100644 --- a/openapi.yaml +++ b/openapi.yaml @@ -1521,7 +1521,8 @@ paths: Supports filtering by tags, name, metadata, and sorting options. operationId: listAssets parameters: - - description: Filter assets that have ALL of these tags + - deprecated: true + description: 'Deprecated alias of tags_all: filter assets that have ALL of these tags' explode: false in: query name: include_tags @@ -1530,7 +1531,8 @@ paths: type: string type: array style: form - - description: Exclude assets that have ANY of these tags + - deprecated: true + description: 'Deprecated alias of tags_none: exclude assets that have ANY of these tags' explode: false in: query name: exclude_tags @@ -1539,6 +1541,33 @@ paths: type: string type: array style: form + - description: Filter assets that have ALL of these tags + explode: false + in: query + name: tags_all + schema: + items: + type: string + type: array + style: form + - description: Filter assets that have AT LEAST ONE of these tags + explode: false + in: query + name: tags_any + schema: + items: + type: string + type: array + style: form + - description: Exclude assets that have ANY of these tags + explode: false + in: query + name: tags_none + schema: + items: + type: string + type: array + style: form - description: Filter assets where name contains this substring (case-insensitive) in: query name: name_contains @@ -2312,7 +2341,8 @@ paths: Only returns tags with non-zero counts (tags that exist on matching assets). operationId: getAssetTagHistogram parameters: - - description: Filter assets that have ALL of these tags + - deprecated: true + description: 'Deprecated alias of tags_all: filter assets that have ALL of these tags' explode: false in: query name: include_tags @@ -2321,7 +2351,8 @@ paths: type: string type: array style: form - - description: Exclude assets that have ANY of these tags + - deprecated: true + description: 'Deprecated alias of tags_none: exclude assets that have ANY of these tags' explode: false in: query name: exclude_tags @@ -2330,6 +2361,33 @@ paths: type: string type: array style: form + - description: Filter assets that have ALL of these tags + explode: false + in: query + name: tags_all + schema: + items: + type: string + type: array + style: form + - description: Filter assets that have AT LEAST ONE of these tags + explode: false + in: query + name: tags_any + schema: + items: + type: string + type: array + style: form + - description: Exclude assets that have ANY of these tags + explode: false + in: query + name: tags_none + schema: + items: + type: string + type: array + style: form - description: Filter assets where name contains this substring (case-insensitive) in: query name: name_contains diff --git a/tests-unit/assets_test/test_list_filter.py b/tests-unit/assets_test/test_list_filter.py index d1cba87b3..04da4f86a 100644 --- a/tests-unit/assets_test/test_list_filter.py +++ b/tests-unit/assets_test/test_list_filter.py @@ -1,10 +1,14 @@ import time import uuid +import warnings import pytest import requests from helpers import assert_hash_fields_consistent +from app.assets.api import routes as assets_routes +from app.assets.api import schemas_in + def test_list_assets_paging_and_sort(http: requests.Session, api_base: str, asset_factory, make_asset_bytes): names = ["a1_u.safetensors", "a2_u.safetensors", "a3_u.safetensors"] @@ -337,3 +341,418 @@ def test_list_assets_name_contains_literal_underscore( assert b["name"] not in names, "Underscore must be escaped — should not match 'fooxbar'" assert c["name"] not in names, "Underscore must be escaped — should not match 'foobar'" assert body["total"] == 1 + + +def test_list_assets_tags_any_alone(http, api_base, asset_factory, make_asset_bytes): + scope = f"lf-any-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + a = asset_factory("any_a.safetensors", [*t, f"{scope}-alpha"], {}, make_asset_bytes("any_a")) + b = asset_factory("any_b.safetensors", [*t, f"{scope}-beta"], {}, make_asset_bytes("any_b")) + c = asset_factory("any_c.safetensors", [*t, f"{scope}-gamma"], {}, make_asset_bytes("any_c")) + + r = http.get( + api_base + "/api/assets", + params={"tags_any": f"{scope}-alpha,{scope}-beta", "limit": "50"}, + timeout=120, + ) + body = r.json() + assert r.status_code == 200, body + names = [x["name"] for x in body["assets"]] + assert a["name"] in names + assert b["name"] in names + assert c["name"] not in names + + +def test_list_assets_tags_any_with_tags_all(http, api_base, asset_factory, make_asset_bytes): + scope = f"lf-anyall-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + alpha, beta = f"{scope}-alpha", f"{scope}-beta" + x = asset_factory("aa_x.safetensors", [*t, alpha], {}, make_asset_bytes("aa_x")) + y = asset_factory("aa_y.safetensors", [*t, beta], {}, make_asset_bytes("aa_y")) + w = asset_factory("aa_w.safetensors", t, {}, make_asset_bytes("aa_w")) + d = asset_factory( + "aa_d.safetensors", + ["models", "model_type:checkpoints", "unit-tests", f"{scope}-other", alpha], + {}, + make_asset_bytes("aa_d"), + ) + + r = http.get( + api_base + "/api/assets", + params={"tags_all": f"unit-tests,{scope}", "tags_any": f"{alpha},{beta}", "limit": "50"}, + timeout=120, + ) + body = r.json() + assert r.status_code == 200, body + names = [a["name"] for a in body["assets"]] + assert x["name"] in names + assert y["name"] in names + assert w["name"] not in names, "asset matching tags_all but not tags_any must be excluded" + assert d["name"] not in names, "asset matching tags_any but not tags_all must be excluded" + + +def test_list_assets_tags_none_wins_over_tags_any(http, api_base, asset_factory, make_asset_bytes): + scope = f"lf-nonewins-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + alpha, beta = f"{scope}-alpha", f"{scope}-beta" + x = asset_factory("nw_x.safetensors", [*t, alpha], {}, make_asset_bytes("nw_x")) + y = asset_factory("nw_y.safetensors", [*t, alpha, beta], {}, make_asset_bytes("nw_y")) + + r = http.get( + api_base + "/api/assets", + params={"tags_any": alpha, "tags_none": beta, "limit": "50"}, + timeout=120, + ) + body = r.json() + assert r.status_code == 200, body + names = [a["name"] for a in body["assets"]] + assert x["name"] in names + assert y["name"] not in names, "tags_none must exclude an asset even when it matches tags_any" + + +def test_list_assets_empty_tag_filter_lists_behave_as_absent(http, api_base, asset_factory, make_asset_bytes): + scope = f"lf-empty-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + a = asset_factory("em_a.safetensors", t, {}, make_asset_bytes("em_a")) + b = asset_factory("em_b.safetensors", t, {}, make_asset_bytes("em_b")) + expected = {a["name"], b["name"]} + + # Empty new-name lists impose no constraint. + r1 = http.get( + api_base + "/api/assets", + params={"tags_all": f"unit-tests,{scope}", "tags_any": "", "tags_none": ""}, + timeout=120, + ) + b1 = r1.json() + assert r1.status_code == 200, b1 + assert {x["name"] for x in b1["assets"]} == expected + + # An empty new-name param alongside old names must not trigger validation. + r2 = http.get( + api_base + "/api/assets", + params={"include_tags": f"unit-tests,{scope}", "tags_any": ""}, + timeout=120, + ) + b2 = r2.json() + assert r2.status_code == 200, b2 + assert {x["name"] for x in b2["assets"]} == expected + + # An empty tags_all next to include_tags is not a mixed-spelling conflict. + r3 = http.get( + api_base + "/api/assets", + params={"include_tags": f"unit-tests,{scope}", "tags_all": ""}, + timeout=120, + ) + b3 = r3.json() + assert r3.status_code == 200, b3 + assert {x["name"] for x in b3["assets"]} == expected + + +def test_list_assets_old_names_match_new_names(http, api_base, asset_factory, make_asset_bytes): + scope = f"lf-alias-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + alpha, beta = f"{scope}-alpha", f"{scope}-beta" + asset_factory("al_a.safetensors", [*t, alpha], {}, make_asset_bytes("al_a")) + asset_factory("al_b.safetensors", [*t, beta], {}, make_asset_bytes("al_b")) + + def names_for(params: dict) -> tuple[list, int]: + r = http.get(api_base + "/api/assets", params={**params, "sort": "name", "order": "asc"}, timeout=120) + body = r.json() + assert r.status_code == 200, body + return [x["name"] for x in body["assets"]], body["total"] + + # include_tags ≡ tags_all + old_names, old_total = names_for({"include_tags": f"unit-tests,{scope}"}) + new_names, new_total = names_for({"tags_all": f"unit-tests,{scope}"}) + assert old_names == new_names + assert old_total == new_total + + # exclude_tags ≡ tags_none (and old/new spellings mix across slots) + old_names, old_total = names_for({"include_tags": f"unit-tests,{scope}", "exclude_tags": alpha}) + new_names, new_total = names_for({"tags_all": f"unit-tests,{scope}", "tags_none": alpha}) + mixed_names, mixed_total = names_for({"include_tags": f"unit-tests,{scope}", "tags_none": alpha}) + assert old_names == new_names == mixed_names == ["al_b.safetensors"] + assert old_total == new_total == mixed_total == 1 + + +@pytest.mark.parametrize( + "params,expected_parameters", + [ + ({"include_tags": "mx-x", "tags_all": "mx-y"}, ["include_tags", "tags_all"]), + ({"exclude_tags": "mx-x", "tags_none": "mx-y"}, ["exclude_tags", "tags_none"]), + ], + ids=["include_tags_with_tags_all", "exclude_tags_with_tags_none"], +) +def test_list_assets_mixed_tag_spellings_rejected(http, api_base, params, expected_parameters): + r = http.get(api_base + "/api/assets", params=params, timeout=120) + body = r.json() + assert r.status_code == 400, body + assert body["error"]["code"] == "INVALID_TAG_FILTER" + assert body["error"]["details"]["parameters"] == expected_parameters + + +@pytest.mark.parametrize( + "params,conflicting,parameters", + [ + ( + {"tags_all": "cf-x", "tags_none": "cf-x"}, + ["cf-x"], + ["tags_all", "tags_none"], + ), + ( + {"include_tags": "cf-x", "tags_none": "cf-x"}, + ["cf-x"], + ["include_tags", "tags_none"], + ), + ( + {"tags_all": "cf-a,cf-b", "tags_none": "cf-b,cf-c"}, + ["cf-b"], + ["tags_all", "tags_none"], + ), + ], + ids=["new_names", "include_tags_remapped", "partial_overlap"], +) +def test_list_assets_all_none_conflict_rejected(http, api_base, params, conflicting, parameters): + r = http.get(api_base + "/api/assets", params=params, timeout=120) + body = r.json() + assert r.status_code == 400, body + assert body["error"]["code"] == "INVALID_TAG_FILTER" + assert body["error"]["details"]["conflicting_tags"] == conflicting + assert body["error"]["details"]["parameters"] == parameters + + +def test_list_assets_any_none_overlap_accepted(http, api_base, asset_factory, make_asset_bytes): + scope = f"lf-deadterm-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + alpha, beta = f"{scope}-alpha", f"{scope}-beta" + x = asset_factory("dt_x.safetensors", [*t, alpha], {}, make_asset_bytes("dt_x")) + y = asset_factory("dt_y.safetensors", [*t, beta], {}, make_asset_bytes("dt_y")) + + # alpha is a dead term (in both tags_any and tags_none) but the query is valid. + r = http.get( + api_base + "/api/assets", + params={"tags_any": f"{alpha},{beta}", "tags_none": alpha, "limit": "50"}, + timeout=120, + ) + body = r.json() + assert r.status_code == 200, body + names = [a["name"] for a in body["assets"]] + assert y["name"] in names + assert x["name"] not in names + + +def test_list_assets_legacy_include_exclude_conflict_still_200(http, api_base, asset_factory, make_asset_bytes): + scope = f"lf-legacy-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + asset_factory("lg_a.safetensors", t, {}, make_asset_bytes("lg_a")) + + # Old names only: the self-contradictory query stays an empty 200, never a 400. + r = http.get( + api_base + "/api/assets", + params={"include_tags": scope, "exclude_tags": scope}, + timeout=120, + ) + body = r.json() + assert r.status_code == 200, body + assert body["assets"] == [] + + +def test_tags_refine_new_tag_filters(http, api_base, asset_factory, make_asset_bytes): + scope = f"rf-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + alpha, beta = f"{scope}-alpha", f"{scope}-beta" + asset_factory("rf_a.safetensors", [*t, alpha], {}, make_asset_bytes("rf_a")) + asset_factory("rf_b.safetensors", [*t, beta], {}, make_asset_bytes("rf_b")) + + r = http.get( + api_base + "/api/assets/tags/refine", + params={"tags_any": f"{alpha},{beta}", "tags_none": alpha}, + timeout=120, + ) + body = r.json() + assert r.status_code == 200, body + counts = body["tag_counts"] + assert counts.get(beta) == 1 + assert alpha not in counts + + r2 = http.get( + api_base + "/api/assets/tags/refine", + params={"tags_all": "rf-x", "tags_none": "rf-x"}, + timeout=120, + ) + body2 = r2.json() + assert r2.status_code == 400, body2 + assert body2["error"]["code"] == "INVALID_TAG_FILTER" + assert body2["error"]["details"]["conflicting_tags"] == ["rf-x"] + + +def test_list_assets_cross_slot_old_new_combinations(http, api_base, asset_factory, make_asset_bytes): + """Old and new spellings of *different* slots combine freely; only + same-slot mixing is rejected.""" + scope = f"lf-cross-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + alpha, beta = f"{scope}-alpha", f"{scope}-beta" + a = asset_factory("cs_a.safetensors", [*t, alpha], {}, make_asset_bytes("cs_a")) + b = asset_factory("cs_b.safetensors", [*t, beta], {}, make_asset_bytes("cs_b")) + + def names_for(params: dict) -> set: + r = http.get(api_base + "/api/assets", params=params, timeout=120) + body = r.json() + assert r.status_code == 200, body + return {x["name"] for x in body["assets"]} + + assert names_for( + {"include_tags": f"unit-tests,{scope}", "tags_any": alpha} + ) == {a["name"]} + assert names_for( + {"tags_all": f"unit-tests,{scope}", "exclude_tags": alpha} + ) == {b["name"]} + assert names_for( + {"tags_any": f"{alpha},{beta}", "exclude_tags": alpha} + ) == {b["name"]} + + +def test_list_assets_repeated_query_keys_concatenate(http, api_base, asset_factory, make_asset_bytes): + """Repeated occurrences of a tag param concatenate before the CSV split + (Core-local behavior, not a cross-platform guarantee).""" + scope = f"lf-repeat-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + alpha, beta = f"{scope}-alpha", f"{scope}-beta" + a = asset_factory("rp_a.safetensors", [*t, alpha], {}, make_asset_bytes("rp_a")) + b = asset_factory("rp_b.safetensors", [*t, beta], {}, make_asset_bytes("rp_b")) + + # requests encodes a list value as repeated keys: tags_any=&tags_any= + r = http.get( + api_base + "/api/assets", + params={"tags_any": [alpha, beta], "limit": "50"}, + timeout=120, + ) + body = r.json() + assert r.status_code == 200, body + names = {x["name"] for x in body["assets"]} + assert {a["name"], b["name"]} <= names + + +def test_list_assets_tags_any_cursor_pagination_consistent(http, api_base, asset_factory, make_asset_bytes): + scope = f"lf-anypage-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + alpha = f"{scope}-alpha" + expected = set() + for i in range(3): + made = asset_factory(f"pg_{i}.safetensors", [*t, alpha], {}, make_asset_bytes(f"pg_{i}")) + expected.add(made["name"]) + + r1 = http.get( + api_base + "/api/assets", + params={"tags_any": alpha, "limit": "2", "sort": "name", "order": "asc"}, + timeout=120, + ) + b1 = r1.json() + assert r1.status_code == 200, b1 + assert b1["total"] == 3 + assert b1["has_more"] is True + assert b1.get("next_cursor"), "expected a keyset cursor on the first page" + + r2 = http.get( + api_base + "/api/assets", + params={ + "tags_any": alpha, + "limit": "2", + "sort": "name", + "order": "asc", + "after": b1["next_cursor"], + }, + timeout=120, + ) + b2 = r2.json() + assert r2.status_code == 200, b2 + assert b2["has_more"] is False + + page1 = {x["name"] for x in b1["assets"]} + page2 = {x["name"] for x in b2["assets"]} + assert not page1 & page2, "cursor pages must not overlap" + assert page1 | page2 == expected + + +def test_tags_refine_mixed_spellings_rejected_and_legacy_conflict_kept(http, api_base): + r = http.get( + api_base + "/api/assets/tags/refine", + params={"include_tags": "rfmx-x", "tags_all": "rfmx-y"}, + timeout=120, + ) + body = r.json() + assert r.status_code == 400, body + assert body["error"]["code"] == "INVALID_TAG_FILTER" + assert body["error"]["details"]["parameters"] == ["include_tags", "tags_all"] + + # Old names only: the refine route keeps legacy behaviour too. + r2 = http.get( + api_base + "/api/assets/tags/refine", + params={"include_tags": "rfmx-z", "exclude_tags": "rfmx-z"}, + timeout=120, + ) + body2 = r2.json() + assert r2.status_code == 200, body2 + assert body2["tag_counts"] == {} + + +def test_list_assets_tag_values_case_sensitive(http, api_base, asset_factory, make_asset_bytes): + """Case-distinct tags are distinct; the all/none conflict check is byte-exact.""" + scope = f"lf-case-{uuid.uuid4().hex[:6]}" + t = ["models", "model_type:checkpoints", "unit-tests", scope] + upper, lower = f"{scope}-ALPHA", f"{scope}-alpha" + a = asset_factory("cx_a.safetensors", [*t, upper], {}, make_asset_bytes("cx_a")) + b = asset_factory("cx_b.safetensors", [*t, lower], {}, make_asset_bytes("cx_b")) + + def names_for(params: dict) -> set: + r = http.get(api_base + "/api/assets", params=params, timeout=120) + body = r.json() + assert r.status_code == 200, body + return {x["name"] for x in body["assets"]} + + assert names_for({"tags_all": f"unit-tests,{scope},{upper}"}) == {a["name"]} + assert names_for({"tags_any": lower, "limit": "50"}) == {b["name"]} + # Case-distinct all/none pair is NOT a conflict — byte-exact comparison. + assert names_for({"tags_all": f"unit-tests,{scope},{upper}", "tags_none": lower}) == {a["name"]} + + +def test_tag_list_cap_applies_to_all_spellings(http, api_base): + """The cap covers the legacy spellings too.""" + big = ",".join(f"cap-{i}" for i in range(101)) + for param in ("tags_any", "include_tags"): + r = http.get(api_base + "/api/assets", params={param: big}, timeout=120) + body = r.json() + assert r.status_code == 400, body + assert body["error"]["code"] == "INVALID_TAG_FILTER" + assert body["error"]["details"]["parameter"] == param + assert body["error"]["details"]["max"] == 100 + + exact = ",".join(f"cap-{i}" for i in range(100)) + r = http.get(api_base + "/api/assets", params={"tags_any": exact}, timeout=120) + assert r.status_code == 200, r.json() + + # The cap counts normalized (deduped) tags, not raw CSV items. + dups = ",".join("cap-dup" for _ in range(150)) + r = http.get(api_base + "/api/assets", params={"tags_any": dups}, timeout=120) + assert r.status_code == 200, r.json() + + +def test_resolve_tag_filters_no_deprecation_warning(): + """The deprecated-field warning is for API clients; the server's own remap + shim must not fire it on every request.""" + for q in ( + schemas_in.ListAssetsQuery(tags_all="a", tags_none="b"), + schemas_in.TagsRefineQuery(tags_any="c"), + ): + with warnings.catch_warnings(): + warnings.simplefilter("error", DeprecationWarning) + assets_routes._resolve_tag_filters(q) + + +def test_tag_filter_alias_fields_marked_deprecated(): + for model in (schemas_in.ListAssetsQuery, schemas_in.TagsRefineQuery): + props = model.model_json_schema()["properties"] + for field in ("include_tags", "exclude_tags"): + assert props[field].get("deprecated") is True, (model.__name__, field) + for field in ("tags_all", "tags_any", "tags_none"): + assert "deprecated" not in props[field], (model.__name__, field) From 6233790c6dff26bf35113d46d6d3367b7041b1d8 Mon Sep 17 00:00:00 2001 From: chelsealong Date: Tue, 11 Aug 2026 11:00:23 +0800 Subject: [PATCH 12/76] Fix VAEDecodeTiled crash on NestedTensor latents (MiniMax H3) (#15477) VAEDecode unwraps a NestedTensor latent (video/audio pair) to its video component before calling vae.decode(). VAEDecodeTiled skipped this unwrap and passed the NestedTensor straight into vae.decode_tiled(), which fails deep in the MiniMax H3 video VAE when a real tensor's .to() is called with the NestedTensor as an argument. Fixes #15468. --- nodes.py | 6 ++++- .../test_vae_decode_tiled_nested.py | 27 +++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) create mode 100644 tests-unit/comfy_test/test_vae_decode_tiled_nested.py diff --git a/nodes.py b/nodes.py index 432f04d89..a7f91720f 100644 --- a/nodes.py +++ b/nodes.py @@ -364,8 +364,12 @@ class VAEDecodeTiled: temporal_size = None temporal_overlap = None + latent = samples["samples"] + if latent.is_nested: + latent = latent.unbind()[0] + compression = vae.spacial_compression_decode() - images = vae.decode_tiled(samples["samples"], tile_x=tile_size // compression, tile_y=tile_size // compression, overlap=overlap // compression, tile_t=temporal_size, overlap_t=temporal_overlap) + images = vae.decode_tiled(latent, tile_x=tile_size // compression, tile_y=tile_size // compression, overlap=overlap // compression, tile_t=temporal_size, overlap_t=temporal_overlap) if len(images.shape) == 5: #Combine batches images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1]) return (images, ) diff --git a/tests-unit/comfy_test/test_vae_decode_tiled_nested.py b/tests-unit/comfy_test/test_vae_decode_tiled_nested.py new file mode 100644 index 000000000..7c4b345b9 --- /dev/null +++ b/tests-unit/comfy_test/test_vae_decode_tiled_nested.py @@ -0,0 +1,27 @@ +from unittest.mock import MagicMock + +import torch + +from comfy.cli_args import args as cli_args + +if not torch.cuda.is_available(): + cli_args.cpu = True + +import comfy.nested_tensor # noqa: E402 +import nodes # noqa: E402 + + +def test_vae_decode_tiled_unwraps_nested_tensor(): + video = torch.zeros(1, 4, 2, 8, 8) + audio = torch.zeros(1, 2, 2, 40) + samples = {"samples": comfy.nested_tensor.NestedTensor((video, audio))} + + vae = MagicMock() + vae.temporal_compression_decode.return_value = None + vae.spacial_compression_decode.return_value = 8 + vae.decode_tiled.return_value = torch.zeros(1, 3, 2, 8, 8) + + nodes.VAEDecodeTiled().decode(vae, samples, tile_size=512) + + decoded_arg = vae.decode_tiled.call_args[0][0] + assert decoded_arg is video From 4f3544d131652678c8070b306f01cce392465cb5 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Mon, 10 Aug 2026 20:25:01 -0700 Subject: [PATCH 13/76] Make cu130 warning more visible. (#15463) --- comfy/quant_ops.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy/quant_ops.py b/comfy/quant_ops.py index 6d9112dbb..18fd2d613 100644 --- a/comfy/quant_ops.py +++ b/comfy/quant_ops.py @@ -40,7 +40,7 @@ try: cuda_version = tuple(map(int, str(torch.version.cuda).split('.'))) if cuda_version < (13,): ck.registry.disable("cuda") - logging.warning("WARNING: You need pytorch with cu130 or higher to use optimized CUDA operations.") + logging.warning("WARNING: You need pytorch with cu130 or higher to use optimized CUDA operations.\nWARNING WARNING WARNING\nIf you are on nvidia 20 series and above it is required that you update your pytorch to cu130 or higher.\n") # On ROCm/AMD the CUDA backend is unavailable, so Triton is the only accelerated # comfy-kitchen backend. Enable it by default there, but only on Triton >= 3.7 AND a From bf4c9a08fc854df6d3b2bef1b92b509e2ef2d2c9 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Mon, 10 Aug 2026 22:03:08 -0700 Subject: [PATCH 14/76] Implement comfy kitchen attention. (#15479) Add a ModelAttentionBackend node to manually select the attention for models in the workflows. Currently supports pytorch attention or comfy kitchen attention. Add --use-ck-attention to enable comfy kitchen attention as the default attention backend for all models (might break some). --- comfy/cli_args.py | 1 + comfy/ldm/minimax/model.py | 8 +- comfy/ldm/modules/attention.py | 109 ++++++++++++++++++++++++++- comfy/model_management.py | 3 + comfy/model_patcher.py | 8 ++ comfy_extras/nodes_model_advanced.py | 37 +++++++++ requirements.txt | 2 +- 7 files changed, 162 insertions(+), 6 deletions(-) diff --git a/comfy/cli_args.py b/comfy/cli_args.py index ee9e1ce9f..9de244087 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -149,6 +149,7 @@ attn_group.add_argument("--use-quad-cross-attention", action="store_true", help= attn_group.add_argument("--use-pytorch-cross-attention", action="store_true", help="Use the new pytorch 2.0 cross attention function.") attn_group.add_argument("--use-sage-attention", action="store_true", help="Use sage attention.") attn_group.add_argument("--use-flash-attention", action="store_true", help="Use FlashAttention.") +attn_group.add_argument("--use-ck-attention", action="store_true", help="Use Comfy Kitchen attention.") parser.add_argument("--disable-xformers", action="store_true", help="Disable xformers.") diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py index bc06288ab..76174483a 100644 --- a/comfy/ldm/minimax/model.py +++ b/comfy/ldm/minimax/model.py @@ -25,7 +25,7 @@ import comfy.model_prefetch import comfy.ops import comfy.patcher_extension import comfy.quant_ops -from comfy.ldm.modules.attention import optimized_attention +from comfy.ldm.modules.attention import AttentionTensorContainer, optimized_attention FRAME_PER_TOKEN = (1, 4, 4, 4, 4) FRAME_RESCALE = 5.0 / 3.0 @@ -165,9 +165,9 @@ class Attention(nn.Module): else: q = self.q_norm(q.view(s, self.heads, self.head_dim)) k = self.k_norm(k.view(s, self.heads, self.head_dim)) - q = q.transpose(0, 1).unsqueeze(0) - k = k.transpose(0, 1).unsqueeze(0) - v = v.transpose(0, 1).unsqueeze(0) + q = AttentionTensorContainer(q.transpose(0, 1).unsqueeze(0)) + k = AttentionTensorContainer(k.transpose(0, 1).unsqueeze(0)) + v = AttentionTensorContainer(v.transpose(0, 1).unsqueeze(0)) out = optimized_attention(q, k, v, self.heads, mask=None, skip_reshape=True, transformer_options=transformer_options) return self.out_proj(out.squeeze(0)) diff --git a/comfy/ldm/modules/attention.py b/comfy/ldm/modules/attention.py index 2c549e095..b22d03d77 100644 --- a/comfy/ldm/modules/attention.py +++ b/comfy/ldm/modules/attention.py @@ -10,6 +10,8 @@ from typing import Optional, Any, Callable, Union import logging import functools +import comfy_kitchen + from .diffusionmodules.util import AlphaBlender, timestep_embedding from .sub_quadratic_attention import efficient_dot_product_attention @@ -49,6 +51,8 @@ except ImportError: logging.error(f"\n\nTo use the `--use-flash-attention` feature, the `flash-attn` package must be installed first.\ncommand:\n\t{sys.executable} -m pip install flash-attn") exit(-1) +COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE = comfy_kitchen.int8_attention_is_available() + REGISTERED_ATTENTION_FUNCTIONS = {} def register_attention_function(name: str, func: Callable): # avoid replacing existing functions @@ -145,9 +149,34 @@ def Normalize(in_channels, dtype=None, device=None): return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, dtype=dtype, device=device) +class AttentionTensorContainer: + """Single-owner tensor input consumed by an optimized attention backend.""" + + __slots__ = ("tensor",) + + def __init__(self, tensor: torch.Tensor): + self.tensor: torch.Tensor | None = tensor + + def peek(self) -> torch.Tensor: + if self.tensor is None: + raise RuntimeError("attention tensor container has already been consumed") + return self.tensor + + def take(self) -> torch.Tensor: + tensor = self.peek() + self.tensor = None + return tensor + + def wrap_attn(func): @functools.wraps(func) def wrapper(*args, **kwargs): + containers = None + if len(args) >= 3 and isinstance(args[0], AttentionTensorContainer): + if not isinstance(args[1], AttentionTensorContainer) or not isinstance(args[2], AttentionTensorContainer): + raise TypeError("q, k, and v must all be attention tensor containers") + containers = args[:3] + remove_attn_wrapper_key = False try: if "_inside_attn_wrapper" not in kwargs: @@ -156,11 +185,22 @@ def wrap_attn(func): kwargs["_inside_attn_wrapper"] = True if transformer_options is not None: if "optimized_attention_override" in transformer_options: - return transformer_options["optimized_attention_override"](func, *args, **kwargs) + optimized_attention_override = transformer_options["optimized_attention_override"] + if containers is not None: + if hasattr(optimized_attention_override, "container_function"): + return optimized_attention_override.container_function(*args, **kwargs) + args = tuple(container.take() for container in containers) + args[3:] + return optimized_attention_override(func, *args, **kwargs) + + if containers is not None: + if wrapper.container_function is not None: + return wrapper.container_function(*args, **kwargs) + args = tuple(container.take() for container in containers) + args[3:] return func(*args, **kwargs) finally: if remove_attn_wrapper_key: del kwargs["_inside_attn_wrapper"] + wrapper.container_function = None return wrapper @wrap_attn @@ -545,6 +585,63 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha ).transpose(1, 2).reshape(-1, q.shape[2], heads * dim_head) return out +def _comfy_kitchen_int8_inputs(q, k, v, heads, mask, skip_reshape, enable_gqa): + dim_head = q.shape[-1] if skip_reshape else q.shape[-1] // heads + b = q.shape[0] + if not skip_reshape: + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, enable_gqa, expand_kv=False) + q, k, v = map(lambda t: t.transpose(1, 2), (q, k, v)) + + if mask is not None: + if mask.ndim == 2: + mask = mask.unsqueeze(0) + if mask.ndim == 3: + mask = mask.unsqueeze(1) + + return q, k, v, mask, b, dim_head + + +@wrap_attn +def attention_comfy_kitchen_int8(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): + q, k, v, mask, b, dim_head = _comfy_kitchen_int8_inputs( + q, k, v, heads, mask, skip_reshape, kwargs.get("enable_gqa", False) + ) + out = comfy_kitchen.int8_attention( + q, + k, + v, + scale=kwargs.get("scale", None), + attn_mask=mask, + ) + if not skip_output_reshape: + out = out.transpose(1, 2).reshape(b, -1, heads * dim_head) + return out + + +def _attention_comfy_kitchen_int8_containers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): + q = q.take() + k = k.take() + v = v.take() + q, k, v, mask, b, dim_head = _comfy_kitchen_int8_inputs( + q, k, v, heads, mask, skip_reshape, kwargs.get("enable_gqa", False) + ) + quantized = comfy_kitchen.prequantize_int8_attention( + q, + k, + v, + scale=kwargs.get("scale", None), + attn_mask=mask, + ) + del q, k, v + out = comfy_kitchen.int8_attention_from_prequantized(quantized) + if not skip_output_reshape: + out = out.transpose(1, 2).reshape(b, -1, heads * dim_head) + return out + + +attention_comfy_kitchen_int8.container_function = _attention_comfy_kitchen_int8_containers + + @wrap_attn def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): if kwargs.get("low_precision_attention", True) is False or (mask is not None and not SAGE_ATTENTION_SUPPORTS_MASK): @@ -775,10 +872,20 @@ else: logging.info("Using sub quadratic optimization for attention, if you have memory or speed issues try using: --use-split-cross-attention") optimized_attention = attention_sub_quad +if model_management.comfy_kitchen_attention_enabled(): + if COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE: + logging.info("Using Comfy Kitchen attention") + optimized_attention = attention_comfy_kitchen_int8 + else: + logging.error("Comfy Kitchen attention is unavailable. Install a Comfy Kitchen build with attention support to use --use-ck-attention.") + exit(-1) + optimized_attention_masked = optimized_attention # register core-supported attention functions +if COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE: + register_attention_function("comfy_kitchen_int8", attention_comfy_kitchen_int8) if SAGE_ATTENTION_IS_AVAILABLE: register_attention_function("sage", attention_sage) if SAGE_ATTENTION3_IS_AVAILABLE: diff --git a/comfy/model_management.py b/comfy/model_management.py index 9f8e7f07b..65599424b 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -1658,6 +1658,9 @@ def unpin_memory(tensor): def sage_attention_enabled(): return args.use_sage_attention +def comfy_kitchen_attention_enabled(): + return args.use_ck_attention + def flash_attention_enabled(): return args.use_flash_attention diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index ae3f0191d..cb44e7394 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -685,6 +685,14 @@ class ModelPatcher: def set_model_attn2_output_patch(self, patch): self.set_model_patch(patch, "attn2_output_patch") + def set_model_optimized_attention(self, optimized_attention): + def optimized_attention_override(_, *args, **kwargs): + return optimized_attention(*args, **kwargs) + + if hasattr(optimized_attention, "container_function") and optimized_attention.container_function is not None: + optimized_attention_override.container_function = optimized_attention.container_function + self.model_options["transformer_options"]["optimized_attention_override"] = optimized_attention_override + def set_model_input_block_patch(self, patch): self.set_model_patch(patch, "input_block_patch") diff --git a/comfy_extras/nodes_model_advanced.py b/comfy_extras/nodes_model_advanced.py index a336ba079..21ea82148 100644 --- a/comfy_extras/nodes_model_advanced.py +++ b/comfy_extras/nodes_model_advanced.py @@ -1,6 +1,9 @@ +import logging + import comfy.sd import comfy.model_sampling import comfy.latent_formats +import comfy.ldm.modules.attention import nodes import torch import node_helpers @@ -346,6 +349,39 @@ class ModelComputeDtype: return (m, ) +class ModelAttentionBackend: + @classmethod + def INPUT_TYPES(s): + backends = ["pytorch attention"] + if comfy.ldm.modules.attention.COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE: + backends.append("comfy kitchen attention") + return {"required": {"model": ("MODEL",), + "attention": (backends,), + }} + + @classmethod + def VALIDATE_INPUTS(s, attention): + return True + + RETURN_TYPES = ("MODEL",) + FUNCTION = "patch" + + CATEGORY = "model/patch" + + def patch(self, model, attention): + attention_name = { + "comfy kitchen attention": "comfy_kitchen_int8", + "pytorch attention": "pytorch", + }.get(attention) + attention_function = comfy.ldm.modules.attention.get_attention_function(attention_name, None) + if attention_function is None: + logging.warning("Attention backend '%s' is unavailable; using PyTorch attention.", attention) + attention_function = comfy.ldm.modules.attention.get_attention_function("pytorch") + m = model.clone() + m.set_model_optimized_attention(attention_function) + return (m, ) + + NODE_CLASS_MAPPINGS = { "ModelSamplingDiscrete": ModelSamplingDiscrete, "ModelSamplingContinuousEDM": ModelSamplingContinuousEDM, @@ -357,4 +393,5 @@ NODE_CLASS_MAPPINGS = { "ModelNoiseScale": ModelNoiseScale, "RescaleCFG": RescaleCFG, "ModelComputeDtype": ModelComputeDtype, + "ModelAttentionBackend": ModelAttentionBackend, } diff --git a/requirements.txt b/requirements.txt index 94cd1c5eb..25eaf8bc7 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.28 +comfy-kitchen==0.2.30 comfy-aimdo==0.4.13 requests simpleeval>=1.0.0 From 62b3c94bd45154f6486c7abf1b9efcacee96ea69 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 11 Aug 2026 02:09:29 -0700 Subject: [PATCH 15/76] Fix peak memory issue with H3. (#15486) --- comfy/ldm/minimax/model.py | 1 + 1 file changed, 1 insertion(+) diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py index 76174483a..f745db884 100644 --- a/comfy/ldm/minimax/model.py +++ b/comfy/ldm/minimax/model.py @@ -165,6 +165,7 @@ class Attention(nn.Module): else: q = self.q_norm(q.view(s, self.heads, self.head_dim)) k = self.k_norm(k.view(s, self.heads, self.head_dim)) + v = v.clone() q = AttentionTensorContainer(q.transpose(0, 1).unsqueeze(0)) k = AttentionTensorContainer(k.transpose(0, 1).unsqueeze(0)) v = AttentionTensorContainer(v.transpose(0, 1).unsqueeze(0)) From 57ce8e1a27dda3d4cedab081081fd3e1359c195b Mon Sep 17 00:00:00 2001 From: Alexis Rolland Date: Tue, 11 Aug 2026 10:47:39 -0700 Subject: [PATCH 16/76] Add support for LTX 2.5 (#15499) --------- Co-authored-by: kijai <40791699+kijai@users.noreply.github.com> --- comfy/ldm/lightricks/av_model.py | 26 +- comfy/ldm/lightricks/duration_head.py | 81 +++ comfy/ldm/lightricks/embeddings_connector.py | 4 + comfy/ldm/lightricks/model.py | 131 ++++- comfy/ldm/lightricks/vae/audio_vae.py | 3 +- .../lightricks/vae/na_diffusion_decoder.py | 515 ++++++++++++++++++ comfy/model_base.py | 8 + comfy/model_detection.py | 1 + comfy/sd.py | 99 +++- comfy/text_encoders/gemma4.py | 17 +- comfy/text_encoders/lt.py | 113 ++-- comfy_extras/nodes_lt.py | 244 +++++++++ comfy_extras/nodes_lt_audio.py | 2 +- comfy_extras/nodes_model_patch.py | 5 + comfy_extras/nodes_textgen.py | 98 +++- 15 files changed, 1259 insertions(+), 88 deletions(-) create mode 100644 comfy/ldm/lightricks/duration_head.py create mode 100644 comfy/ldm/lightricks/vae/na_diffusion_decoder.py diff --git a/comfy/ldm/lightricks/av_model.py b/comfy/ldm/lightricks/av_model.py index 8e360f6a8..c60148e2a 100644 --- a/comfy/ldm/lightricks/av_model.py +++ b/comfy/ldm/lightricks/av_model.py @@ -96,6 +96,8 @@ class BasicAVTransformerBlock(nn.Module): attn_precision=None, apply_gated_attention=False, cross_attention_adaln=False, + ff_bias=True, + audio_ff_bias=True, dtype=None, device=None, operations=None, @@ -178,10 +180,10 @@ class BasicAVTransformerBlock(nn.Module): ) self.ff = FeedForward( - v_dim, dim_out=v_dim, glu=True, dtype=dtype, device=device, operations=operations + v_dim, dim_out=v_dim, glu=True, ff_bias=ff_bias, dtype=dtype, device=device, operations=operations ) self.audio_ff = FeedForward( - a_dim, dim_out=a_dim, glu=True, dtype=dtype, device=device, operations=operations + a_dim, dim_out=a_dim, glu=True, ff_bias=audio_ff_bias, dtype=dtype, device=device, operations=operations ) num_ada_params = ADALN_CROSS_ATTN_PARAMS_COUNT if cross_attention_adaln else ADALN_BASE_PARAMS_COUNT @@ -413,12 +415,16 @@ class LTXAVModel(LTXVModel): apply_gated_attention=False, caption_proj_before_connector=False, cross_attention_adaln=False, + ff_bias=True, + audio_ff_bias=True, + use_prompt_adaln_single=True, dtype=None, device=None, operations=None, **kwargs, ): # Store audio-specific parameters + self.audio_ff_bias = audio_ff_bias self.audio_in_channels = audio_in_channels self.audio_cross_attention_dim = audio_cross_attention_dim self.audio_attention_head_dim = audio_attention_head_dim @@ -451,6 +457,8 @@ class LTXAVModel(LTXVModel): timestep_scale_multiplier=timestep_scale_multiplier, caption_proj_before_connector=caption_proj_before_connector, cross_attention_adaln=cross_attention_adaln, + ff_bias=ff_bias, + use_prompt_adaln_single=use_prompt_adaln_single, dtype=dtype, device=device, operations=operations, @@ -475,7 +483,7 @@ class LTXAVModel(LTXVModel): operations=self.operations, ) - if self.cross_attention_adaln: + if self.cross_attention_adaln and self.use_prompt_adaln_single: self.audio_prompt_adaln_single = AdaLayerNormSingle( self.audio_inner_dim, embedding_coefficient=2, @@ -606,6 +614,8 @@ class LTXAVModel(LTXVModel): a_context_dim=self.audio_cross_attention_dim, apply_gated_attention=self.apply_gated_attention, cross_attention_adaln=self.cross_attention_adaln, + ff_bias=self.ff_bias, + audio_ff_bias=self.audio_ff_bias, dtype=dtype, device=device, operations=self.operations, @@ -924,9 +934,15 @@ class LTXAVModel(LTXVModel): blocks_replace = patches_replace.get("dit", {}) prefetch_queue = comfy.model_prefetch.make_prefetch_queue(list(self.transformer_blocks), vx.device, transformer_options) + # Blocks whose self-attention should be perturbed to a value-passthrough (STG). + stg_self_attn_blocks = transformer_options.get("stg_self_attn_blocks", ()) + # Process transformer blocks for i, block in enumerate(self.transformer_blocks): comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, vx.device, block) + block_transformer_options = transformer_options + if i in stg_self_attn_blocks: + block_transformer_options = {**transformer_options, "stg_skip_self_attn": True} if ("double_block", i) in blocks_replace: def block_wrap(args): @@ -969,7 +985,7 @@ class LTXAVModel(LTXVModel): "a_cross_scale_shift_timestep": av_ca_audio_scale_shift_timestep, "v_cross_gate_timestep": av_ca_a2v_gate_noise_timestep, "a_cross_gate_timestep": av_ca_v2a_gate_noise_timestep, - "transformer_options": transformer_options, + "transformer_options": block_transformer_options, "self_attention_mask": self_attention_mask, "v_prompt_timestep": v_prompt_timestep, "a_prompt_timestep": a_prompt_timestep, @@ -993,7 +1009,7 @@ class LTXAVModel(LTXVModel): a_cross_scale_shift_timestep=av_ca_audio_scale_shift_timestep, v_cross_gate_timestep=av_ca_a2v_gate_noise_timestep, a_cross_gate_timestep=av_ca_v2a_gate_noise_timestep, - transformer_options=transformer_options, + transformer_options=block_transformer_options, self_attention_mask=self_attention_mask, v_prompt_timestep=v_prompt_timestep, a_prompt_timestep=a_prompt_timestep, diff --git a/comfy/ldm/lightricks/duration_head.py b/comfy/ldm/lightricks/duration_head.py new file mode 100644 index 000000000..7d45d1fa7 --- /dev/null +++ b/comfy/ldm/lightricks/duration_head.py @@ -0,0 +1,81 @@ +"""LTX 2.4 DurationHead: predicts the natural shot duration (in seconds) from +the caption connector token outputs, without running the diffusion pipeline. +""" + +import torch +import torch.nn.functional as F +from torch import nn + + +class AttentionPooler(nn.Module): + """Cross-attend ``num_queries`` learnable tokens against ``tokens``.""" + + def __init__(self, hidden_dim=256, num_queries=1, num_heads=4): + super().__init__() + self.num_queries = num_queries + self.query_tokens = nn.Parameter(torch.empty(num_queries, hidden_dim)) + self.cross_attn = nn.MultiheadAttention(embed_dim=hidden_dim, num_heads=num_heads, batch_first=True) + + def forward(self, tokens): + queries = self.query_tokens.unsqueeze(0).expand(tokens.shape[0], -1, -1) + pooled, _ = self.cross_attn(queries, tokens, tokens, need_weights=False) + return pooled + + +class DurationHead(nn.Module): + """Predict duration in seconds from one or both connector outputs.""" + + def __init__( + self, + video_cross_attention_dim=4096, + audio_cross_attention_dim=2048, + pooler_hidden_dim=256, + num_queries=1, + num_pooler_heads=4, + mlp_hidden=256, + ): + super().__init__() + self.video_input_proj = nn.Linear(video_cross_attention_dim, pooler_hidden_dim) + self.video_modality_emb = nn.Parameter(torch.empty(pooler_hidden_dim)) + self.audio_input_proj = nn.Linear(audio_cross_attention_dim, pooler_hidden_dim) + self.audio_modality_emb = nn.Parameter(torch.empty(pooler_hidden_dim)) + self.attention_pooler = AttentionPooler( + hidden_dim=pooler_hidden_dim, num_queries=num_queries, num_heads=num_pooler_heads) + self.mlp_hidden = nn.Linear(pooler_hidden_dim * num_queries, mlp_hidden) + self.mlp_out = nn.Linear(mlp_hidden, 1) + + def forward(self, video_tokens=None, audio_tokens=None): + """``video_tokens``: (B, T_v, 4096), ``audio_tokens``: (B, T_a, 2048); + at least one required. Returns duration in seconds, shape (B,).""" + token_groups = [] + if video_tokens is not None: + token_groups.append(self.video_input_proj(video_tokens) + self.video_modality_emb) + if audio_tokens is not None: + token_groups.append(self.audio_input_proj(audio_tokens) + self.audio_modality_emb) + if not token_groups: + raise ValueError("DurationHead requires at least one of video_tokens / audio_tokens") + pooled = self.attention_pooler(torch.cat(token_groups, dim=1)) + pooled = pooled.reshape(pooled.shape[0], -1) + hidden = F.gelu(self.mlp_hidden(pooled), approximate="tanh") + return self.mlp_out(hidden).squeeze(-1).exp() + + +def normalize_state_dict(sd): + for prefix in ("model.diffusion_model.duration_head.", "duration_head."): + stripped = {k[len(prefix):]: v for k, v in sd.items() if k.startswith(prefix)} + if stripped: + return stripped + return sd + + +def seconds_to_num_frames(seconds, frame_rate, min_seconds, max_seconds, time_scale=8): + """Convert seconds to a frame count clamped to ``[min_seconds, max_seconds]`` + and snapped (floor) to the VAE's ``8k + 1`` causal temporal grid; snapping + that undershoots the minimum bumps up to the next grid point instead.""" + min_frames = max(1, round(min_seconds * frame_rate)) + max_frames = round(max_seconds * frame_rate) + raw_frames = max(min_frames, min(round(seconds * frame_rate), max_frames)) + frames = (raw_frames - 1) // time_scale * time_scale + 1 + if frames < min_frames: + frames = min(-(-(min_frames - 1) // time_scale) * time_scale + 1, max_frames) + return frames diff --git a/comfy/ldm/lightricks/embeddings_connector.py b/comfy/ldm/lightricks/embeddings_connector.py index 1a6ddcc8d..9c412827f 100644 --- a/comfy/ldm/lightricks/embeddings_connector.py +++ b/comfy/ldm/lightricks/embeddings_connector.py @@ -50,6 +50,7 @@ class BasicTransformerBlock1D(nn.Module): context_dim=None, attn_precision=None, apply_gated_attention=False, + ff_bias=True, dtype=None, device=None, operations=None, @@ -74,6 +75,7 @@ class BasicTransformerBlock1D(nn.Module): dim, dim_out=dim, glu=True, + ff_bias=ff_bias, dtype=dtype, device=device, operations=operations, @@ -123,6 +125,7 @@ class Embeddings1DConnector(nn.Module): causal_temporal_positioning=False, num_learnable_registers: Optional[int] = 128, apply_gated_attention=False, + connector_ff_bias=True, dtype=None, device=None, operations=None, @@ -148,6 +151,7 @@ class Embeddings1DConnector(nn.Module): attention_head_dim, context_dim=cross_attention_dim, apply_gated_attention=apply_gated_attention, + ff_bias=connector_ff_bias, dtype=dtype, device=device, operations=operations, diff --git a/comfy/ldm/lightricks/model.py b/comfy/ldm/lightricks/model.py index f80bffba7..dcbfa43ad 100644 --- a/comfy/ldm/lightricks/model.py +++ b/comfy/ldm/lightricks/model.py @@ -303,22 +303,22 @@ class NormSingleLinearTextProjection(nn.Module): class GELU_approx(nn.Module): - def __init__(self, dim_in, dim_out, dtype=None, device=None, operations=None): + def __init__(self, dim_in, dim_out, bias=True, dtype=None, device=None, operations=None): super().__init__() - self.proj = operations.Linear(dim_in, dim_out, dtype=dtype, device=device) + self.proj = operations.Linear(dim_in, dim_out, bias=bias, dtype=dtype, device=device) def forward(self, x): return torch.nn.functional.gelu(self.proj(x), approximate="tanh") class FeedForward(nn.Module): - def __init__(self, dim, dim_out, mult=4, glu=False, dropout=0.0, dtype=None, device=None, operations=None): + def __init__(self, dim, dim_out, mult=4, glu=False, dropout=0.0, ff_bias=True, dtype=None, device=None, operations=None): super().__init__() inner_dim = int(dim * mult) - project_in = GELU_approx(dim, inner_dim, dtype=dtype, device=device, operations=operations) + project_in = GELU_approx(dim, inner_dim, bias=ff_bias, dtype=dtype, device=device, operations=operations) self.net = nn.Sequential( - project_in, nn.Dropout(dropout), operations.Linear(inner_dim, dim_out, dtype=dtype, device=device) + project_in, nn.Dropout(dropout), operations.Linear(inner_dim, dim_out, bias=ff_bias, dtype=dtype, device=device) ) def forward(self, x): @@ -462,28 +462,34 @@ class CrossAttention(nn.Module): ) def forward(self, x, context=None, mask=None, pe=None, k_pe=None, transformer_options={}): + self_attn = context is None q = self.to_q(x) context = x if context is None else context k = self.to_k(context) v = self.to_v(context) - q = self.q_norm(q) - 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 k_pe is None and q.shape == k.shape: - 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: - out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options) - elif isinstance(mask, GuideAttentionMask): - out = _attention_with_guide_mask(q, k, v, self.heads, mask, attn_precision=self.attn_precision, transformer_options=transformer_options) + # Spatio-Temporal Guidance (STG) perturbation: for the flagged self-attention + # layers, the attention degrades to a passthrough of the value projection (out = V). + if self_attn and transformer_options.get("stg_skip_self_attn", False): + out = v else: - out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, mask=mask, attn_precision=self.attn_precision, transformer_options=transformer_options) + q = self.q_norm(q) + 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 k_pe is None and q.shape == k.shape: + 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: + out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options) + elif isinstance(mask, GuideAttentionMask): + out = _attention_with_guide_mask(q, k, v, self.heads, mask, attn_precision=self.attn_precision, transformer_options=transformer_options) + else: + out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, mask=mask, attn_precision=self.attn_precision, transformer_options=transformer_options) # Apply per-head gating if enabled if self.to_gate_logits is not None: @@ -502,7 +508,7 @@ ADALN_CROSS_ATTN_PARAMS_COUNT = 9 class BasicTransformerBlock(nn.Module): def __init__( - self, dim, n_heads, d_head, context_dim=None, attn_precision=None, cross_attention_adaln=False, dtype=None, device=None, operations=None + self, dim, n_heads, d_head, context_dim=None, attn_precision=None, cross_attention_adaln=False, ff_bias=True, dtype=None, device=None, operations=None ): super().__init__() @@ -518,7 +524,7 @@ class BasicTransformerBlock(nn.Module): device=device, operations=operations, ) - self.ff = FeedForward(dim, dim_out=dim, glu=True, dtype=dtype, device=device, operations=operations) + self.ff = FeedForward(dim, dim_out=dim, glu=True, ff_bias=ff_bias, dtype=dtype, device=device, operations=operations) self.attn2 = CrossAttention( query_dim=dim, @@ -717,6 +723,9 @@ class LTXBaseModel(torch.nn.Module, ABC): caption_proj_before_connector=False, cross_attention_adaln=False, caption_projection_first_linear=True, + ff_bias=True, + use_prompt_adaln_single=True, + use_keyframes_abs_pos_embedding=False, dtype=None, device=None, operations=None, @@ -746,6 +755,9 @@ class LTXBaseModel(torch.nn.Module, ABC): self.caption_proj_before_connector = caption_proj_before_connector self.cross_attention_adaln = cross_attention_adaln self.caption_projection_first_linear = caption_projection_first_linear + self.ff_bias = ff_bias + self.use_prompt_adaln_single = use_prompt_adaln_single + self.use_keyframes_abs_pos_embedding = use_keyframes_abs_pos_embedding # Common dimensions self.inner_dim = num_attention_heads * attention_head_dim @@ -773,12 +785,17 @@ class LTXBaseModel(torch.nn.Module, ABC): self.in_channels, self.inner_dim, bias=True, dtype=dtype, device=device ) + if self.use_keyframes_abs_pos_embedding: + self.keyframes_abs_pos_embedding = nn.Parameter(torch.zeros(1, self.inner_dim, dtype=dtype, device=device)) + else: + self.keyframes_abs_pos_embedding = None + embedding_coefficient = ADALN_CROSS_ATTN_PARAMS_COUNT if self.cross_attention_adaln else ADALN_BASE_PARAMS_COUNT self.adaln_single = AdaLayerNormSingle( self.inner_dim, embedding_coefficient=embedding_coefficient, use_additional_conditions=False, dtype=dtype, device=device, operations=self.operations ) - if self.cross_attention_adaln: + if self.cross_attention_adaln and self.use_prompt_adaln_single: self.prompt_adaln_single = AdaLayerNormSingle( self.inner_dim, embedding_coefficient=2, use_additional_conditions=False, dtype=dtype, device=device, operations=self.operations ) @@ -1070,6 +1087,7 @@ class LTXVModel(LTXBaseModel): self.attention_head_dim, context_dim=self.cross_attention_dim, cross_attention_adaln=self.cross_attention_adaln, + ff_bias=self.ff_bias, dtype=dtype, device=device, operations=self.operations, @@ -1099,6 +1117,15 @@ class LTXVModel(LTXBaseModel): grid_mask = None if keyframe_idxs is not None and keyframe_idxs.shape[2] > 0: + tokens_per_frame = self.tokens_per_latent_frame(additional_args["orig_shape"]) + if keyframe_idxs.shape[2] % tokens_per_frame != 0: + raise ValueError( + f"keyframe_idxs holds {keyframe_idxs.shape[2]} tokens, which is not a whole number of " + f"{tokens_per_frame}-token latent frames. The appended frames were recorded against a " + "different spatial resolution than the latent being sampled, so their positions would land " + "on the wrong tokens. Crop the guides and separate the generated keyframes before " + "upscaling the latent." + ) additional_args.update({ "orig_patchified_shape": list(x.shape)}) denoise_mask = self.patchifier.patchify(denoise_mask)[0] grid_mask = ~torch.any(denoise_mask < 0, dim=-1)[0] @@ -1141,8 +1168,64 @@ class LTXVModel(LTXBaseModel): additional_args["num_guide_tokens"] = keyframe_idxs.shape[2] x = self.patchify_proj(x) + x = self.apply_keyframes_abs_pos_embedding( + x, + pixel_coords, + orig_shape=additional_args["orig_shape"], + grid_mask=grid_mask, + num_guide_tokens=additional_args.get("num_guide_tokens", 0), + generated_keyframes=kwargs.get("generated_keyframes", None), + ) return x, pixel_coords, additional_args + def tokens_per_latent_frame(self, orig_shape): + """Token count of a single latent frame at the given latent shape.""" + patch_size = self.patchifier.patch_size + return (orig_shape[3] // patch_size[1]) * (orig_shape[4] // patch_size[2]) + + def keyframes_abs_pos_mask(self, pixel_coords, orig_shape, grid_mask, num_guide_tokens, generated_keyframes): + """Per-token mask selecting the latents that encode a single standalone pixel frame. + + Returns a (batch, tokens) boolean mask over the already grid-filtered token sequence. + """ + temporal_start = pixel_coords[:, 0] + if temporal_start.ndim == 3: # (batch, tokens, [start, end]) + temporal_start = temporal_start[..., 0] + mask = temporal_start == 0 + if num_guide_tokens > 0: + mask[:, -num_guide_tokens:] = False + + if generated_keyframes is not None: + # The temporal patch size is always 1, so one latent frame is one row of tokens. + tokens_per_frame = self.tokens_per_latent_frame(orig_shape) + if generated_keyframes["tokens_per_frame"] != tokens_per_frame: + raise ValueError( + f"The generated keyframes were recorded at {generated_keyframes['tokens_per_frame']} tokens " + f"per latent frame but this latent has {tokens_per_frame}. Separate the generated keyframes " + "before upscaling the latent." + ) + first_token = generated_keyframes["first_latent_frame"] * tokens_per_frame + num_slot_tokens = generated_keyframes["num_keyframes"] * tokens_per_frame + slots = torch.zeros(orig_shape[2] * tokens_per_frame, dtype=torch.bool, device=mask.device) + slots[first_token:first_token + num_slot_tokens] = True + if grid_mask is not None: + slots = slots[grid_mask] + mask = mask | slots + + return mask + + def apply_keyframes_abs_pos_embedding(self, x, pixel_coords, orig_shape, grid_mask, num_guide_tokens, generated_keyframes): + """Add the learned keyframe marker to the single-pixel-frame tokens. + + A no-op for every checkpoint built without the parameter. + """ + if self.keyframes_abs_pos_embedding is None: + return x + + mask = self.keyframes_abs_pos_mask(pixel_coords, orig_shape, grid_mask, num_guide_tokens, generated_keyframes) + embedding = self.keyframes_abs_pos_embedding.to(device=x.device, dtype=x.dtype) + return x + mask.unsqueeze(-1).to(x.dtype) * embedding + def _build_guide_self_attention_mask(self, x, transformer_options, merged_args): """Build self-attention mask for per-guide attention attenuation. diff --git a/comfy/ldm/lightricks/vae/audio_vae.py b/comfy/ldm/lightricks/vae/audio_vae.py index b4a8c7524..f5b1756d3 100644 --- a/comfy/ldm/lightricks/vae/audio_vae.py +++ b/comfy/ldm/lightricks/vae/audio_vae.py @@ -1,6 +1,5 @@ import json from dataclasses import dataclass -import math import torch import torchaudio @@ -186,7 +185,7 @@ class AudioVAE(torch.nn.Module): ) def num_of_latents_from_frames(self, frames_number: int, frame_rate: float) -> int: - return math.ceil((float(frames_number) / frame_rate) * self.latents_per_second) + return round((float(frames_number) / frame_rate) * self.latents_per_second) def run_vocoder(self, mel_spec: torch.Tensor) -> torch.Tensor: audio_channels = self.autoencoder.decoder.out_ch diff --git a/comfy/ldm/lightricks/vae/na_diffusion_decoder.py b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py new file mode 100644 index 000000000..8a172e101 --- /dev/null +++ b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py @@ -0,0 +1,515 @@ +"""LTX 2.4 diffusion video VAE decoder (NADiffusionDecoder). + +Port of the reference ``DiffusionVideoDecoder`` without the NATTEN dependency: +``natten.na3d`` is replaced by ``comfy_kitchen.na3d``, which reproduces +NATTEN's semantics (window of exactly ``kernel_size`` per query, shifted +inward at grid boundaries, dilation 1) and dispatches cuda/triton/eager per +device and dtype (the eager backend covers CPU and fp32). + +Stages 1-4 deterministically upsample the latent into a context volume via +NA transformer blocks + linear pixel-shuffle upsamples. Stage 5 runs +``DiffusionNABlock``s that denoise patchified noised pixels ``x_t`` guided by +that context through AdaLN-Zero scale/shift. The 2.4 checkpoint is single-step +``x0``: one forward pass yields the pixels directly, no Euler loop. + +State dict keys match the shipped checkpoints directly (fused ``attn.qkv``, +``t_embedder.mlp.{0,2}``, ``shared_adaln.proj``); no rename pass is needed. +""" + +import math + +import torch +import torch.nn.functional as F +from einops import rearrange +from torch import nn + +from comfy.ldm.lightricks.model import get_timestep_embedding +from .causal_video_autoencoder import Encoder, processor + +import comfy_kitchen + +# Token chunk for the SwiGLU MLP (bounds the [chunk, hidden] workspace). +MLP_TOKEN_CHUNK = 65536 + + +def rms_norm(x, weight, eps=1e-6): + if hasattr(F, "rms_norm"): + return F.rms_norm(x, (x.shape[-1],), weight=weight.to(x.dtype), eps=eps) + x_f = x.float() + x_f = x_f * torch.rsqrt(x_f.pow(2).mean(-1, keepdim=True) + eps) + return (x_f * weight.float()).to(x.dtype) + + +class RMSNorm(nn.Module): + def __init__(self, dim, eps=1e-6): + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim)) + + def forward(self, x): + return rms_norm(x, self.weight, self.eps) + + +def patchify(x, patch_size_hw, patch_size_t=1): + if patch_size_hw == 1 and patch_size_t == 1: + return x + return rearrange(x, "b c (f p) (h q) (w r) -> b (c p r q) f h w", p=patch_size_t, q=patch_size_hw, r=patch_size_hw) + + +def unpatchify(x, patch_size_hw, patch_size_t=1): + if patch_size_hw == 1 and patch_size_t == 1: + return x + return rearrange(x, "b (c p r q) f h w -> b c (f p) (h q) (w r)", p=patch_size_t, q=patch_size_hw, r=patch_size_hw) + + +# --- Absolute per-axis RoPE (matches ltx-core rope.py numerics) --- + +def default_rope_dim_split(head_dim): + d_t = (head_dim // 4) // 2 * 2 + d_hw = (head_dim - d_t) // 2 + if d_hw % 2 != 0: + d_t -= 2 + d_hw = (head_dim - d_t) // 2 + return (d_t, d_hw, d_hw) + + +def rope_inv_freqs(dim, base=10000.0, device=None): + exponents = torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim + return (1.0 / torch.pow(torch.tensor(float(base), dtype=torch.float64, device=device), exponents)).to(torch.float32) + + +def _rope_tables(lengths, inv_freqs, device): + """Precompute per-axis fp32 cos/sin tables for global 0-based positions.""" + tables = [] + for length, inv in zip(lengths, inv_freqs): + pos = torch.arange(length, dtype=torch.float32, device=device) + ang = pos[:, None] * inv[None, :] + tables.append((ang.cos(), ang.sin())) + return tables + + +def _rope_matrices_slice(tables, t0, t1, h, w): + """Per-token rotation matrices ``(1, ts*h*w, 1, hd/2, 2, 2)`` fp32 for + ``comfy_kitchen.rms_rope_`` (interleaved-pair convention), covering global + frames ``[t0, t1)`` of the axis-factorized tables.""" + parts = [] + for (c, s), sl in zip(tables, (slice(t0, t1), slice(None), slice(None))): + c, s = c[sl], s[sl] + parts.append(torch.stack([c, -s, s, c], dim=-1).reshape(c.shape[0], 1, 1, c.shape[1], 2, 2)) + ts = t1 - t0 + freqs = torch.cat([ + parts[0].expand(ts, h, w, -1, 2, 2), + parts[1].transpose(0, 1).expand(ts, h, w, -1, 2, 2), + parts[2].movedim(0, 2).expand(ts, h, w, -1, 2, 2), + ], dim=3) + return freqs.reshape(1, ts * h * w, 1, -1, 2, 2) + + +class NeighborhoodAttention3D(nn.Module): + """QKV (fused, matching checkpoint keys) + q/k RMSNorm + abs RoPE + NA.""" + + def __init__(self, dim, kernel_size, head_dim=64, rope_base=10000.0): + super().__init__() + self.dim = dim + self.num_heads = dim // head_dim + self.head_dim = head_dim + self.kernel_size = tuple(kernel_size) + self.scale = head_dim ** -0.5 + self.rope_split = default_rope_dim_split(head_dim) + self.rope_base = rope_base + + self.qkv = nn.Linear(dim, dim * 3, bias=True) + self.proj = nn.Linear(dim, dim, bias=True) + self.q_norm = RMSNorm(head_dim, eps=1e-6) + self.k_norm = RMSNorm(head_dim, eps=1e-6) + + def forward(self, x, pre=None, add_to=None): + """``pre`` (per-token norm/modulate) is applied slice-wise so the full + pre-attention tensor is never materialized; ``add_to`` streams the + output projection into it in place (residual add) and returns it. + Both bound peak memory without changing results.""" + batch, t, h, w, _ = x.shape + inv_freqs = tuple(rope_inv_freqs(d, self.rope_base, device=x.device) for d in self.rope_split) + tables = _rope_tables((t, h, w), inv_freqs, x.device) + shape = (batch, t, h, w, self.num_heads, self.head_dim) + q = torch.empty(shape, dtype=x.dtype, device=x.device) + k = torch.empty(shape, dtype=x.dtype, device=x.device) + v = torch.empty(shape, dtype=x.dtype, device=x.device) + q_weight = (self.q_norm.weight.detach() * self.scale).to(x.dtype) # scale commutes with the rotation + k_weight = self.k_norm.weight.detach().to(x.dtype) + chunk = max(1, (2 ** 25) // max(h * w * self.dim, 1)) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + sl = x[:, t0:t1] if pre is None else pre(x[:, t0:t1]) + qc, kc, vc = self.qkv(sl).chunk(3, dim=-1) + cshape = (batch, t1 - t0, h, w, self.num_heads, self.head_dim) + q[:, t0:t1] = qc.reshape(cshape) + k[:, t0:t1] = kc.reshape(cshape) + v[:, t0:t1] = vc.reshape(cshape) + freqs = _rope_matrices_slice(tables, t0, t1, h, w) + nt = (t1 - t0) * h * w + for b in range(batch): + comfy_kitchen.rms_rope_( + q[b, t0:t1].view(1, nt, self.num_heads, self.head_dim), + k[b, t0:t1].view(1, nt, self.num_heads, self.head_dim), + freqs, q_weight, k_weight) + out = comfy_kitchen.na3d(q, k, v, list(self.kernel_size), None, 1.0) + del q, k, v + out = out.reshape(batch, t, h, w, self.dim) + res = add_to if add_to is not None else torch.empty_like(out) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + if add_to is not None: + res[:, t0:t1] += self.proj(out[:, t0:t1]) + else: + res[:, t0:t1] = self.proj(out[:, t0:t1]) + return res + + +class SwiGLU(nn.Module): + """``w_down(silu(w_gate(x)) * w_up(x))``, chunked over tokens to bound the + ``[chunk, hidden]`` workspace.""" + + def __init__(self, dim, hidden_dim): + super().__init__() + self.w_up = nn.Linear(dim, hidden_dim, bias=False) + self.w_gate = nn.Linear(dim, hidden_dim, bias=False) + self.w_down = nn.Linear(hidden_dim, dim, bias=False) + + def forward(self, x, pre=None, add_to=None): + """``pre``/``add_to`` as in ``NeighborhoodAttention3D.forward``.""" + _, t, h, w, _ = x.shape + chunk = max(1, MLP_TOKEN_CHUNK // max(h * w, 1)) + out = add_to if add_to is not None else torch.empty_like(x) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + sl = x[:, t0:t1] if pre is None else pre(x[:, t0:t1]) + y = self.w_down(F.silu(self.w_gate(sl)) * self.w_up(sl)) + if add_to is not None: + out[:, t0:t1] += y + else: + out[:, t0:t1] = y + return out + + +class NABlock(nn.Module): + """Pre-norm transformer block: NA -> SwiGLU MLP with residual adds.""" + + def __init__(self, dim, kernel_size, head_dim=64, mlp_ratio=4.0): + super().__init__() + self.norm1 = RMSNorm(dim, eps=1e-6) + self.attn = NeighborhoodAttention3D(dim, kernel_size, head_dim=head_dim) + self.norm2 = RMSNorm(dim, eps=1e-6) + hidden = (int(dim * mlp_ratio) + 15) // 16 * 16 + self.mlp = SwiGLU(dim, hidden) + + def forward(self, x): + x = self.attn(x, pre=self.norm1, add_to=x) + return self.mlp(x, pre=self.norm2, add_to=x) + + +def modulate(x, scale, shift): + return x * (1.0 + scale) + shift + + +class AdaLNZero(nn.Module): + """``t_emb`` -> 7 (scale/shift/gate) chunks; gate slots unused (folded at export).""" + + NUM_CHUNKS = 7 + + def __init__(self, dim, t_emb_dim): + super().__init__() + self.proj = nn.Linear(t_emb_dim, self.NUM_CHUNKS * dim, bias=True) + + def forward(self, t_emb): + h = self.proj(F.silu(t_emb)) + return tuple(c[:, None, None, None, :] for c in h.chunk(self.NUM_CHUNKS, dim=-1)) + + +class DiffusionNABlock(nn.Module): + """NA + SwiGLU with shared AdaLN-Zero scale/shift (ungated residuals).""" + + def __init__(self, dim, kernel_size, context_channels, head_dim=64, mlp_ratio=4.0): + super().__init__() + self.context_proj = nn.Linear(context_channels, dim, bias=True) + self.scale_shift_table = nn.Parameter(torch.zeros(AdaLNZero.NUM_CHUNKS, dim)) + self.norm1 = RMSNorm(dim, eps=1e-6) + self.attn = NeighborhoodAttention3D(dim, kernel_size, head_dim=head_dim) + self.norm2 = RMSNorm(dim, eps=1e-6) + hidden = (int(dim * mlp_ratio) + 15) // 16 * 16 + self.mlp = SwiGLU(dim, hidden) + + def forward(self, x, latent_context, modulation): + scale_msa, shift_msa, _, scale_mlp, shift_mlp, _, _ = [ + modulation[i] + self.scale_shift_table[i].view(1, 1, 1, 1, -1) for i in range(AdaLNZero.NUM_CHUNKS) + ] + chunk = max(1, MLP_TOKEN_CHUNK // max(x.shape[2] * x.shape[3], 1)) + for t0 in range(0, x.shape[1], chunk): + x[:, t0:t0 + chunk] += self.context_proj(latent_context[:, t0:t0 + chunk]) + x = self.attn(x, pre=lambda s: modulate(self.norm1(s), scale_msa, shift_msa), add_to=x) + return self.mlp(x, pre=lambda s: modulate(self.norm2(s), scale_mlp, shift_mlp), add_to=x) + + +class LinearPixelShuffleUpsample(nn.Module): + """Linear channel-expand, then channels-last pixel shuffle.""" + + def __init__(self, in_channels, stride, out_channels_reduction_factor=1): + super().__init__() + self.stride = tuple(stride) + proj_out_channels = math.prod(stride) * in_channels // out_channels_reduction_factor + self.out_channels = proj_out_channels // math.prod(stride) + self.proj = nn.Linear(in_channels, proj_out_channels, bias=True) + + def forward(self, x, drop_leading_frame=True): + batch, t, h, w, _ = x.shape + p1, p2, p3 = self.stride + out = torch.empty((batch, t * p1, h * p2, w * p3, self.out_channels), dtype=x.dtype, device=x.device) + chunk = max(1, MLP_TOKEN_CHUNK // max(h * w, 1)) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + out[:, t0 * p1:t1 * p1] = rearrange( + self.proj(x[:, t0:t1]), "b t h w (c p1 p2 p3) -> b (t p1) (h p2) (w p3) c", + p1=p1, p2=p2, p3=p3, + ) + if p1 == 2 and drop_leading_frame: + # The causal temporal pixel-shuffle duplicates the leading frame. + out = out[:, 1:] + return out + + +class TimestepEmbedder(nn.Module): + """Sinusoidal(256) -> MLP. ``mlp.{0,2}`` naming matches the checkpoint.""" + + def __init__(self, t_emb_dim=384, freq_dim=256): + super().__init__() + self.freq_dim = freq_dim + self.mlp = nn.Sequential( + nn.Linear(freq_dim, t_emb_dim, bias=True), + nn.SiLU(), + nn.Linear(t_emb_dim, t_emb_dim, bias=True), + ) + + def forward(self, timestep, dtype): + emb = get_timestep_embedding(timestep.flatten(), self.freq_dim, flip_sin_to_cos=True, + downscale_freq_shift=0, scale=1) + return self.mlp(emb.to(dtype)) + + +class NADiffusionDecoder(nn.Module): + """Stages 1-4 (deterministic NA upsample) + stage-5 diffusion blocks. + + Input latent must already be un-normalized (the wrapper applies + ``per_channel_statistics.un_normalize``, same as the conv VAE path). + """ + + def __init__( + self, + in_channels=128, + out_channels=3, + patch_size=4, + head_dim=64, + stage_channels=(2048, 1024, 512, 512, 256), + stage_depths=(4, 6, 4, 2, 8), + stage_kernels=((3, 7, 7), (3, 7, 7), (3, 5, 5), (3, 5, 5), (11, 11, 11)), + upsamples=(((1, 2, 2), 2), ((2, 1, 1), 2), ((2, 2, 2), 1), ((2, 2, 2), 2)), + stage5_kernel=(11, 11, 11), + t_emb_dim=384, + default_num_inference_steps=1, + timestep_scale_multiplier=1000.0, + model_output_type="x0", + ): + super().__init__() + self.patch_size = patch_size + self.out_channels = out_channels + self.timestep_scale_multiplier = timestep_scale_multiplier + self.model_output_type = model_output_type + self.register_buffer( + "default_inference_timesteps", + torch.linspace(1.0, 1.0 / default_num_inference_steps, default_num_inference_steps), + persistent=False, + ) + self.temporal_upscale = math.prod(s[0] for s, _ in upsamples) + self.spatial_upscale = math.prod(s[1] for s, _ in upsamples) * patch_size + # NATTEN-style last-frame border mitigation: replicate the last latent + # frame through stages 1-4, crop the appendix off the context after. + self.trailing_pad_latent_frames = (stage_kernels[0][0] // 2) * 2 + + self.conv_in = nn.Linear(in_channels, stage_channels[0], bias=True) + + self.det_stages = nn.ModuleList() + self.upsamples = nn.ModuleList() + for stage_i in range(len(stage_channels) - 1): + c = stage_channels[stage_i] + self.det_stages.append(nn.ModuleList( + [NABlock(c, stage_kernels[stage_i], head_dim=head_dim) for _ in range(stage_depths[stage_i])] + )) + stride, reduction = upsamples[stage_i] + self.upsamples.append(LinearPixelShuffleUpsample(c, stride, out_channels_reduction_factor=reduction)) + + self.t_embedder = TimestepEmbedder(t_emb_dim=t_emb_dim) + + c5 = stage_channels[-1] + self.context_channels = c5 + noised_pixel_channels = out_channels * (patch_size ** 2) + self.conv_in_x_t = nn.Linear(noised_pixel_channels, c5, bias=True) + self.shared_adaln = AdaLNZero(c5, t_emb_dim) + self.diff_blocks = nn.ModuleList([ + DiffusionNABlock(c5, stage5_kernel, context_channels=c5, head_dim=head_dim) + for _ in range(stage_depths[-1]) + ]) + self.norm_out = RMSNorm(c5, eps=1e-6) + self.conv_out = nn.Linear(c5, noised_pixel_channels, bias=True) + + def forward_pre_diffusion(self, z, drop_leading_frame=True, pad_trailing=True): + """Stages 1-4: latent -> stage-5 context, channels-last. + + ``drop_leading_frame`` must be True only when ``z`` contains the + latent's true temporal origin (t=0); tiled callers decoding a later + temporal chunk pass False (the duplicate leading frame belongs solely + to the origin chunk). ``pad_trailing`` only for chunks containing the + latent's last frame.""" + n = self.trailing_pad_latent_frames if pad_trailing else 0 + if n > 0: + z = torch.cat([z, z[:, :, -1:].expand(-1, -1, n, -1, -1)], dim=2) + x = z.permute(0, 2, 3, 4, 1) + x = self.conv_in(x) + for stage_i, blocks in enumerate(self.det_stages): + for block in blocks: + x = block(x) + x = self.upsamples[stage_i](x, drop_leading_frame=drop_leading_frame) + if n > 0: + x = x[:, :-(n * self.temporal_upscale)] + return x + + def forward_diff_step(self, context, x_t, t): + x = patchify(x_t, patch_size_hw=self.patch_size, patch_size_t=1) + x = self.conv_in_x_t(x.permute(0, 2, 3, 4, 1)) + t_emb = self.t_embedder(self.timestep_scale_multiplier * t, dtype=x.dtype) + modulation = self.shared_adaln(t_emb) + for block in self.diff_blocks: + x = block(x, context, modulation) + x = self.norm_out(x) + x = self.conv_out(x) + x = x.permute(0, 4, 1, 2, 3) + return unpatchify(x, patch_size_hw=self.patch_size, patch_size_t=1) + + def forward(self, z, generator=None, drop_leading_frame=True, pad_trailing=True): + context = self.forward_pre_diffusion(z, drop_leading_frame=drop_leading_frame, pad_trailing=pad_trailing) + batch, t5, h5, w5, _ = context.shape + pixel_shape = (batch, self.out_channels, t5, h5 * self.patch_size, w5 * self.patch_size) + x_t = torch.randn(pixel_shape, dtype=z.dtype, device=z.device, generator=generator) + + timesteps = self.default_inference_timesteps.to(z.device) + num_steps = timesteps.shape[0] + for i in range(num_steps): + t_now = timesteps[i].expand(batch) + model_out = self.forward_diff_step(context, x_t, t_now) + if self.model_output_type == "x0": + x0 = model_out + if i == num_steps - 1: + return x0 + velocity = (x_t.float() - x0.float()) / timesteps[i] + else: # "v" + velocity = model_out.float() + if i == num_steps - 1: + return (x_t.float() - timesteps[i] * velocity).to(z.dtype) + t_next = timesteps[i + 1] if i + 1 < num_steps else torch.zeros_like(timesteps[i]) + x_t = (x_t.float() - (timesteps[i] - t_next) * velocity).to(z.dtype) + return x_t + + +LTX_24_VAE_CONFIG = { + "_class_name": "CausalDiffusionVAE", + "dims": 3, + "model_output_type": "x0", + "encoder": { + "dims": 3, + "in_channels": 3, + "out_channels": 128, + "blocks": [ + ["res_x", {"num_layers": 4}], + ["compress_space_res", {"multiplier": 2}], + ["res_x", {"num_layers": 6}], + ["compress_time_res", {"multiplier": 2}], + ["res_x", {"num_layers": 4}], + ["compress_all_res", {"multiplier": 2}], + ["res_x", {"num_layers": 2}], + ["compress_all_res", {"multiplier": 1}], + ["res_x", {"num_layers": 2}], + ], + "patch_size": 4, + "latent_log_var": "constant", + "norm_layer": "pixel_norm", + "base_channels": 128, + "spatial_padding_mode": "zeros", + }, + "decoder": { + "in_channels": 128, + "out_channels": 3, + "patch_size": 4, + "head_dim": 64, + "stage_channels": [2048, 1024, 512, 512, 256], + "stage_depths": [4, 6, 4, 2, 8], + "stage_kernels": [[3, 7, 7], [3, 7, 7], [3, 5, 5], [3, 5, 5], [11, 11, 11]], + "upsamples": [[[1, 2, 2], 2], [[2, 1, 1], 2], [[2, 2, 2], 1], [[2, 2, 2], 2]], + "stage5_kernel": [11, 11, 11], + "timestep_scale_multiplier": 1000.0, + "default_num_inference_steps": 1, + }, +} + + +class CausalDiffusionVAE(nn.Module): + """LTX 2.4 video VAE: conv encoder (shared with the 2.0 arch) + NA + diffusion decoder. Interface mirrors ``causal_video_autoencoder.VideoVAE``. + """ + + def __init__(self, config=None): + super().__init__() + if config is None: + config = LTX_24_VAE_CONFIG + self.config = config + enc = config.get("encoder", LTX_24_VAE_CONFIG["encoder"]) + dec = config.get("decoder", LTX_24_VAE_CONFIG["decoder"]) + dec_defaults = LTX_24_VAE_CONFIG["decoder"] + + self.encoder = Encoder( + dims=enc.get("dims", 3), + in_channels=enc.get("in_channels", 3), + out_channels=enc.get("out_channels", 128), + blocks=enc.get("blocks", LTX_24_VAE_CONFIG["encoder"]["blocks"]), + patch_size=enc.get("patch_size", 4), + latent_log_var=enc.get("latent_log_var", "constant"), + norm_layer=enc.get("norm_layer", "pixel_norm"), + spatial_padding_mode=enc.get("spatial_padding_mode", "zeros"), + base_channels=enc.get("base_channels", 128), + ) + + self.decoder = NADiffusionDecoder( + in_channels=dec.get("in_channels", 128), + out_channels=dec.get("out_channels", 3), + patch_size=dec.get("patch_size", 4), + head_dim=dec.get("head_dim", 64), + stage_channels=tuple(dec.get("stage_channels", dec_defaults["stage_channels"])), + stage_depths=tuple(dec.get("stage_depths", dec_defaults["stage_depths"])), + stage_kernels=tuple(tuple(k) for k in dec.get("stage_kernels", dec_defaults["stage_kernels"])), + upsamples=tuple((tuple(s), r) for s, r in dec.get("upsamples", dec_defaults["upsamples"])), + stage5_kernel=tuple(dec.get("stage5_kernel", dec_defaults["stage5_kernel"])), + t_emb_dim=dec.get("t_emb_dim", 384), + default_num_inference_steps=dec.get("default_num_inference_steps", 1), + timestep_scale_multiplier=dec.get("timestep_scale_multiplier", 1000.0), + model_output_type=config.get("model_output_type", "x0"), + ) + + self.per_channel_statistics = processor() + + def encode(self, x, device=None): + x = x[:, :, :max(1, 1 + ((x.shape[2] - 1) // 8) * 8), :, :] + means, logvar = torch.chunk(self.encoder(x, device=device), 2, dim=1) + return self.per_channel_statistics.normalize(means) + + def decode(self, x): + # Fixed-seed noise so decodes are reproducible TODO: expose? + generator = torch.Generator(device=x.device) + generator.manual_seed(0) + return self.decoder(self.per_channel_statistics.un_normalize(x), generator=generator) diff --git a/comfy/model_base.py b/comfy/model_base.py index 469d301ea..7d855f5a1 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -1153,6 +1153,10 @@ class LTXV(BaseModel): if guide_attention_entries is not None: out['guide_attention_entries'] = comfy.conds.CONDConstant(guide_attention_entries) + generated_keyframes = kwargs.get("generated_keyframes", None) + if generated_keyframes is not None: + out['generated_keyframes'] = comfy.conds.CONDConstant(generated_keyframes) + return out def process_timestep(self, timestep, x, denoise_mask=None, **kwargs): @@ -1213,6 +1217,10 @@ class LTXAV(BaseModel): if ref_audio is not None: out['ref_audio'] = comfy.conds.CONDConstant(ref_audio) + generated_keyframes = kwargs.get("generated_keyframes", None) + if generated_keyframes is not None: + out['generated_keyframes'] = comfy.conds.CONDConstant(generated_keyframes) + return out def process_timestep(self, timestep, x, denoise_mask=None, audio_denoise_mask=None, **kwargs): diff --git a/comfy/model_detection.py b/comfy/model_detection.py index 103680fd1..bc7b2b9f8 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -397,6 +397,7 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["cross_attention_dim"] = shape[1] if metadata is not None and "config" in metadata: dit_config.update(json.loads(metadata["config"]).get("transformer", {})) + dit_config["use_keyframes_abs_pos_embedding"] = '{}keyframes_abs_pos_embedding'.format(key_prefix) in state_dict_keys return dit_config if '{}genre_embedder.weight'.format(key_prefix) in state_dict_keys: #ACE-Step model diff --git a/comfy/sd.py b/comfy/sd.py index 5fed4ca9a..8bae76768 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -11,6 +11,7 @@ from .ldm.cascade.stage_c_coder import StageC_coder from .ldm.audio.autoencoder import AudioOobleckVAE import comfy.ldm.genmo.vae.model import comfy.ldm.lightricks.vae.causal_video_autoencoder +import comfy.ldm.lightricks.vae.na_diffusion_decoder import comfy.ldm.lightricks.vae.audio_vae import comfy.ldm.cosmos.vae import comfy.ldm.wan.vae @@ -583,6 +584,22 @@ class VAE: 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_x_t.weight" in sd: # lightricks LTX 2.4 diffusion VAE decoder + vae_config = None + if metadata is not None and "config" in metadata: + vae_config = json.loads(metadata["config"]).get("vae", None) + self.first_stage_model = comfy.ldm.lightricks.vae.na_diffusion_decoder.CausalDiffusionVAE(config=vae_config) + self.latent_channels = sd["decoder.conv_in.weight"].shape[1] + self.latent_dim = 3 + self.disable_offload = True + self.crop_input = False # generic crop would narrow the frame axis by the 32x spatial ratio + self.memory_used_decode = lambda shape, dtype: (1700 * shape[2] * shape[3] * shape[4] * (8 * 8 * 8)) * model_management.dtype_size(dtype) + self.memory_used_encode = lambda shape, dtype: (80 * max(shape[2], 7) * shape[3] * shape[4]) * model_management.dtype_size(dtype) + self.upscale_ratio = (lambda a: max(0, a * 8 - 7), 32, 32) + self.upscale_index_formula = (8, 32, 32) + self.downscale_ratio = (lambda a: max(0, math.floor((a + 7) / 8)), 32, 32) + self.downscale_index_formula = (8, 32, 32) + self.working_dtypes = [torch.bfloat16, torch.float32] elif "decoder.conv_in.weight" in sd: if sd['decoder.conv_in.weight'].shape[1] == 64: ddconfig = {"block_out_channels": [128, 256, 512, 512, 1024, 1024], "in_channels": 3, "out_channels": 3, "num_res_blocks": 2, "ffactor_spatial": 32, "downsample_match_channel": True, "upsample_match_channel": True} @@ -1222,16 +1239,46 @@ class VAE: tile = 256 // self.spacial_compression_decode() overlap = tile // 4 if self.handles_tiling: + memory_used = self.memory_used_decode(self._tile_bounded_shape(samples_in.shape, tile, tile, None), self.vae_dtype) + model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload) pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap) else: - pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) + # Reserve as much as an untiled decode could use (capped by what the device can provide), then size the tiles to fill that reservation: + # shrink the temporal tile until one tile fits, then grow the spatial tile while it still fits. + budget = min(memory_used, int(model_management.get_total_memory(self.device) * 0.8)) + model_management.load_models_gpu([self.patcher], memory_required=budget, force_full_load=self.disable_offload) + tile_t = samples_in.shape[2] + est = lambda tt, txy: self.memory_used_decode(self._tile_bounded_shape(samples_in.shape, txy, txy, tt), self.vae_dtype) + while tile_t > 2 and est(tile_t, tile) > budget: + tile_t = -(-tile_t // 2) + while tile * 2 <= max(samples_in.shape[3], samples_in.shape[4]) and est(tile_t, tile * 2) <= budget: + tile *= 2 + overlap = tile // 4 + pixel_samples = self.decode_tiled_3d(samples_in, tile_t=tile_t, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) pixel_samples = pixel_samples.to(self.output_device).movedim(1,-1) return pixel_samples + def _tile_bounded_shape(self, shape, tile_x, tile_y, tile_t): + """Clamp a latent shape to one tile for memory estimates: peak memory of a tiled decode is per-tile. Only caller-provided tile dims are clamped.""" + s = list(shape) + if len(s) == 5: + if tile_t is not None: + s[2] = min(s[2], tile_t) + if tile_y is not None: + s[3] = min(s[3], tile_y) + if tile_x is not None: + s[4] = min(s[4], tile_x) + else: + if tile_y is not None: + s[2] = min(s[2], tile_y) + if tile_x is not None: + s[3] = min(s[3], tile_x) + return tuple(s) + def decode_tiled(self, samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): self.throw_exception_if_invalid() - memory_used = self.memory_used_decode(samples.shape, self.vae_dtype) #TODO: calculate mem required for tile + memory_used = self.memory_used_decode(self._tile_bounded_shape(samples.shape, tile_x, tile_y, tile_t), self.vae_dtype) model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload) dims = samples.ndim - 2 args = {} @@ -1702,12 +1749,21 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None) elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B, TEModel.GEMMA_4_12B): - variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, - TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, - TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, - TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model] - clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant) - clip_target.tokenizer = variant.tokenizer + if te_model == TEModel.GEMMA_4_12B and "text_embedding_projection.video_aggregate_embed.weight" in clip_data[0]: + clip_target.clip = comfy.text_encoders.lt.ltxav_te( + **llama_detect(clip_data), + **comfy.text_encoders.lt.sd_detect(clip_data), + text_encoder_model=comfy.text_encoders.gemma4.gemma4_text_encoder_model(comfy.text_encoders.gemma4.Gemma4_12B), + text_encoder_key="gemma4", + ) + clip_target.tokenizer = comfy.text_encoders.lt.ltxav_gemma4_tokenizer(comfy.text_encoders.gemma4.Gemma4_12B.tokenizer) + else: + variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, + TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, + TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, + TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model] + clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant) + clip_target.tokenizer = variant.tokenizer tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None) elif te_model == TEModel.GEMMA_2_2B: if clip_type == CLIPType.PIXELDIT: @@ -1875,9 +1931,30 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.clip = comfy.text_encoders.kandinsky5.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.kandinsky5.Kandinsky5TokenizerImage elif clip_type == CLIPType.LTXV: - clip_target.clip = comfy.text_encoders.lt.ltxav_te(**llama_detect(clip_data), **comfy.text_encoders.lt.sd_detect(clip_data)) - clip_target.tokenizer = comfy.text_encoders.lt.LTXAVGemmaTokenizer - tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None) + te_models = [detect_te_model(sd) for sd in clip_data] + gemma4_models = { + TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, + TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, + TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, + TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B, + } + gemma4_type = next((model for model in te_models if model in gemma4_models), None) + if gemma4_type is None: + clip_target.clip = comfy.text_encoders.lt.ltxav_te(**llama_detect(clip_data), **comfy.text_encoders.lt.sd_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.lt.LTXAVGemmaTokenizer + gemma_sd = clip_data[te_models.index(TEModel.GEMMA_3_12B)] if TEModel.GEMMA_3_12B in te_models else clip_data[0] + tokenizer_data["spiece_model"] = gemma_sd.get("spiece_model", None) + else: + variant = gemma4_models[gemma4_type] + clip_target.clip = comfy.text_encoders.lt.ltxav_te( + **llama_detect(clip_data), + **comfy.text_encoders.lt.sd_detect(clip_data), + text_encoder_model=comfy.text_encoders.gemma4.gemma4_text_encoder_model(variant), + text_encoder_key="gemma4", + ) + clip_target.tokenizer = comfy.text_encoders.lt.ltxav_gemma4_tokenizer(variant.tokenizer) + gemma_sd = clip_data[te_models.index(gemma4_type)] + tokenizer_data["tokenizer_json"] = gemma_sd.get("tokenizer_json", None) elif clip_type == CLIPType.NEWBIE: clip_target.clip = comfy.text_encoders.newbie.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.newbie.NewBieTokenizer diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index 5163c1676..fc62bc7cc 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -1443,6 +1443,10 @@ class Gemma4UnifiedTokenizer(Gemma4Tokenizer): class Gemma4Model(sd1_clip.SDClipModel): model_class = None def __init__(self, device="cpu", layer="all", layer_idx=None, dtype=None, attention_mask=True, model_options={}): + llama_quantization_metadata = model_options.get("llama_quantization_metadata", None) + if llama_quantization_metadata is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = llama_quantization_metadata self.dtypes = set() self.dtypes.add(dtype) super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, dtype=dtype, special_tokens={"start": 2, "pad": 0}, layer_norm_hidden_state=False, model_class=self.model_class, enable_attention_masks=attention_mask, return_attention_masks=attention_mask, model_options=model_options) @@ -1474,8 +1478,19 @@ class Gemma4Model(sd1_clip.SDClipModel): 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_clip_model(model_class): + return type('Gemma4Model_', (Gemma4Model,), {'model_class': model_class}) + + +def gemma4_text_encoder_model(model_class): + return type('Gemma4TextEncoderModel_', (Gemma4Model,), { + 'model_class': model_class, + 'process_tokens': sd1_clip.SDClipModel.process_tokens, + }) + + def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None): - clip_model = type('Gemma4Model_', (Gemma4Model,), {'model_class': model_class}) + clip_model = gemma4_clip_model(model_class) class Gemma4TEModel_(sd1_clip.SD1ClipModel): def __init__(self, device="cpu", dtype=None, model_options={}): if llama_quantization_metadata is not None: diff --git a/comfy/text_encoders/lt.py b/comfy/text_encoders/lt.py index bc5cbae28..c512a7d48 100644 --- a/comfy/text_encoders/lt.py +++ b/comfy/text_encoders/lt.py @@ -81,6 +81,17 @@ class LTXAVGemmaTokenizer(sd1_clip.SD1Tokenizer): super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma3_12b", tokenizer=Gemma3_12BTokenizer) +def ltxav_gemma4_tokenizer(tokenizer): + class LTXAVGemma4Tokenizer(tokenizer): + def __init__(self, embedding_directory=None, tokenizer_data={}): + super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data) + gemma_tokenizer = getattr(self, self.clip) + if gemma_tokenizer.min_length == 1: + gemma_tokenizer.min_length = 1024 + + return LTXAVGemma4Tokenizer + + class Gemma3_12BModel(sd1_clip.SDClipModel): def __init__(self, device="cpu", layer="all", layer_idx=None, dtype=None, attention_mask=True, model_options={}): llama_quantization_metadata = model_options.get("llama_quantization_metadata", None) @@ -97,10 +108,10 @@ class Gemma3_12BModel(sd1_clip.SDClipModel): return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, stop_tokens=[106], presence_penalty=presence_penalty) # 106 is class DualLinearProjection(torch.nn.Module): - def __init__(self, in_dim, out_dim_video, out_dim_audio, dtype=None, device=None, operations=None): + def __init__(self, in_dim, out_dim_video, out_dim_audio, video_bias=True, audio_bias=True, dtype=None, device=None, operations=None): super().__init__() - self.audio_aggregate_embed = operations.Linear(in_dim, out_dim_audio, bias=True, dtype=dtype, device=device) - self.video_aggregate_embed = operations.Linear(in_dim, out_dim_video, bias=True, dtype=dtype, device=device) + self.audio_aggregate_embed = operations.Linear(in_dim, out_dim_audio, bias=audio_bias, dtype=dtype, device=device) + self.video_aggregate_embed = operations.Linear(in_dim, out_dim_video, bias=video_bias, dtype=dtype, device=device) def forward(self, x): source_dim = x.shape[-1] @@ -112,22 +123,28 @@ class DualLinearProjection(torch.nn.Module): return torch.cat((video, audio), dim=-1) class LTXAVTEModel(torch.nn.Module): - def __init__(self, dtype_llama=None, device="cpu", dtype=None, text_projection_type="single_linear", model_options={}): + def __init__(self, dtype_llama=None, device="cpu", dtype=None, text_projection_type="single_linear", text_encoder_model=Gemma3_12BModel, text_encoder_key="gemma3_12b", video_projection_dim=3840, audio_projection_dim=2048, video_projection_bias=None, audio_projection_bias=True, model_options={}): super().__init__() self.dtypes = set() self.dtypes.add(dtype) self.compat_mode = False self.text_projection_type = text_projection_type + self.text_encoder_key = text_encoder_key + self.execution_device = None - self.gemma3_12b = Gemma3_12BModel(device=device, dtype=dtype_llama, model_options=model_options, layer="all", layer_idx=None) + self.gemma3_12b = text_encoder_model(device=device, dtype=dtype_llama, model_options=model_options, layer="all", layer_idx=None) self.dtypes.add(dtype_llama) operations = self.gemma3_12b.operations # TODO + text_encoder_config = self.gemma3_12b.transformer.model.config + projection_in_dim = text_encoder_config.hidden_size * (text_encoder_config.num_hidden_layers + 1) + if video_projection_bias is None: + video_projection_bias = self.text_projection_type == "dual_linear" if self.text_projection_type == "single_linear": - self.text_embedding_projection = operations.Linear(3840 * 49, 3840, bias=False, dtype=dtype, device=device) + self.text_embedding_projection = operations.Linear(projection_in_dim, video_projection_dim, bias=video_projection_bias, dtype=dtype, device=device) elif self.text_projection_type == "dual_linear": - self.text_embedding_projection = DualLinearProjection(3840 * 49, 4096, 2048, dtype=dtype, device=device, operations=operations) + self.text_embedding_projection = DualLinearProjection(projection_in_dim, video_projection_dim, audio_projection_dim, video_bias=video_projection_bias, audio_bias=audio_projection_bias, dtype=dtype, device=device, operations=operations) def enable_compat_mode(self): # TODO: remove @@ -161,7 +178,7 @@ class LTXAVTEModel(torch.nn.Module): self.execution_device = None def encode_token_weights(self, token_weight_pairs): - token_weight_pairs = token_weight_pairs["gemma3_12b"] + token_weight_pairs = token_weight_pairs[self.text_encoder_key] out, pooled, extra = self.gemma3_12b.encode_token_weights(token_weight_pairs) out = out[:, :, -torch.sum(extra["attention_mask"]).item():] @@ -189,51 +206,54 @@ class LTXAVTEModel(torch.nn.Module): return out.to(device=out_device, dtype=torch.float), pooled, extra def generate(self, tokens, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty): - return self.gemma3_12b.generate(tokens["gemma3_12b"], do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty) + return self.gemma3_12b.generate(tokens[self.text_encoder_key], do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty) def load_sd(self, sd): - if "model.layers.47.self_attn.q_norm.weight" in sd: - return self.gemma3_12b.load_sd(sd) - else: - sdo = comfy.utils.state_dict_prefix_replace(sd, {"text_embedding_projection.aggregate_embed.weight": "text_embedding_projection.weight", "text_embedding_projection.": "text_embedding_projection."}, filter_keys=True) - if len(sdo) == 0: - sdo = sd + missing_all = [] + unexpected_all = [] - missing_all = [] - unexpected_all = [] + if "model.layers.0.self_attn.q_norm.weight" in sd: + gemma_sd = {k: v for k, v in sd.items() if not k.startswith("text_embedding_projection.")} + missing, unexpected = self.gemma3_12b.load_sd(gemma_sd) + missing_all.extend(missing) + unexpected_all.extend(unexpected) - for prefix, component in [("text_embedding_projection.", self.text_embedding_projection)]: - component_sd = {k.replace(prefix, ""): v for k, v in sdo.items() if k.startswith(prefix)} - if component_sd: - missing, unexpected = component.load_state_dict(component_sd, strict=False, assign=getattr(self, "can_assign_sd", False)) - missing_all.extend([f"{prefix}{k}" for k in missing]) - unexpected_all.extend([f"{prefix}{k}" for k in unexpected]) + sdo = comfy.utils.state_dict_prefix_replace(sd, {"text_embedding_projection.aggregate_embed.": "text_embedding_projection.", "text_embedding_projection.": "text_embedding_projection."}, filter_keys=True) + if len(sdo) == 0: + sdo = sd - if "model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.2.attn1.to_q.bias" not in sd: # TODO: remove - ww = sd.get("model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.0.attn1.to_q.bias", None) - if ww is not None: - if ww.shape[0] == 3840: - self.enable_compat_mode() - sdv = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.video_embeddings_connector.": ""}, filter_keys=True) - self.video_embeddings_connector.load_state_dict(sdv, strict=False, assign=getattr(self, "can_assign_sd", False)) - sda = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.audio_embeddings_connector.": ""}, filter_keys=True) - self.audio_embeddings_connector.load_state_dict(sda, strict=False, assign=getattr(self, "can_assign_sd", False)) + for prefix, component in [("text_embedding_projection.", self.text_embedding_projection)]: + component_sd = {k.replace(prefix, ""): v for k, v in sdo.items() if k.startswith(prefix)} + if component_sd: + missing, unexpected = component.load_state_dict(component_sd, strict=False, assign=getattr(self, "can_assign_sd", False)) + missing_all.extend([f"{prefix}{k}" for k in missing]) + unexpected_all.extend([f"{prefix}{k}" for k in unexpected]) - return (missing_all, unexpected_all) + if "model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.2.attn1.to_q.bias" not in sd: # TODO: remove + ww = sd.get("model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.0.attn1.to_q.bias", None) + if ww is not None: + if ww.shape[0] == 3840: + self.enable_compat_mode() + sdv = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.video_embeddings_connector.": ""}, filter_keys=True) + self.video_embeddings_connector.load_state_dict(sdv, strict=False, assign=getattr(self, "can_assign_sd", False)) + sda = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.audio_embeddings_connector.": ""}, filter_keys=True) + self.audio_embeddings_connector.load_state_dict(sda, strict=False, assign=getattr(self, "can_assign_sd", False)) + + return (missing_all, unexpected_all) def memory_estimation_function(self, token_weight_pairs, device=None): constant = 6.0 if comfy.model_management.should_use_bf16(device): constant /= 2.0 - token_weight_pairs = token_weight_pairs.get("gemma3_12b", []) + token_weight_pairs = token_weight_pairs.get(self.text_encoder_key, []) m = min([sum(1 for _ in itertools.takewhile(lambda x: x[0] == 0, sub)) for sub in token_weight_pairs]) num_tokens = sum(map(lambda a: len(a), token_weight_pairs)) - m num_tokens = max(num_tokens, 642) return num_tokens * constant * 1024 * 1024 -def ltxav_te(dtype_llama=None, llama_quantization_metadata=None, text_projection_type="single_linear"): +def ltxav_te(dtype_llama=None, llama_quantization_metadata=None, text_projection_type="single_linear", text_encoder_model=Gemma3_12BModel, text_encoder_key="gemma3_12b", video_projection_dim=3840, audio_projection_dim=2048, video_projection_bias=None, audio_projection_bias=True): class LTXAVTEModel_(LTXAVTEModel): def __init__(self, device="cpu", dtype=None, model_options={}): if llama_quantization_metadata is not None: @@ -241,16 +261,29 @@ def ltxav_te(dtype_llama=None, llama_quantization_metadata=None, text_projection model_options["llama_quantization_metadata"] = llama_quantization_metadata if dtype_llama is not None: dtype = dtype_llama - super().__init__(dtype_llama=dtype_llama, device=device, dtype=dtype, text_projection_type=text_projection_type, model_options=model_options) + super().__init__(dtype_llama=dtype_llama, device=device, dtype=dtype, text_projection_type=text_projection_type, text_encoder_model=text_encoder_model, text_encoder_key=text_encoder_key, video_projection_dim=video_projection_dim, audio_projection_dim=audio_projection_dim, video_projection_bias=video_projection_bias, audio_projection_bias=audio_projection_bias, model_options=model_options) return LTXAVTEModel_ def sd_detect(state_dict_list, prefix=""): for sd in state_dict_list: - if "{}text_embedding_projection.audio_aggregate_embed.bias".format(prefix) in sd: - return {"text_projection_type": "dual_linear"} - if "{}text_embedding_projection.weight".format(prefix) in sd or "{}text_embedding_projection.aggregate_embed.weight".format(prefix) in sd: - return {"text_projection_type": "single_linear"} + video_key = "{}text_embedding_projection.video_aggregate_embed.weight".format(prefix) + audio_key = "{}text_embedding_projection.audio_aggregate_embed.weight".format(prefix) + if video_key in sd and audio_key in sd: + return { + "text_projection_type": "dual_linear", + "video_projection_dim": sd[video_key].shape[0], + "audio_projection_dim": sd[audio_key].shape[0], + "video_projection_bias": "{}text_embedding_projection.video_aggregate_embed.bias".format(prefix) in sd, + "audio_projection_bias": "{}text_embedding_projection.audio_aggregate_embed.bias".format(prefix) in sd, + } + for key in ("{}text_embedding_projection.weight".format(prefix), "{}text_embedding_projection.aggregate_embed.weight".format(prefix)): + if key in sd: + return { + "text_projection_type": "single_linear", + "video_projection_dim": sd[key].shape[0], + "video_projection_bias": key.removesuffix("weight") + "bias" in sd, + } return {} diff --git a/comfy_extras/nodes_lt.py b/comfy_extras/nodes_lt.py index 8c85c92b1..a6e5c5d27 100644 --- a/comfy_extras/nodes_lt.py +++ b/comfy_extras/nodes_lt.py @@ -2,11 +2,14 @@ import nodes import node_helpers import torch import torchaudio +import comfy.ldm.lightricks.duration_head import comfy.model_management import comfy.model_sampling import comfy.samplers import comfy.utils +import logging import math +import re import numpy as np import av from io import BytesIO @@ -934,6 +937,243 @@ class LTXVReferenceAudio(io.ComfyNode): return io.NodeOutput(m, positive, negative) +class LTXVSpatioTemporalGuidance(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVSpatioTemporalGuidance", + display_name="LTXV Spatio-Temporal Guidance (STG)", + category="advanced/guidance", + description="Runs one extra pass per step with the self-attention of the selected blocks degraded to a value-passthrough, " + "then guides away from it - improving spatial detail and motion coherence.", + inputs=[ + io.Model.Input("model"), + io.Float.Input("scale", default=1.0, min=0.0, max=100.0, step=0.01, round=0.01), + io.String.Input("blocks", default="29", tooltip="Comma-separated transformer block indices to perturb."), + io.Float.Input("start_percent", default=0.0, min=0.0, max=1.0, step=0.001, advanced=True), + io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.001, advanced=True), + ], + outputs=[io.Model.Output()], + ) + + @classmethod + def execute(cls, model, scale, blocks, start_percent, end_percent) -> io.NodeOutput: + block_set = frozenset(int(b) for b in re.findall(r"\d+", blocks)) + + m = model.clone() + model_sampling = m.get_model_object("model_sampling") + sigma_start = model_sampling.percent_to_sigma(start_percent) + sigma_end = model_sampling.percent_to_sigma(end_percent) + + def post_cfg_function(args): + if scale == 0 or not block_set: + return args["denoised"] + + sigma_ = args["sigma"][0].item() + if sigma_ > sigma_start or sigma_ < sigma_end: + return args["denoised"] + + cond_pred = args["cond_denoised"] + cond = args["cond"] + cfg_result = args["denoised"] + x = args["input"] + + model_options = args["model_options"].copy() + transformer_options = model_options.get("transformer_options", {}).copy() + transformer_options["stg_self_attn_blocks"] = block_set + model_options["transformer_options"] = transformer_options + + (perturbed,) = comfy.samplers.calc_cond_batch(args["model"], [cond], x, args["sigma"], model_options) + + return cfg_result + (cond_pred - perturbed) * scale + + m.set_model_sampler_post_cfg_function(post_cfg_function) + return io.NodeOutput(m) + + +class LTXVModalityGuidance(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVModalityGuidance", + display_name="LTXV Modality Guidance (A/V coupling)", + category="advanced/guidance", + description="Cross-modal (audio-video) guidance for LTXV-AV. Runs one extra forward " + "pass per step with the a2v/v2a cross-attention severed, then pushes the " + "result toward the coupled prediction - strengthening audio-visual sync " + "(e.g. lip-sync). Reference default modality_scale is 3.0. Stacks with the " + "dual-CFG guider and STG. Set to 1.0 to disable (no extra pass).", + inputs=[ + io.Model.Input("model"), + io.Float.Input("modality_scale", default=3.0, min=1.0, max=100.0, step=0.1, round=0.01), + io.Float.Input("start_percent", default=0.0, min=0.0, max=1.0, step=0.001, advanced=True), + io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.001, advanced=True), + ], + outputs=[io.Model.Output()], + ) + + @classmethod + def execute(cls, model, modality_scale, start_percent, end_percent) -> io.NodeOutput: + m = model.clone() + model_sampling = m.get_model_object("model_sampling") + sigma_start = model_sampling.percent_to_sigma(start_percent) + sigma_end = model_sampling.percent_to_sigma(end_percent) + + def post_cfg_function(args): + if math.isclose(modality_scale, 1.0): + return args["denoised"] + + sigma_ = args["sigma"][0].item() + if sigma_ > sigma_start or sigma_ < sigma_end: + return args["denoised"] + + cond_pred = args["cond_denoised"] + cond = args["cond"] + cfg_result = args["denoised"] + x = args["input"] + + # Extra pass with audio-video cross-attention severed (both directions) + model_options = args["model_options"].copy() + transformer_options = model_options.get("transformer_options", {}).copy() + transformer_options["a2v_cross_attn"] = False + transformer_options["v2a_cross_attn"] = False + model_options["transformer_options"] = transformer_options + + (mod_pred,) = comfy.samplers.calc_cond_batch( + args["model"], [cond], x, args["sigma"], model_options + ) + + # (modality_scale - 1) * (cond - uncond_modality), per the reference guider. + return cfg_result + (cond_pred - mod_pred) * (modality_scale - 1.0) + + m.set_model_sampler_post_cfg_function(post_cfg_function) + return io.NodeOutput(m) + + +class Guider_LTXAVDualCFG(comfy.samplers.CFGGuider): + """CFG guider that applies separate guidance scales to the video and audio + modalities of a packed LTXV-AV latent. + """ + + def set_conds(self, positive, negative): + self.inner_set_conds({"positive": positive, "negative": negative}) + + def set_cfg(self, video_cfg, audio_cfg): + self.video_cfg = video_cfg + self.audio_cfg = audio_cfg + self.cfg = max(video_cfg, audio_cfg) + + def sample(self, noise, latent_image, *args, **kwargs): + # Capture the video/audio split from the nested latent before it is packed. + self._v_numel = None + if getattr(latent_image, "is_nested", False): + parts = latent_image.unbind() + if len(parts) >= 2: + self._v_numel = math.prod(parts[0].shape[1:]) + return super().sample(noise, latent_image, *args, **kwargs) + + def predict_noise(self, x, timestep, model_options={}, seed=None): + v = getattr(self, "_v_numel", None) + if v is None or math.isclose(self.video_cfg, self.audio_cfg): + # Not an AV latent, or equal scales: fall back to standard single-CFG. + self.cfg = self.video_cfg + return super().predict_noise(x, timestep, model_options, seed) + + video_cfg, audio_cfg = self.video_cfg, self.audio_cfg + + def dual_cfg(args): + # Noise-space: cond = x - cond_pred, uncond = x - uncond_pred; the + # returned tensor is subtracted from x by cfg_function. + cond, uncond = args["cond"], args["uncond"] + out = uncond + (cond - uncond) * video_cfg + out[..., v:] = uncond[..., v:] + (cond[..., v:] - uncond[..., v:]) * audio_cfg + return out + + # disable_cfg1_optimization so the uncond pass always runs even if one of the two scales is 1.0. + model_options = {**model_options, "sampler_cfg_function": dual_cfg, "disable_cfg1_optimization": True} + return super().predict_noise(x, timestep, model_options, seed) + + +class LTXVDualCFGGuider(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVDualCFGGuider", + display_name="LTXV Dual CFG Guider", + category="model/sampling/guiders", + description="Separate CFG scales for the video and audio modalities of a packed LTXV-AV latent.", + inputs=[ + io.Model.Input("model"), + io.Conditioning.Input("positive"), + io.Conditioning.Input("negative"), + io.Float.Input("video_cfg", default=3.0, min=0.0, max=100.0, step=0.1, round=0.01), + io.Float.Input("audio_cfg", default=7.0, min=0.0, max=100.0, step=0.1, round=0.01), + ], + outputs=[io.Guider.Output()], + ) + + @classmethod + def execute(cls, model, positive, negative, video_cfg, audio_cfg) -> io.NodeOutput: + guider = Guider_LTXAVDualCFG(model) + guider.set_conds(positive, negative) + guider.set_cfg(video_cfg, audio_cfg) + return io.NodeOutput(guider) + + +class LTXVDurationPredictor(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVDurationPredictor", + display_name="LTXV Duration Predictor", + category="conditioning/video_models", + description="Predicts the natural shot duration for a prompt using the LTX 2.4 duration " + "head (loaded with ModelPatchLoader), and snaps it to the VAE's 8k+1 frame grid.", + search_aliases=["auto duration", "duration head", "num_frames"], + inputs=[ + io.Model.Input("model"), + io.Conditioning.Input("positive"), + io.Custom("MODEL_PATCH").Input("duration_head", + tooltip="LTX 2.4 duration head loaded with ModelPatchLoader."), + io.Float.Input("frame_rate", default=24.0, min=1.0, max=120.0, step=0.01), + io.Float.Input("min_seconds", default=1.0, min=0.5, max=120.0, step=0.1), + io.Float.Input("max_seconds", default=20.0, min=0.5, max=120.0, step=0.1), + ], + outputs=[ + io.Int.Output(display_name="num_frames"), + io.Float.Output(display_name="seconds", tooltip="Raw (unclamped) predicted duration."), + ], + ) + + @classmethod + def execute(cls, model, positive, duration_head, frame_rate, min_seconds, max_seconds) -> io.NodeOutput: + dm = model.model.diffusion_model + head = duration_head.model + if not isinstance(head, comfy.ldm.lightricks.duration_head.DurationHead): + raise ValueError("The connected model_patch is not an LTX duration head.") + + context = positive[0][0] + meta = positive[0][1] + if context.shape[0] != 1: + context = context[:1] + + # Run the caption connectors exactly the way sampling does. + comfy.model_management.load_models_gpu([model, duration_head]) + device = model.load_device + head = head.to(device) + with torch.no_grad(): + context = context.to(device=device, dtype=model.model.get_dtype_inference()) + processed = dm.preprocess_text_embeds(context, unprocessed=meta.get("unprocessed_ltxav_embeds", False)) + video_tokens = processed[..., :dm.cross_attention_dim].float() + audio_tokens = processed[..., dm.cross_attention_dim:].float() + seconds = float(head(video_tokens, audio_tokens)[0]) + + num_frames = comfy.ldm.lightricks.duration_head.seconds_to_num_frames( + seconds, frame_rate, min_seconds, max_seconds) + logging.info("LTXV duration head predicted %.2fs -> %d frames @ %.2f fps", seconds, num_frames, frame_rate) + return io.NodeOutput(num_frames, seconds) + + class LtxvExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[io.ComfyNode]]: @@ -951,6 +1191,10 @@ class LtxvExtension(ComfyExtension): LTXVConcatAVLatent, LTXVSeparateAVLatent, LTXVReferenceAudio, + LTXVDualCFGGuider, + LTXVModalityGuidance, + LTXVSpatioTemporalGuidance, + LTXVDurationPredictor, ] diff --git a/comfy_extras/nodes_lt_audio.py b/comfy_extras/nodes_lt_audio.py index 3ff18d8d4..0924f3e9e 100644 --- a/comfy_extras/nodes_lt_audio.py +++ b/comfy_extras/nodes_lt_audio.py @@ -173,7 +173,7 @@ class LTXAVTextEncoderLoader(io.ComfyNode): node_id="LTXAVTextEncoderLoader", display_name="Load LTXV Audio Text Encoder", category="model/loaders", - description="Recipes:\nltxav: gemma 3 12B", + description="Recipes:\nltxav: gemma 3 12B or matching gemma 4 model", inputs=[ io.Combo.Input( "text_encoder", diff --git a/comfy_extras/nodes_model_patch.py b/comfy_extras/nodes_model_patch.py index 4d7bf7476..d81112932 100644 --- a/comfy_extras/nodes_model_patch.py +++ b/comfy_extras/nodes_model_patch.py @@ -10,6 +10,7 @@ import comfy.ldm.lumina.controlnet import comfy.ldm.supir.supir_modules import comfy.ldm.anima.lllite import comfy.ldm.wan.uni3c +import comfy.ldm.lightricks.duration_head from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel from comfy_api.latest import io from comfy.ldm.supir.supir_patch import SUPIRPatch @@ -296,6 +297,10 @@ class ModelPatchLoader: device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast) + elif any(k.endswith("duration_head.attention_pooler.query_tokens") for k in sd) or "attention_pooler.query_tokens" in sd: + sd = comfy.ldm.lightricks.duration_head.normalize_state_dict(sd) + sd = {k: v.float() for k, v in sd.items()} # tiny head, keep fp32 + model = comfy.ldm.lightricks.duration_head.DurationHead() elif "audio_proj.proj1.weight" in sd: model = MultiTalkModelPatch( audio_window=5, context_tokens=32, vae_scale=4, diff --git a/comfy_extras/nodes_textgen.py b/comfy_extras/nodes_textgen.py index 5a947d5c5..40004652c 100644 --- a/comfy_extras/nodes_textgen.py +++ b/comfy_extras/nodes_textgen.py @@ -1,3 +1,4 @@ +import re from comfy_api.latest import ComfyExtension, io from typing_extensions import override @@ -152,6 +153,64 @@ You are a Creative Assistant writing concise, action-focused image-to-video prom Style: realistic - cinematic - The woman glances at her watch and smiles warmly. She speaks in a cheerful, friendly voice, "I think we're right on time!" In the background, a café barista prepares drinks at the counter. The barista calls out in a clear, upbeat tone, "Two cappuccinos ready!" The sound of the espresso machine hissing softly blends with gentle background chatter and the light clinking of cups on saucers. """ +LTX24_T2V_SYSTEM_PROMPT = """You are given a user's short text-to-video request. Write a single, highly detailed audio-visual caption describing the video that best fulfills that request, in the EXACT style of the training captions used for this video model. The generated video is scored against the user's ORIGINAL request, so preserve every element the user stated; expand faithfully into the full caption style without contradicting or dropping anything they asked for. + +Match this captioning style precisely: + +1. Begin immediately with the action or visual detail. Do NOT use "The scene opens…", "We see…", "There is…". + +2. Objective, observable description only. Do not infer emotions or intentions — describe what is visible and audible (e.g. not "he looks sad" but "his eyebrows angle downward and his lips are pressed together"). + +3. Full visual detail: environment (materials, textures, lighting, colors), character appearance (clothing, posture, facial details), and the spatial positioning of all elements. When a human appears, identify them specifically (gendered terms when clearly implied; differentiate multiple people consistently) and describe visible physical attributes — apparent gender presentation, skin tone, estimated age group, hair color/length/style, build, clothing and accessories. Do not infer ethnicity, nationality, religion, or culture. + +4. Precise motion and cinematic description. For every shot you MUST include, woven naturally into the prose (never as tags or labels): + - Shot type (exactly one: extreme wide shot / wide shot / medium shot / medium close-up / close-up / extreme close-up) + - Camera motion (always stated; if none, explicitly say the camera remains static). Camera movement is expected and good — match the user if they specified it, otherwise choose the treatment that best presents the requested scene. + - Camera viewpoint relative to subject (front-facing / back-facing / side view / over-the-shoulder / top-down / low-angle / high-angle). + Express these as flowing prose: "a medium shot frames…, captured from a front-facing angle as the camera slowly pans…". Never as "medium shot, static camera —". + +5. Complete soundscape, integrated naturally: any dialogue (quote it exactly, in the original language), tone of voice, background music (type, mood, volume changes), and environmental sounds (footsteps, wind, traffic, animals). If the request implies sound, describe it plausibly. + +6. Strict chronological, real-time flow using transitions like "Initially…", "A moment later…", "Simultaneously…". Keep every stated action in motion. + +7. One single continuous paragraph. No bullet points, no section headers, no labels like "Audio:" or "Visual:". Exhaustive and lossless — include background elements, subtle movements, lighting, secondary sounds — detailed enough to reconstruct the scene. Aim for a rich, complete paragraph (roughly 150–220 words). + +If the user wrote in another language, produce the English caption of the same content. Output ONLY the caption text — no JSON, no preamble. + +AESTHETIC QUALITY (in addition to the above, without breaking the objective caption style): render the described scene with strong visual production value — cinematic, film-grade color and contrast, beautiful natural lighting, crisp fine detail and texture, pleasing composition and depth. Weave these quality descriptors naturally into the same observable prose (e.g. "warm cinematic lighting", "richly saturated film-grade color", "crisp high-resolution detail") — describe how the exact requested scene LOOKS at its most visually striking, never adding new objects or actions. Keep everything else (framing triple, soundscape, chronological single paragraph, faithfulness) exactly as specified. +""" + + +LTX24_I2V_SYSTEM_PROMPT = """You are given a REFERENCE IMAGE (the exact first frame of the video) and a user's short image-to-video request. Write a single, highly detailed audio-visual caption describing the video that BEGINS from this exact reference image and best fulfills that request, in the EXACT style of the training captions used for this video model. The generated video is scored against the user's ORIGINAL request, so preserve every element the user stated; expand faithfully into the full caption style without contradicting or dropping anything they asked for. + +FIRST-FRAME / IMAGE GROUNDING (do this first): the opening of your caption must match the reference image exactly — same subject(s), identity, appearance, clothing, setting, lighting, and composition as shown. The video starts on this frame; describe it faithfully, then narrate chronologically as the user's requested action unfolds from it. Never contradict, replace, or invent things not consistent with the image. Single continuous take — no hard cuts. + +Match this captioning style precisely: + +1. Begin immediately with the action or visual detail. Do NOT use "The scene opens…", "We see…", "There is…". + +2. Objective, observable description only. Do not infer emotions or intentions — describe what is visible and audible (e.g. not "he looks sad" but "his eyebrows angle downward and his lips are pressed together"). + +3. Full visual detail: environment (materials, textures, lighting, colors), character appearance (clothing, posture, facial details), and the spatial positioning of all elements — grounded in and consistent with the reference image. When a human appears, identify them specifically (gendered terms when clearly implied; differentiate multiple people consistently) and describe visible physical attributes — apparent gender presentation, skin tone, estimated age group, hair color/length/style, build, clothing and accessories. Do not infer ethnicity, nationality, religion, or culture. + +4. Precise motion and cinematic description. For every shot you MUST include, woven naturally into the prose (never as tags or labels): + - Shot type (exactly one: extreme wide shot / wide shot / medium shot / medium close-up / close-up / extreme close-up) — consistent with how the reference image is framed at the start. + - Camera motion (always stated; if none, explicitly say the camera remains static). Camera movement is expected and good — match the user if they specified it, otherwise choose the treatment that best presents the requested scene starting from this frame. + - Camera viewpoint relative to subject (front-facing / back-facing / side view / over-the-shoulder / top-down / low-angle / high-angle) — matching the reference image's viewpoint at the opening. + Express these as flowing prose: "a medium shot frames…, captured from a front-facing angle as the camera slowly pans…". Never as "medium shot, static camera —". + +5. Complete soundscape, integrated naturally: any dialogue (quote it exactly, in the original language), tone of voice, background music (type, mood, volume changes), and environmental sounds (footsteps, wind, traffic, animals). If the request implies sound, describe it plausibly. + +6. Strict chronological, real-time flow using transitions like "Initially…", "A moment later…", "Simultaneously…". Keep the user's requested motion/action central and in motion throughout. + +7. One single continuous paragraph. No bullet points, no section headers, no labels like "Audio:" or "Visual:". Exhaustive and lossless — include background elements, subtle movements, lighting, secondary sounds — detailed enough to reconstruct the scene. Aim for a rich, complete paragraph (roughly 150–220 words). + +If the user wrote in another language, produce the English caption of the same content. Output ONLY the caption text — no JSON, no preamble. + +AESTHETIC QUALITY (in addition to the above, without breaking the objective caption style or contradicting the reference image): render the described scene with strong visual production value — cinematic, film-grade color and contrast, beautiful natural lighting, crisp fine detail and texture, pleasing composition and depth. Weave these quality descriptors naturally into the same observable prose (e.g. "warm cinematic lighting", "richly saturated film-grade color", "crisp high-resolution detail") — describe how the exact requested scene, starting from this frame, LOOKS at its most visually striking, never adding new objects or actions and never contradicting the first frame. Keep everything else (first-frame grounding, framing triple, soundscape, chronological single paragraph, faithfulness) exactly as specified. +""" + + class TextGenerateLTX2Prompt(TextGenerate): @classmethod def define_schema(cls): @@ -167,11 +226,42 @@ class TextGenerateLTX2Prompt(TextGenerate): @classmethod def execute(cls, clip, prompt, max_length, sampling_mode, image=None, thinking=False, use_default_template=True, video=None, audio=None) -> io.NodeOutput: - if image is None: - formatted_prompt = f"system\n{LTX2_T2V_SYSTEM_PROMPT.strip()}\nuser\nUser Raw Input Prompt: {prompt}.\nmodel\n" + # Gemma 3 and Gemma 4 use different chat-turn markers and image tokens. + # The Gemma 4 text encoder is the LTX 2.4 path; Gemma 3 is LTX 2.0. + is_gemma4 = "gemma4" in getattr(clip.tokenizer, "clip_name", "") + + if is_gemma4: + if image is not None: + system = LTX24_I2V_SYSTEM_PROMPT.strip() + user_text = f"User Raw Input Prompt: {prompt}." + else: + system = LTX24_T2V_SYSTEM_PROMPT.strip() + user_text = f"user prompt: {prompt}" + think_prefix = "<|think|>\n" if thinking else "" + model_open = "" if thinking else "<|channel>final\n" + media = "<|image><|image|>\n\n" if image is not None else "" + formatted_prompt = ( + f"<|turn>system\n{think_prefix}{system}\n" + f"<|turn>user\n{media}{user_text}\n" + f"<|turn>model\n{model_open}" + ) else: - formatted_prompt = f"system\n{LTX2_I2V_SYSTEM_PROMPT.strip()}\nuser\n\n\n\nUser Raw Input Prompt: {prompt}.\nmodel\n" - return super().execute(clip, formatted_prompt, max_length, sampling_mode, image=image, thinking=thinking, use_default_template=use_default_template, video=video, audio=audio) + system = (LTX2_I2V_SYSTEM_PROMPT if image is not None else LTX2_T2V_SYSTEM_PROMPT).strip() + media = "\n\n" if image is not None else "" + formatted_prompt = ( + f"system\n{system}\n" + f"user\n{media}\nUser Raw Input Prompt: {prompt}.\n" + f"model\n" + ) + + out = super().execute(clip, formatted_prompt, max_length, sampling_mode, image=image, thinking=thinking, use_default_template=use_default_template, video=video, audio=audio) + + text = out.args[0] + text = re.sub(r".*?", "", text, flags=re.DOTALL) + if "" in text: # unclosed/truncated reasoning: keep what follows the last close + text = text.rsplit("", 1)[-1] + text = re.sub(r"|<\|channel>\w*\n?||<\|turn>\w*\n?", "", text).strip() + return io.NodeOutput(text) class TextgenExtension(ComfyExtension): From ce4fc130943162759a34f0c96935d3ab5e7e1bb6 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 11 Aug 2026 20:50:42 +0300 Subject: [PATCH 17/76] [Partner Nodes] feat(LTX): add new nodes for model version 2.5 (#15501) * [Partner Nodes] feat(LTX): add new nodes for model version 2.5 Signed-off-by: Alexander Piskun --- comfy_api_nodes/nodes_ltxv.py | 380 ++++++++++++++++++++++++++++++++++ 1 file changed, 380 insertions(+) diff --git a/comfy_api_nodes/nodes_ltxv.py b/comfy_api_nodes/nodes_ltxv.py index 878e04b4e..44723dde2 100644 --- a/comfy_api_nodes/nodes_ltxv.py +++ b/comfy_api_nodes/nodes_ltxv.py @@ -6,8 +6,12 @@ from typing_extensions import override from comfy_api.latest import IO, ComfyExtension, Input, InputImpl from comfy_api_nodes.util import ( ApiEndpoint, + download_url_to_video_output, get_number_of_images, + poll_op, + sync_op, sync_op_raw, + upload_audio_to_comfyapi, upload_images_to_comfyapi, validate_string, ) @@ -17,6 +21,11 @@ MODELS_MAP = { "LTX-2 (Fast)": "ltx-2-fast", } +V25_MODELS_MAP = { + "LTX-2.5 (Fast)": "ltx-2-5-fast", + "LTX-2.5 (Pro)": "ltx-2-5-pro", +} + class ExecuteTaskRequest(BaseModel): prompt: str = Field(...) @@ -26,6 +35,48 @@ class ExecuteTaskRequest(BaseModel): fps: int | None = Field(25) generate_audio: bool | None = Field(True) image_uri: str | None = Field(None) + last_frame_uri: str | None = Field(None) + + +class AudioToVideoRequest(BaseModel): + prompt: str = Field(...) + model: str = Field(...) + resolution: str = Field(...) + audio_uri: str = Field(...) + image_uri: str | None = Field(None) + + +class Ltx25SubmitResponse(BaseModel): + id: str = Field(...) + + +class Ltx25JobResult(BaseModel): + video_url: str | None = Field(None) + + +class Ltx25JobStatusResponse(BaseModel): + id: str = Field(...) + status: str = Field(...) + result: Ltx25JobResult | None = Field(None) + + +async def _v25_submit_and_poll(cls: type[IO.ComfyNode], route: str, data: BaseModel) -> IO.NodeOutput: + submit = await sync_op( + cls, + ApiEndpoint(f"/proxy/ltx/v2/{route}", "POST"), + response_model=Ltx25SubmitResponse, + data=data, + max_retries=1, + ) + job = await poll_op( + cls, + ApiEndpoint(f"/proxy/ltx/v2/{route}/{submit.id}"), + response_model=Ltx25JobStatusResponse, + status_extractor=lambda r: r.status, + ) + if not job.result or not job.result.video_url: + raise RuntimeError(f"LTX job {job.id} completed without a video URL.") + return IO.NodeOutput(await download_url_to_video_output(job.result.video_url, cls=cls)) PRICE_BADGE = IO.PriceBadge( @@ -43,6 +94,128 @@ PRICE_BADGE = IO.PriceBadge( """, ) +V25_PRICE_BADGE = IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model", "model.duration", "model.resolution"]), + expr=""" + ( + $prices := { + "ltx-2.5 (fast)": { + "1280x720":0.1287,"720x1280":0.1287, + "1920x1080":0.1859,"1080x1920":0.1859, + "2560x1440":0.2717,"1440x2560":0.2717, + "3840x2160":0.429,"2160x3840":0.429 + }, + "ltx-2.5 (pro)": { + "1280x720":0.1716,"720x1280":0.1716, + "1920x1080":0.2431,"1080x1920":0.2431 + } + }; + $model := $lookup(widgets, "model"); + $table := $type($model) = "string" ? $lookup($prices, $model) : undefined; + $res := $lookup(widgets, "model.resolution"); + $pps := $type($table) = "object" and $type($res) = "string" ? $lookup($table, $res) : undefined; + $durRaw := $lookup(widgets, "model.duration"); + $dur := $type($durRaw) in ["string", "number"] ? $number($durRaw) : undefined; + $type($pps) = "number" and $type($dur) = "number" + ? {"type":"usd","usd": $pps * $dur} + : undefined + ) + """, +) + +V25_A2V_PRICE_BADGE = IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model"]), + expr=""" + ( + $rates := {"ltx-2.5 (fast)":0.1859, "ltx-2.5 (pro)":0.2431}; + $model := $lookup(widgets, "model"); + $rate := $type($model) = "string" ? $lookup($rates, $model) : undefined; + $type($rate) = "number" + ? {"type":"usd","usd": $rate, "format":{"suffix":"/second"}} + : undefined + ) + """, +) + + +def _v25_generation_inputs( + durations: list[str], resolutions: list[str], fps_options: list[str], tooltip: str | None +) -> list: + return [ + IO.Combo.Input( + "duration", + options=durations, + default="8", + tooltip=tooltip, + ), + IO.Combo.Input( + "resolution", + options=resolutions, + default="1920x1080", + ), + IO.Combo.Input("fps", options=fps_options, default="25"), + IO.Boolean.Input( + "generate_audio", + default=True, + tooltip="When true, the generated video will include AI-generated audio matching the scene.", + advanced=True, + ), + ] + + +def _v25_model_combo() -> IO.DynamicCombo.Input: + return IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "LTX-2.5 (Fast)", + _v25_generation_inputs( + ["2", "3", "4", "5", "6", "8", "10", "12", "14", "16", "18", "20"], + [ + "1280x720", + "720x1280", + "1920x1080", + "1080x1920", + "2560x1440", + "1440x2560", + "3840x2160", + "2160x3840", + ], + ["24", "25", "48", "50"], + "Video duration in seconds. Durations over 10s require a 720p/1080p resolution and 24/25 FPS.", + ), + ), + IO.DynamicCombo.Option( + "LTX-2.5 (Pro)", + _v25_generation_inputs( + ["2", "3", "4", "5", "6", "8", "10"], + ["1280x720", "720x1280", "1920x1080", "1080x1920"], + ["24", "25", "50"], + "Video duration in seconds.", + ), + ), + ], + ) + + +def _v25_seed_input() -> IO.Int.Input: + return IO.Int.Input( + "seed", + default=42, + min=0, + max=0xFFFFFFFF, + control_after_generate=True, + tooltip="Seed to determine if node should re-run; " + "actual results are nondeterministic regardless of seed.", + ) + + +def _v25_validate_settings(model: dict) -> None: + if int(model["duration"]) > 10 and ( + int(model["fps"]) > 25 or model["resolution"] in ("2560x1440", "1440x2560", "3840x2160", "2160x3840") + ): + raise ValueError("Durations over 10s require a 720p or 1080p resolution and 24/25 FPS.") + class TextToVideoNode(IO.ComfyNode): @classmethod @@ -86,6 +259,7 @@ class TextToVideoNode(IO.ComfyNode): IO.Hidden.unique_id, ], is_api_node=True, + is_deprecated=True, price_badge=PRICE_BADGE, ) @@ -164,6 +338,7 @@ class ImageToVideoNode(IO.ComfyNode): IO.Hidden.unique_id, ], is_api_node=True, + is_deprecated=True, price_badge=PRICE_BADGE, ) @@ -203,12 +378,217 @@ class ImageToVideoNode(IO.ComfyNode): return IO.NodeOutput(InputImpl.VideoFromFile(BytesIO(response))) +class Ltx25TextToVideoNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="LtxApi25TextToVideo", + display_name="LTX 2.5 Text To Video", + category="partner/video/LTXV", + description="Professional-quality videos with customizable duration and resolution.", + inputs=[ + _v25_model_combo(), + IO.String.Input( + "prompt", + multiline=True, + default="", + ), + _v25_seed_input(), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=V25_PRICE_BADGE, + ) + + @classmethod + async def execute( + cls, + model: dict, + prompt: str, + seed: int = 42, + ) -> IO.NodeOutput: + validate_string(prompt, min_length=1, max_length=10000) + _v25_validate_settings(model) + return await _v25_submit_and_poll( + cls, + "text-to-video", + ExecuteTaskRequest( + prompt=prompt, + model=V25_MODELS_MAP[model["model"]], + duration=int(model["duration"]), + resolution=model["resolution"], + fps=int(model["fps"]), + generate_audio=model["generate_audio"], + ), + ) + + +class Ltx25ImageToVideoNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="LtxApi25ImageToVideo", + display_name="LTX 2.5 Image To Video", + category="partner/video/LTXV", + description="Professional-quality videos with customizable duration and resolution based on start image.", + inputs=[ + IO.Image.Input("image", tooltip="First frame to be used for the video."), + _v25_model_combo(), + IO.String.Input( + "prompt", + multiline=True, + default="", + ), + _v25_seed_input(), + IO.Image.Input( + "last_frame", + optional=True, + tooltip="Last frame to be used for the video.", + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=V25_PRICE_BADGE, + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + model: dict, + prompt: str, + seed: int = 42, + last_frame: Input.Image | None = None, + ) -> IO.NodeOutput: + validate_string(prompt, min_length=1, max_length=10000) + _v25_validate_settings(model) + if get_number_of_images(image) != 1: + raise ValueError("Currently only one input image is supported.") + last_frame_uri = None + if last_frame is not None: + if get_number_of_images(last_frame) != 1: + raise ValueError("Currently only one last frame image is supported.") + last_frame_uri = (await upload_images_to_comfyapi(cls, last_frame, max_images=1, mime_type="image/png"))[0] + return await _v25_submit_and_poll( + cls, + "image-to-video", + ExecuteTaskRequest( + image_uri=(await upload_images_to_comfyapi(cls, image, max_images=1, mime_type="image/png"))[0], + last_frame_uri=last_frame_uri, + prompt=prompt, + model=V25_MODELS_MAP[model["model"]], + duration=int(model["duration"]), + resolution=model["resolution"], + fps=int(model["fps"]), + generate_audio=model["generate_audio"], + ), + ) + + +class Ltx25AudioToVideoNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="LtxApi25AudioToVideo", + display_name="LTX 2.5 Audio To Video", + category="partner/video/LTXV", + description="Generate a video driven by an audio track, with an optional first frame image.", + inputs=[ + IO.Audio.Input( + "audio", + tooltip="Audio track driving the video. Its length (2-20 seconds) sets the video duration.", + ), + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "LTX-2.5 (Fast)", + [IO.Combo.Input("resolution", options=["1920x1080", "1080x1920"])], + ), + IO.DynamicCombo.Option( + "LTX-2.5 (Pro)", + [IO.Combo.Input("resolution", options=["1920x1080", "1080x1920"])], + ), + ], + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + ), + _v25_seed_input(), + IO.Image.Input( + "image", + optional=True, + tooltip="Optional first frame to be used for the video.", + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=V25_A2V_PRICE_BADGE, + ) + + @classmethod + async def execute( + cls, + audio: Input.Audio, + model: dict, + prompt: str, + seed: int = 42, + image: Input.Image | None = None, + ) -> IO.NodeOutput: + validate_string(prompt, min_length=1, max_length=10000) + audio_duration = audio["waveform"].shape[-1] / audio["sample_rate"] + if not 2 <= audio_duration <= 20: + raise ValueError(f"Audio duration must be between 2 and 20 seconds, got {audio_duration:.1f}s.") + image_uri = None + if image is not None: + if get_number_of_images(image) != 1: + raise ValueError("Currently only one input image is supported.") + image_uri = (await upload_images_to_comfyapi(cls, image, max_images=1, mime_type="image/png"))[0] + return await _v25_submit_and_poll( + cls, + "audio-to-video", + AudioToVideoRequest( + prompt=prompt, + model=V25_MODELS_MAP[model["model"]], + resolution=model["resolution"], + audio_uri=await upload_audio_to_comfyapi(cls, audio), + image_uri=image_uri, + ), + ) + + class LtxvApiExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: return [ TextToVideoNode, ImageToVideoNode, + Ltx25TextToVideoNode, + Ltx25ImageToVideoNode, + Ltx25AudioToVideoNode, ] From 2a19bbf0140743553d396d2ac49a2c73439195cc Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 11 Aug 2026 11:07:32 -0700 Subject: [PATCH 18/76] Fix for broken tiled audio decode. (#15502) --- comfy/sd.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/comfy/sd.py b/comfy/sd.py index 8bae76768..46c9acba1 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -1269,11 +1269,13 @@ class VAE: s[3] = min(s[3], tile_y) if tile_x is not None: s[4] = min(s[4], tile_x) - else: + elif len(s) == 4 and self.extra_1d_channel is None: if tile_y is not None: s[2] = min(s[2], tile_y) if tile_x is not None: s[3] = min(s[3], tile_x) + elif tile_x is not None: + s[-1] = min(s[-1], tile_x) return tuple(s) def decode_tiled(self, samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): From eb4a7b4fcfcedba4aba66b7297de4137ce0e1b2f Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Wed, 12 Aug 2026 03:20:39 +0800 Subject: [PATCH 19/76] chore: update workflow templates to v0.11.39 (#15504) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 25eaf8bc7..925b2eca4 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.48.7 -comfyui-workflow-templates==0.11.37 +comfyui-workflow-templates==0.11.39 comfyui-embedded-docs==0.5.9 torch torchsde From 2eaf09f50d1f8c2cfd382302bed36396a3633331 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 11 Aug 2026 22:38:08 +0300 Subject: [PATCH 20/76] [Partner Nodes] feat(Grok): add Grok Imagine Image 2.0 model (#15496) Signed-off-by: Alexander Piskun --- comfy_api_nodes/apis/grok.py | 2 + comfy_api_nodes/nodes_grok.py | 95 ++++++++++++++++++++++++++++------- 2 files changed, 78 insertions(+), 19 deletions(-) diff --git a/comfy_api_nodes/apis/grok.py b/comfy_api_nodes/apis/grok.py index 526d8c8ab..82dfc8b82 100644 --- a/comfy_api_nodes/apis/grok.py +++ b/comfy_api_nodes/apis/grok.py @@ -9,6 +9,7 @@ class ImageGenerationRequest(BaseModel): seed: int = Field(...) response_format: str = Field("url") resolution: str = Field(...) + quality: str | None = Field(None) class InputUrlObject(BaseModel): @@ -28,6 +29,7 @@ class ImageEditRequest(BaseModel): seed: int = Field(...) response_format: str = Field("url") aspect_ratio: str | None = Field(...) + quality: str | None = Field(None) class VideoGenerationRequest(BaseModel): diff --git a/comfy_api_nodes/nodes_grok.py b/comfy_api_nodes/nodes_grok.py index 672a3e537..c8f0f20d5 100644 --- a/comfy_api_nodes/nodes_grok.py +++ b/comfy_api_nodes/nodes_grok.py @@ -36,6 +36,26 @@ _GROK_VIDEO_MODEL_API_IDS = { "grok-imagine-video-1.5": "grok-imagine-video-1.5", } +_GROK_IMAGE_MODEL_API_IDS = { + "grok-imagine-image-2.0": "grok-imagine-image-2.0", +} + +_GROK_IMAGE_QUALITY_MODELS = {"grok-imagine-image-2.0"} + +_GROK_IMAGE_QUALITY_OPTIONS = ["medium", "low"] + +_GROK_IMAGE_EDIT_MAX_IMAGES = { + "grok-imagine-image-2.0": 3, + "grok-imagine-image-pro": 1, + "grok-imagine-image-quality": 3, + "grok-imagine-image": 3, +} + +_GROK_IMAGE_EDIT_ASPECT_RATIO_NEEDS_MULTIPLE = { + "grok-imagine-image-quality", + "grok-imagine-image", +} + _GROK_VOICE_OPTIONS = [ "none", "ara", @@ -132,6 +152,7 @@ class GrokImageNode(IO.ComfyNode): IO.Combo.Input( "model", options=[ + "grok-imagine-image-2.0", "grok-imagine-image-quality", "grok-imagine-image-pro", "grok-imagine-image", @@ -181,6 +202,12 @@ class GrokImageNode(IO.ComfyNode): "actual results are nondeterministic regardless of seed.", ), IO.Combo.Input("resolution", options=["1K", "2K"], optional=True), + IO.Combo.Input( + "quality", + options=_GROK_IMAGE_QUALITY_OPTIONS, + optional=True, + tooltip="Quality level, supported only by the grok-imagine-image-2.0 model.", + ), ], outputs=[ IO.Image.Output(), @@ -192,12 +219,15 @@ class GrokImageNode(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["model", "number_of_images", "resolution"]), + depends_on=IO.PriceBadgeDepends(widgets=["model", "number_of_images", "resolution", "quality"]), expr=""" ( - $rate := widgets.model = "grok-imagine-image-quality" - ? (widgets.resolution = "1k" ? 0.05 : 0.07) - : ($contains(widgets.model, "pro") ? 0.07 : 0.02); + $is1k := widgets.resolution = "1k"; + $rate := widgets.model = "grok-imagine-image-2.0" + ? (widgets.quality = "low" ? ($is1k ? 0.04 : 0.06) : ($is1k ? 0.06 : 0.08)) + : (widgets.model = "grok-imagine-image-quality" + ? ($is1k ? 0.05 : 0.07) + : ($contains(widgets.model, "pro") ? 0.07 : 0.02)); {"type":"usd","usd": $rate * widgets.number_of_images} ) """, @@ -213,18 +243,20 @@ class GrokImageNode(IO.ComfyNode): number_of_images: int, seed: int, resolution: str = "1K", + quality: str = "medium", ) -> IO.NodeOutput: validate_string(prompt, strip_whitespace=True, min_length=1) response = await sync_op( cls, ApiEndpoint(path="/proxy/xai/v1/images/generations", method="POST"), data=ImageGenerationRequest( - model=model, + model=_GROK_IMAGE_MODEL_API_IDS.get(model, model), prompt=prompt, aspect_ratio=aspect_ratio, n=number_of_images, seed=seed, resolution=resolution.lower(), + quality=quality if model in _GROK_IMAGE_QUALITY_MODELS else None, ), response_model=ImageGenerationResponse, ) @@ -255,7 +287,9 @@ _GROK_IMAGE_EDIT_ASPECT_RATIO_OPTIONS = [ ] -def _grok_image_edit_model_inputs(*, max_ref_images: int, with_aspect_ratio: bool): +def _grok_image_edit_model_inputs( + *, max_ref_images: int, with_aspect_ratio: bool, with_quality: bool = False, aspect_ratio_needs_multiple: bool = True +): inputs = [ IO.Autogrow.Input( "images", @@ -281,12 +315,18 @@ def _grok_image_edit_model_inputs(*, max_ref_images: int, with_aspect_ratio: boo display_mode=IO.NumberDisplay.number, ), ] + if with_quality: + inputs.append(IO.Combo.Input("quality", options=_GROK_IMAGE_QUALITY_OPTIONS)) if with_aspect_ratio: inputs.append( IO.Combo.Input( "aspect_ratio", options=_GROK_IMAGE_EDIT_ASPECT_RATIO_OPTIONS, - tooltip="Only allowed when multiple images are connected.", + tooltip=( + "Only allowed when multiple images are connected." + if aspect_ratio_needs_multiple + else "Aspect ratio of the edited image." + ), ) ) return inputs @@ -451,6 +491,15 @@ class GrokImageEditNodeV2(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "grok-imagine-image-2.0", + _grok_image_edit_model_inputs( + max_ref_images=3, + with_aspect_ratio=True, + with_quality=True, + aspect_ratio_needs_multiple=False, + ), + ), IO.DynamicCombo.Option( "grok-imagine-image-quality", _grok_image_edit_model_inputs(max_ref_images=3, with_aspect_ratio=True), @@ -488,18 +537,23 @@ class GrokImageEditNodeV2(IO.ComfyNode): is_api_node=True, price_badge=IO.PriceBadge( depends_on=IO.PriceBadgeDepends( - widgets=["model", "model.resolution", "model.number_of_images"], + widgets=["model", "model.resolution", "model.number_of_images", "model.quality"], ), expr=""" ( - $isQualityModel := widgets.model = "grok-imagine-image-quality"; + $is20 := widgets.model = "grok-imagine-image-2.0"; $isPro := $contains(widgets.model, "pro"); $res := $lookup(widgets, "model.resolution"); $n := $lookup(widgets, "model.number_of_images"); - $rate := $isQualityModel - ? ($res = "1k" ? 0.05 : 0.07) - : ($isPro ? 0.07 : 0.02); - $base := $isQualityModel ? 0.01 : 0.002; + $is1k := $res = "1k"; + $rate := $is20 + ? ($lookup(widgets, "model.quality") = "low" + ? ($is1k ? 0.04 : 0.06) + : ($is1k ? 0.06 : 0.08)) + : (widgets.model = "grok-imagine-image-quality" + ? ($is1k ? 0.05 : 0.07) + : ($isPro ? 0.07 : 0.02)); + $base := ($is20 or widgets.model = "grok-imagine-image-quality") ? 0.01 : 0.002; $output := $rate * $n; $isPro ? {"type":"usd","usd": $base + $output} @@ -525,13 +579,15 @@ class GrokImageEditNodeV2(IO.ComfyNode): image_tensors: list[Input.Image] = [t for t in images_dict.values() if t is not None] n_images = sum(get_number_of_images(t) for t in image_tensors) + max_images = _GROK_IMAGE_EDIT_MAX_IMAGES.get(model_id, 3) if n_images < 1: raise ValueError("At least one image is required for editing.") - if model_id == "grok-imagine-image-pro" and n_images > 1: - raise ValueError("The pro model supports only 1 input image.") - if model_id != "grok-imagine-image-pro" and n_images > 3: - raise ValueError("A maximum of 3 input images is supported.") - if aspect_ratio != "auto" and n_images == 1: + if n_images > max_images: + raise ValueError( + f"The {model_id} model supports at most {max_images} input " + f"image{'s' if max_images > 1 else ''}; {n_images} are connected." + ) + if aspect_ratio != "auto" and model_id in _GROK_IMAGE_EDIT_ASPECT_RATIO_NEEDS_MULTIPLE and n_images == 1: raise ValueError( "Custom aspect ratio is only allowed when multiple images are connected to the image input." ) @@ -547,7 +603,7 @@ class GrokImageEditNodeV2(IO.ComfyNode): cls, ApiEndpoint(path="/proxy/xai/v1/images/edits", method="POST"), data=ImageEditRequest( - model=model_id, + model=_GROK_IMAGE_MODEL_API_IDS.get(model_id, model_id), images=[ InputUrlObject(url=f"data:image/png;base64,{tensor_to_base64_string(i)}") for i in flat_tensors ], @@ -556,6 +612,7 @@ class GrokImageEditNodeV2(IO.ComfyNode): n=number_of_images, seed=seed, aspect_ratio=None if aspect_ratio == "auto" else aspect_ratio, + quality=model.get("quality") if model_id in _GROK_IMAGE_QUALITY_MODELS else None, ), response_model=ImageGenerationResponse, ) From d9f9d2ba1291ebc35ea3cce445f2f0c64deaceb2 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 11 Aug 2026 12:54:15 -0700 Subject: [PATCH 21/76] Fix some clip vision regression. (#15506) --- comfy/clip_model.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/comfy/clip_model.py b/comfy/clip_model.py index d7d3f994c..26cc5d7ee 100644 --- a/comfy/clip_model.py +++ b/comfy/clip_model.py @@ -314,13 +314,18 @@ class CLIPVisionModelProjection(torch.nn.Module): if "projection_dim" in config_dict: self.visual_projection = operations.Linear(config_dict["hidden_size"], config_dict["projection_dim"], bias=False) else: - self.visual_projection = lambda a: a + self.visual_projection = torch.nn.Identity() if "llava3" == config_dict.get("projector_type", None): self.multi_modal_projector = LlavaProjector(config_dict["hidden_size"], 4096, dtype, device, operations) else: self.multi_modal_projector = None + def _load_from_state_dict(self, state_dict, prefix, *args, **kwargs): + if "{}visual_projection.weight".format(prefix) not in state_dict: + self.visual_projection = torch.nn.Identity() + super()._load_from_state_dict(state_dict, prefix, *args, **kwargs) + def forward(self, *args, **kwargs): x = self.vision_model(*args, **kwargs) out = self.visual_projection(x[2]) From bbb4b04caa37b4608db32163899fdda148a041e5 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 11 Aug 2026 12:54:44 -0700 Subject: [PATCH 22/76] Don't depend on transformers for mistral and llama tokenizers. (#15503) --- comfy/text_encoders/bpe_tokenizer.py | 333 +++++++++++++++++++++++++++ comfy/text_encoders/flux.py | 43 +--- comfy/text_encoders/hunyuan_video.py | 2 +- 3 files changed, 339 insertions(+), 39 deletions(-) create mode 100644 comfy/text_encoders/bpe_tokenizer.py diff --git a/comfy/text_encoders/bpe_tokenizer.py b/comfy/text_encoders/bpe_tokenizer.py new file mode 100644 index 000000000..e49e36ca0 --- /dev/null +++ b/comfy/text_encoders/bpe_tokenizer.py @@ -0,0 +1,333 @@ +""" +Pure-Python byte-level BPE tokenizer. +Supports loading from HuggingFace tokenizer.json (LLaMA-style) +and from Mistral tekken JSON blobs. +No dependency on the `transformers`, `tokenizers`, or `regex` packages. +""" +import base64 +import json +import os +import re +import unicodedata + + +# This is also the default pattern used by the previous MistralConverter path. +_LLAMA_PATTERN = r"""(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+""" +_CONTRACTIONS = ("'re", "'ve", "'ll", "'s", "'t", "'m", "'d") + + +def _is_letter(c): + return unicodedata.category(c)[0] == "L" + + +def _is_number(c): + return unicodedata.category(c)[0] == "N" + + +def _is_whitespace(c): + return c in " \t\n\r\v\f\x85\u2028\u2029" or unicodedata.category(c) == "Zs" + + +def _split_llama(text): + pieces = [] + i = 0 + while i < len(text): + contraction = None + if text[i] == "'": + for suffix in _CONTRACTIONS: + if text[i:i + len(suffix)].casefold() == suffix: + contraction = text[i:i + len(suffix)] + break + if contraction is not None: + pieces.append(contraction) + i += len(contraction) + continue + + j = i + if text[j] not in "\r\n" and not _is_letter(text[j]) and not _is_number(text[j]): + j += 1 + if j < len(text) and _is_letter(text[j]): + j += 1 + while j < len(text) and _is_letter(text[j]): + j += 1 + pieces.append(text[i:j]) + i = j + continue + + if _is_number(text[i]): + j = i + 1 + while j < len(text) and j - i < 3 and _is_number(text[j]): + j += 1 + pieces.append(text[i:j]) + i = j + continue + + j = i + if text[j] == " ": + j += 1 + punct_start = j + while j < len(text) and not _is_whitespace(text[j]) and not _is_letter(text[j]) and not _is_number(text[j]): + j += 1 + if j > punct_start: + while j < len(text) and text[j] in "\r\n": + j += 1 + pieces.append(text[i:j]) + i = j + continue + + if _is_whitespace(text[i]): + j = i + 1 + while j < len(text) and _is_whitespace(text[j]): + j += 1 + last_newline = max(text.rfind("\r", i, j), text.rfind("\n", i, j)) + if last_newline >= i: + j = last_newline + 1 + elif j < len(text) and j - i > 1: + j -= 1 + pieces.append(text[i:j]) + i = j + continue + + pieces.append(text[i]) + i += 1 + return pieces + + +def _make_split_pattern(pattern_str): + if pattern_str != _LLAMA_PATTERN: + raise ValueError(f"Unsupported tokenizer split pattern: {pattern_str}") + return _split_llama + + +def _bytes_to_unicode(): + bs = (list(range(ord("!"), ord("~") + 1)) + + list(range(ord("¡"), ord("¬") + 1)) + + list(range(ord("®"), ord("ÿ") + 1))) + cs = bs[:] + n = 0 + for b in range(2**8): + if b not in bs: + bs.append(b) + cs.append(2**8 + n) + n += 1 + cs = [chr(n) for n in cs] + return dict(zip(bs, cs)) + + +class BPETokenizer: + """Byte-level BPE tokenizer with optional BOS prepending.""" + + def __init__(self, vocab, merges_by_pair, special_token_ids, pattern_str, + byte_encoder, byte_decoder, bos_id=None): + self._vocab = vocab # str -> int + self._inv_vocab = {v: k for k, v in vocab.items()} + self._merges = merges_by_pair # (str, str) -> priority int + self._special_token_ids = special_token_ids # str -> int + self._special_ids = set(special_token_ids.values()) + self._byte_encoder = byte_encoder + self._byte_decoder = byte_decoder + self._bos_id = bos_id + + self._split = _make_split_pattern(pattern_str) + sorted_specials = sorted(special_token_ids.keys(), key=len, reverse=True) + if sorted_specials: + self._special_split = re.compile( + '(' + '|'.join(re.escape(s) for s in sorted_specials) + ')' + ) + else: + self._special_split = None + + def _bpe_encode_piece(self, chars): + if len(chars) <= 1: + return chars + while True: + min_rank = float('inf') + best_pair = None + for i in range(len(chars) - 1): + r = self._merges.get((chars[i], chars[i + 1]), float('inf')) + if r < min_rank: + min_rank = r + best_pair = (chars[i], chars[i + 1]) + if best_pair is None: + break + merged = best_pair[0] + best_pair[1] + new_chars = [] + i = 0 + while i < len(chars): + if i < len(chars) - 1 and chars[i] == best_pair[0] and chars[i + 1] == best_pair[1]: + new_chars.append(merged) + i += 2 + else: + new_chars.append(chars[i]) + i += 1 + chars = new_chars + if len(chars) == 1: + break + return chars + + def _encode_raw(self, text): + ids = [] + parts = self._special_split.split(text) if self._special_split else [text] + for part in parts: + if not part: + continue + if part in self._special_token_ids: + ids.append(self._special_token_ids[part]) + else: + for piece in self._split(part): + byte_chars = [self._byte_encoder[b] for b in piece.encode('utf-8')] + for tok in self._bpe_encode_piece(byte_chars): + ids.append(self._vocab[tok]) + return ids + + def __call__(self, text): + ids = self._encode_raw(text) + if self._bos_id is not None: + ids = [self._bos_id] + ids + return {"input_ids": ids} + + def get_vocab(self): + return dict(self._vocab) + + def decode(self, token_ids, skip_special_tokens=True): + buf = bytearray() + for tid in token_ids: + s = self._inv_vocab.get(tid, '') + if tid in self._special_ids: + if not skip_special_tokens: + buf.extend(s.encode('utf-8')) + else: + for c in s: + buf.append(self._byte_decoder[c]) + return buf.decode('utf-8', errors='replace') + + +def _extract_pattern(pretok): + if pretok.get('type') == 'Sequence': + for sub in pretok.get('pretokenizers', []): + if sub.get('type') == 'Split': + pat = sub.get('pattern', {}) + if 'Regex' in pat: + return pat['Regex'] + elif pretok.get('type') == 'Split': + pat = pretok.get('pattern', {}) + if 'Regex' in pat: + return pat['Regex'] + return None + + +def _extract_bos_id(post_processor, special_token_ids): + if post_processor.get('type') == 'TemplateProcessing': + single = post_processor.get('single', []) + if single and 'SpecialToken' in single[0]: + bos_str = single[0]['SpecialToken']['id'] + return special_token_ids.get(bos_str) + return None + + +def from_tokenizer_json(path): + """Load a BPETokenizer from a directory containing tokenizer.json.""" + tok_file = os.path.join(path, 'tokenizer.json') + with open(tok_file, encoding='utf-8') as f: + data = json.load(f) + + vocab = dict(data['model']['vocab']) # str -> int + + merges_by_pair = {} + for i, merge_str in enumerate(data['model'].get('merges', [])): + a, b = merge_str.split(' ', 1) + if (a, b) not in merges_by_pair: + merges_by_pair[(a, b)] = i + + special_token_ids = {} + for tok in data.get('added_tokens', []): + special_token_ids[tok['content']] = tok['id'] + vocab[tok['content']] = tok['id'] # include in vocab for inv_vocab decode + + pattern = _extract_pattern(data.get('pre_tokenizer', {})) + if pattern is None: + raise ValueError(f"Could not extract regex pattern from {tok_file}") + + bos_id = _extract_bos_id(data.get('post_processor', {}), special_token_ids) + + byte_encoder = _bytes_to_unicode() + byte_decoder = {v: k for k, v in byte_encoder.items()} + + return BPETokenizer(vocab, merges_by_pair, special_token_ids, pattern, + byte_encoder, byte_decoder, bos_id=bos_id) + + +def from_tekken_json(data): + """Build a BPETokenizer from a Mistral tekken JSON blob (bytes or str).""" + mistral_vocab = json.loads(data) + config = mistral_vocab["config"] + + byte_encoder = _bytes_to_unicode() + byte_decoder = {v: k for k, v in byte_encoder.items()} + + def tbts(b): + return "".join(byte_encoder[ord(c)] for c in b.decode("latin-1")) + + special_token_offset = config["default_num_special_tokens"] + max_vocab = config["default_vocab_size"] - special_token_offset + + raw_vocab = {} + for w in mistral_vocab["vocab"]: + r = w["rank"] + if r >= max_vocab: + continue + raw_vocab[base64.b64decode(w["token_bytes"])] = r + special_token_offset + + special_tokens_dict = {} + for w in mistral_vocab["special_tokens"]: + if "token_bytes" in w: + special_tokens_dict[base64.b64decode(w["token_bytes"])] = w["rank"] + else: + special_tokens_dict[w["token_str"]] = w["rank"] + + all_special = list(special_tokens_dict.keys()) + combined = dict(special_tokens_dict) + combined.update(raw_vocab) + + bpe_vocab = {} + merge_triples = [] + for token, rank in combined.items(): + if token not in all_special: + bpe_vocab[tbts(token)] = rank + if len(token) == 1: + continue + local = [] + for i in range(1, len(token)): + pl, pr = token[:i], token[i:] + if pl in combined and pr in combined and (pl + pr) in combined: + local.append((pl, pr, rank)) + local.sort(key=lambda x: (combined[x[0]], combined[x[1]])) + merge_triples.extend(local) + else: + tok_str = token.decode("utf-8", errors="replace") if isinstance(token, bytes) else token + bpe_vocab[tok_str] = rank + + merge_triples.sort(key=lambda v: v[2]) + + merges_by_pair = {} + for i, (pl, pr, _) in enumerate(merge_triples): + pair = (tbts(pl), tbts(pr)) + if pair not in merges_by_pair: + merges_by_pair[pair] = i + + special_str_ids = {} + for tok in all_special: + tok_str = tok.decode("utf-8", errors="replace") if isinstance(tok, bytes) else tok + if tok_str in bpe_vocab: + special_str_ids[tok_str] = bpe_vocab[tok_str] + + return BPETokenizer(bpe_vocab, merges_by_pair, special_str_ids, _LLAMA_PATTERN, + byte_encoder, byte_decoder, bos_id=None) + + +class LlamaTokenizerFast: + """Drop-in replacement for transformers.LlamaTokenizerFast (read-only use).""" + + @staticmethod + def from_pretrained(path, **kwargs): + return from_tokenizer_json(path) diff --git a/comfy/text_encoders/flux.py b/comfy/text_encoders/flux.py index d5eb91dcb..fbdb1d13a 100644 --- a/comfy/text_encoders/flux.py +++ b/comfy/text_encoders/flux.py @@ -3,11 +3,10 @@ import comfy.text_encoders.t5 import comfy.text_encoders.sd3_clip import comfy.text_encoders.llama import comfy.model_management -from transformers import T5TokenizerFast, LlamaTokenizerFast, Qwen2Tokenizer +from transformers import T5TokenizerFast, Qwen2Tokenizer +from .bpe_tokenizer import from_tekken_json import torch import os -import json -import base64 class T5XXLTokenizer(sd1_clip.SDTokenizer): def __init__(self, embedding_directory=None, tokenizer_data={}): @@ -75,45 +74,13 @@ def flux_clip(dtype_t5=None, t5_quantization_metadata=None): def load_mistral_tokenizer(data): if torch.is_tensor(data): data = data.numpy().tobytes() + return {"tokenizer_object": from_tekken_json(data)} - try: - from transformers.integrations.mistral import MistralConverter - except ModuleNotFoundError: - from transformers.models.pixtral.convert_pixtral_weights_to_hf import MistralConverter - - mistral_vocab = json.loads(data) - - special_tokens = {} - vocab = {} - - max_vocab = mistral_vocab["config"]["default_vocab_size"] - max_vocab -= len(mistral_vocab["special_tokens"]) - - for w in mistral_vocab["vocab"]: - r = w["rank"] - if r >= max_vocab: - continue - - vocab[base64.b64decode(w["token_bytes"])] = r - - for w in mistral_vocab["special_tokens"]: - if "token_bytes" in w: - special_tokens[base64.b64decode(w["token_bytes"])] = w["rank"] - else: - special_tokens[w["token_str"]] = w["rank"] - - all_special = [] - for v in special_tokens: - all_special.append(v) - - special_tokens.update(vocab) - vocab = special_tokens - return {"tokenizer_object": MistralConverter(vocab=vocab, additional_special_tokens=all_special).converted(), "legacy": False} class MistralTokenizerClass: @staticmethod - def from_pretrained(path, **kwargs): - return LlamaTokenizerFast(**kwargs) + def from_pretrained(path, tokenizer_object=None, **kwargs): + return tokenizer_object class Mistral3Tokenizer(sd1_clip.SDTokenizer): def __init__(self, embedding_directory=None, embedding_size=5120, embedding_key='mistral3_24b', tokenizer_data={}): diff --git a/comfy/text_encoders/hunyuan_video.py b/comfy/text_encoders/hunyuan_video.py index 2ddb4da60..932a3d49b 100644 --- a/comfy/text_encoders/hunyuan_video.py +++ b/comfy/text_encoders/hunyuan_video.py @@ -2,7 +2,7 @@ from comfy import sd1_clip import comfy.model_management import comfy.text_encoders.llama from .hunyuan_image import HunyuanImageTokenizer -from transformers import LlamaTokenizerFast +from .bpe_tokenizer import LlamaTokenizerFast import torch import os import numbers From 024cbc5fc1c779ea7905356d3f3239b90dd0dae3 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 11 Aug 2026 13:45:52 -0700 Subject: [PATCH 23/76] Remove potentially problematic process_tokens method. (#15507) --- comfy/text_encoders/gemma4.py | 4 ---- comfy/text_encoders/lumina2.py | 4 ---- 2 files changed, 8 deletions(-) diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index fc62bc7cc..5b0b968c9 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -1451,10 +1451,6 @@ class Gemma4Model(sd1_clip.SDClipModel): self.dtypes.add(dtype) super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, dtype=dtype, special_tokens={"start": 2, "pad": 0}, layer_norm_hidden_state=False, model_class=self.model_class, enable_attention_masks=attention_mask, return_attention_masks=attention_mask, model_options=model_options) - def process_tokens(self, tokens, device): - embeds, _, _, _ = super().process_tokens(tokens, device) - return embeds - def generate(self, tokens, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty=0.0): if isinstance(tokens, dict): tokens = next(iter(tokens.values())) diff --git a/comfy/text_encoders/lumina2.py b/comfy/text_encoders/lumina2.py index b1f1dbb9f..e44920203 100644 --- a/comfy/text_encoders/lumina2.py +++ b/comfy/text_encoders/lumina2.py @@ -49,10 +49,6 @@ class Gemma3_4B_Vision_Model(sd1_clip.SDClipModel): super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, dtype=dtype, special_tokens={"start": 2, "pad": 0}, layer_norm_hidden_state=False, model_class=comfy.text_encoders.llama.Gemma3_4B_Vision, enable_attention_masks=attention_mask, return_attention_masks=attention_mask, model_options=model_options) - def process_tokens(self, tokens, device): - embeds, _, _, _ = super().process_tokens(tokens, device) - return embeds - class LuminaModel(sd1_clip.SD1ClipModel): def __init__(self, device="cpu", dtype=None, model_options={}, name="gemma2_2b", clip_model=Gemma2_2BModel): super().__init__(device=device, dtype=dtype, name=name, clip_model=clip_model, model_options=model_options) From c2bcbecd82ec5ae66594340b395c24ef0217b238 Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Tue, 11 Aug 2026 16:46:59 -0400 Subject: [PATCH 24/76] ComfyUI v0.32.0 --- comfyui_version.py | 2 +- pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/comfyui_version.py b/comfyui_version.py index 7358bab8f..568e75b88 100644 --- a/comfyui_version.py +++ b/comfyui_version.py @@ -1,3 +1,3 @@ # This file is automatically generated by the build process when version is # updated in pyproject.toml. -__version__ = "0.31.0" +__version__ = "0.32.0" diff --git a/pyproject.toml b/pyproject.toml index 4f80dc5cf..941b04d45 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ComfyUI" -version = "0.31.0" +version = "0.32.0" readme = "README.md" license = { file = "LICENSE" } requires-python = ">=3.10" From 27bca654eb9a70237d93f56a6ea336ab55f8925d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Wed, 12 Aug 2026 00:58:00 +0300 Subject: [PATCH 25/76] Fix KSamplerAdvanced with add_noise disabled on nested latents (#15447) --- comfy/sample.py | 5 +++++ comfy_extras/nodes_custom_sampler.py | 10 +--------- nodes.py | 2 +- 3 files changed, 7 insertions(+), 10 deletions(-) diff --git a/comfy/sample.py b/comfy/sample.py index 2be0cae5f..617816882 100644 --- a/comfy/sample.py +++ b/comfy/sample.py @@ -37,6 +37,11 @@ def prepare_noise(latent_image, seed, noise_inds=None): return noises +def prepare_empty_noise(latent_image): + if latent_image.is_nested: + return comfy.nested_tensor.NestedTensor([torch.zeros_like(t, device="cpu") for t in latent_image.unbind()]) + return torch.zeros_like(latent_image, device="cpu") + def fix_empty_latent_channels(model, latent_image, downscale_ratio_spacial=None, downscale_ratio_temporal=None): if latent_image.is_nested: return latent_image diff --git a/comfy_extras/nodes_custom_sampler.py b/comfy_extras/nodes_custom_sampler.py index d5aa730d2..c73a8f6dc 100644 --- a/comfy_extras/nodes_custom_sampler.py +++ b/comfy_extras/nodes_custom_sampler.py @@ -718,15 +718,7 @@ class Noise_EmptyNoise: self.seed = 0 def generate_noise(self, input_latent): - latent_image = input_latent["samples"] - if latent_image.is_nested: - tensors = latent_image.unbind() - zeros = [] - for t in tensors: - zeros.append(torch.zeros(t.shape, dtype=t.dtype, layout=t.layout, device="cpu")) - return comfy.nested_tensor.NestedTensor(zeros) - else: - return torch.zeros(latent_image.shape, dtype=latent_image.dtype, layout=latent_image.layout, device="cpu") + return comfy.sample.prepare_empty_noise(input_latent["samples"]) class Noise_RandomNoise: diff --git a/nodes.py b/nodes.py index a7f91720f..ec298e1de 100644 --- a/nodes.py +++ b/nodes.py @@ -1570,7 +1570,7 @@ def common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, latent_image = comfy.sample.fix_empty_latent_channels(model, latent_image, latent.get("downscale_ratio_spacial", None), latent.get("downscale_ratio_temporal", None)) if disable_noise: - noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu") + noise = comfy.sample.prepare_empty_noise(latent_image) else: batch_inds = latent["batch_index"] if "batch_index" in latent else None noise = comfy.sample.prepare_noise(latent_image, seed, batch_inds) From 1108f2ac5e412b27accb0e5d51c90ef2ba39784d Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Wed, 12 Aug 2026 13:24:24 +0800 Subject: [PATCH 26/76] chore: update workflow templates to v0.11.40 (#15522) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 925b2eca4..f461e3b76 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.48.7 -comfyui-workflow-templates==0.11.39 +comfyui-workflow-templates==0.11.40 comfyui-embedded-docs==0.5.9 torch torchsde From 945ffca32e4d6c05d6f1bf37101601a30318b396 Mon Sep 17 00:00:00 2001 From: Robin Huang Date: Tue, 11 Aug 2026 23:21:08 -0700 Subject: [PATCH 27/76] chore: replace api nodes -> partner nodes in README (#15519) --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 1721cefa5..c4dfc8be1 100644 --- a/README.md +++ b/README.md @@ -37,7 +37,7 @@ ComfyUI is the AI creation engine for visual professionals who demand control over every model, every parameter, and every output. Its powerful and modular node graph interface empowers creatives to generate images, videos, 3D models, audio, and more... - ComfyUI natively supports the latest open-source state of the art models. -- API nodes provide access to the best closed source models such as Nano Banana, Seedance, Hunyuan3D, etc. +- [Partner nodes](https://docs.comfy.org/tutorials/partner-nodes/overview#partner-nodes) provide access to the best closed source models such as Nano Banana, Seedance, Hunyuan3D, etc. - It is available on Windows, Linux, and macOS, locally with our [desktop application](https://www.comfy.org/download), our [portable install](#installing) or on our [cloud](https://www.comfy.org/cloud). - The most sophisticated workflows can be exposed through a simple UI thanks to App Mode. - It integrates seamlessly into production pipelines with our API endpoints. From 26d7f8556822d9d08c2d3e1878636ac3b4969af9 Mon Sep 17 00:00:00 2001 From: Christian Byrne Date: Tue, 11 Aug 2026 23:42:35 -0700 Subject: [PATCH 28/76] Fix PreviewAny escaping non-ASCII text in dict and list previews (#15513) --- comfy_extras/nodes_preview_any.py | 2 +- .../nodes_preview_any_test.py | 30 +++++++++++++++++++ 2 files changed, 31 insertions(+), 1 deletion(-) create mode 100644 tests-unit/comfy_extras_test/nodes_preview_any_test.py diff --git a/comfy_extras/nodes_preview_any.py b/comfy_extras/nodes_preview_any.py index d985f3287..18bf7c1cd 100644 --- a/comfy_extras/nodes_preview_any.py +++ b/comfy_extras/nodes_preview_any.py @@ -29,7 +29,7 @@ class PreviewAny(): value = str(source) elif source is not None: try: - value = json.dumps(source, indent=4) + value = json.dumps(source, indent=4, ensure_ascii=False) except Exception: try: value = str(source) diff --git a/tests-unit/comfy_extras_test/nodes_preview_any_test.py b/tests-unit/comfy_extras_test/nodes_preview_any_test.py new file mode 100644 index 000000000..563c1f1b0 --- /dev/null +++ b/tests-unit/comfy_extras_test/nodes_preview_any_test.py @@ -0,0 +1,30 @@ +from unittest.mock import patch, MagicMock + +mock_nodes = MagicMock() +mock_nodes.MAX_RESOLUTION = 16384 +mock_server = MagicMock() + +with patch.dict("sys.modules", {"nodes": mock_nodes, "server": mock_server}): + from comfy_extras.nodes_preview_any import PreviewAny + + +class TestPreviewAnyMain: + @staticmethod + def _exec(source) -> dict: + return PreviewAny().main(source) + + def test_dict_keeps_non_ascii(self): + result = self._exec({"greeting": "你好"}) + assert "你好" in result["ui"]["text"][0] + assert "\\u" not in result["ui"]["text"][0] + assert result["result"][0] == result["ui"]["text"][0] + + def test_list_keeps_non_ascii(self): + result = self._exec(["你好", "こんにちは"]) + assert "こんにちは" in result["result"][0] + assert "\\u" not in result["result"][0] + + def test_string_passthrough(self): + result = self._exec("你好") + assert result["ui"]["text"][0] == "你好" + assert result["result"][0] == "你好" From bd34f338ac505ea79e43968753968a464060e609 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 12 Aug 2026 00:55:08 -0700 Subject: [PATCH 29/76] Fix float64 device in ltx diffusion decoder. (#15516) --- comfy/ldm/lightricks/vae/na_diffusion_decoder.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/comfy/ldm/lightricks/vae/na_diffusion_decoder.py b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py index 8a172e101..ec539c665 100644 --- a/comfy/ldm/lightricks/vae/na_diffusion_decoder.py +++ b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py @@ -22,6 +22,7 @@ import torch import torch.nn.functional as F from einops import rearrange from torch import nn +import comfy.model_management from comfy.ldm.lightricks.model import get_timestep_embedding from .causal_video_autoencoder import Encoder, processor @@ -74,8 +75,12 @@ def default_rope_dim_split(head_dim): def rope_inv_freqs(dim, base=10000.0, device=None): + out_device = device + if not comfy.model_management.supports_fp64(device): + device = torch.device("cpu") + exponents = torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim - return (1.0 / torch.pow(torch.tensor(float(base), dtype=torch.float64, device=device), exponents)).to(torch.float32) + return (1.0 / torch.pow(torch.tensor(float(base), dtype=torch.float64, device=device), exponents)).to(dtype=torch.float32, device=out_device) def _rope_tables(lengths, inv_freqs, device): From 725e6ec60621c6f001af04769173e7dbb3c53541 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 12 Aug 2026 13:22:50 -0700 Subject: [PATCH 30/76] Support anima tunes with extra blocks. (#15555) --- comfy/model_detection.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index bc7b2b9f8..aec205290 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -830,11 +830,10 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["use_adaln_lora"] = True dit_config["adaln_lora_dim"] = 256 + dit_config["num_blocks"] = count_blocks(state_dict_keys, '{}blocks.'.format(key_prefix) + '{}.') if dit_config["model_channels"] == 2048: - dit_config["num_blocks"] = 28 dit_config["num_heads"] = 16 elif dit_config["model_channels"] == 5120: - dit_config["num_blocks"] = 36 dit_config["num_heads"] = 40 if dit_config["in_channels"] == 16: From 6b30dc206805829cef6353ba0c82a6a7c9b3944a Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 12 Aug 2026 18:39:00 -0700 Subject: [PATCH 31/76] Don't disable dynamic vram on WSL. (#15562) --- main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main.py b/main.py index 361b1fc89..eb64726ea 100644 --- a/main.py +++ b/main.py @@ -248,7 +248,7 @@ import hook_breaker_ac10a0 import comfy.memory_management import comfy.model_patcher -if args.enable_dynamic_vram or (enables_dynamic_vram() and comfy.model_management.is_nvidia() and not comfy.model_management.is_wsl()): +if args.enable_dynamic_vram or (enables_dynamic_vram() and comfy.model_management.is_nvidia()): if (not args.enable_dynamic_vram) and (comfy.model_management.torch_version_numeric < (2, 8)): logging.warning("Unsupported Pytorch detected. DynamicVRAM support requires Pytorch version 2.8 or later. Falling back to legacy ModelPatcher. VRAM estimates may be unreliable especially on Windows") else: From 2220d111c8b036f094eb465400fdf962626e4afa Mon Sep 17 00:00:00 2001 From: Alex Harper Date: Wed, 12 Aug 2026 21:40:11 -0400 Subject: [PATCH 32/76] Query pytorch for aotriton support instead of listing its lib directory (#15412) --- comfy/model_management.py | 36 ++++++++++++++++++++++-------------- 1 file changed, 22 insertions(+), 14 deletions(-) diff --git a/comfy/model_management.py b/comfy/model_management.py index 65599424b..15c03dc77 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -490,28 +490,36 @@ try: except: rocm_version = (6, -1) - def aotriton_supported(gpu_arch): - path = torch.__path__[0] - path = os.path.join(os.path.join(path, "lib"), "aotriton.images") - gfx = set(map(lambda a: a[4:], filter(lambda a: a.startswith("amd-gfx"), os.listdir(path)))) - if gpu_arch in gfx: - return True - if "{}x".format(gpu_arch[:-1]) in gfx: - return True - if "{}xx".format(gpu_arch[:-2]) in gfx: - return True - return False + def aotriton_supported(): + """Whether pytorch reports flash attention as usable on this gpu. + + can_use_flash_attention() evaluates runtime eligibility for the given + parameters; on a ROCm build that includes checking the gpu arch against the + kernel images AOTriton was compiled for. Querying it avoids assuming where + those images live inside the torch install. The probe tensor is shaped and + typed to pass the unrelated SDPA checks, so False means no hardware support + rather than a rejected shape. + """ + try: + if not torch.backends.cuda.is_flash_attention_available(): # not built with flash attention + return False + q = torch.empty((1, 1, 8, 64), dtype=torch.float16, device=get_torch_device()) + params = torch.backends.cuda.SDPAParams(q, q, q, None, 0.0, False, False) + return torch.backends.cuda.can_use_flash_attention(params, False) + except (AttributeError, RuntimeError, TypeError) as e: + logging.warning("Could not query aotriton support: {}".format(e)) + return False logging.info("AMD arch: {}".format(arch)) logging.info("ROCm version: {}".format(rocm_version)) if args.use_split_cross_attention == False and args.use_quad_cross_attention == False: - if aotriton_supported(arch): # AMD efficient attention implementation depends on aotriton. + if aotriton_supported(): # AMD efficient attention implementation depends on aotriton. if torch_version_numeric >= (2, 7): # works on 2.6 but doesn't actually seem to improve much if any((a in arch) for a in ["gfx90a", "gfx942", "gfx950", "gfx1100", "gfx1101", "gfx1150", "gfx1151"]): # TODO: more arches, TODO: gfx950 ENABLE_PYTORCH_ATTENTION = True if rocm_version >= (7, 0): - if any((a in arch) for a in ["gfx1200", "gfx1201"]): - ENABLE_PYTORCH_ATTENTION = True + if any((a in arch) for a in ["gfx1200", "gfx1201"]): + ENABLE_PYTORCH_ATTENTION = True if torch_version_numeric >= (2, 7) and rocm_version >= (6, 4): if any((a in arch) for a in ["gfx1200", "gfx1201", "gfx950"]): # TODO: more arches, "gfx942" gives error on pytorch nightly 2.10 1013 rocm7.0 SUPPORT_FP8_OPS = True From addd479729c7c266a1351fb88493aa5c1657f49b Mon Sep 17 00:00:00 2001 From: Alexander Brown Date: Wed, 12 Aug 2026 19:44:14 -0700 Subject: [PATCH 33/76] Fix Generate Text ignoring thinking=false on Gemma4 E2B/E4B (#15278) --- comfy/text_encoders/gemma4.py | 14 +++-- tests-unit/comfy_test/gemma4_template_test.py | 61 +++++++++++++++++++ 2 files changed, 69 insertions(+), 6 deletions(-) create mode 100644 tests-unit/comfy_test/gemma4_template_test.py diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index 5b0b968c9..606f8993e 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -1183,6 +1183,7 @@ def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, po class Gemma4_Tokenizer(): tokenizer_json_data = None + prime_empty_thought = False def state_dict(self): if self.tokenizer_json_data is not None: @@ -1333,8 +1334,8 @@ class Gemma4_Tokenizer(): num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1] n_audio_tokens = self._audio_token_count(num_samples) media += "<|audio>" + "<|audio|>" * n_audio_tokens + "" - # Non-thinking mode primes an empty thought channel so the model answers directly. - model_open = "" if thinking else "<|channel>thought\n" + # 12B/31B prime a closed thought block for non-thinking mode, E2B/E4B must not: it cues them into reasoning inline. + model_open = "<|channel>thought\n" if self.prime_empty_thought and not thinking else "" llama_text = f"{system}<|turn>user\n{text}{media}\n<|turn>model\n{model_open}" text_tokens = super().tokenize_with_weights(llama_text, return_word_ids) @@ -1418,6 +1419,7 @@ class Gemma4Tokenizer(sd1_clip.SD1Tokenizer): class Gemma4UnifiedSDTokenizer(Gemma4SDTokenizer): """Encoder-free (gemma4_unified) audio: raw 16kHz waveform frames instead of mel spectrogram.""" embedding_size = 3840 + prime_empty_thought = True def _extract_audio_features(self, waveform, sample_rate): audio = self._resample_16k(waveform, sample_rate) @@ -1500,7 +1502,7 @@ def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=No # Variants -def _make_variant(config_cls): +def _make_variant(config_cls, prime_empty_thought=False): audio = config_cls.audio_config is not None bases = (Gemma4AudioMixin, Gemma4Base) if audio else (Gemma4Base,) class Variant(*bases): @@ -1510,8 +1512,8 @@ def _make_variant(config_cls): if audio: self._init_audio(self.model.config, dtype, device, operations) embedding_size = config_cls.hidden_size - if embedding_size != Gemma4SDTokenizer.embedding_size: - tok_cls = type('T', (Gemma4SDTokenizer,), {'embedding_size': embedding_size}) + if embedding_size != Gemma4SDTokenizer.embedding_size or prime_empty_thought: + tok_cls = type('T', (Gemma4SDTokenizer,), {'embedding_size': embedding_size, 'prime_empty_thought': prime_empty_thought}) class Tokenizer(Gemma4Tokenizer): tokenizer_class = tok_cls Variant.tokenizer = Tokenizer @@ -1521,7 +1523,7 @@ def _make_variant(config_cls): Gemma4_E4B = _make_variant(Gemma4Config) Gemma4_E2B = _make_variant(Gemma4_E2B_Config) -Gemma4_31B = _make_variant(Gemma4_31B_Config) +Gemma4_31B = _make_variant(Gemma4_31B_Config, prime_empty_thought=True) # Gemma4 12B Unified: encoder-free multimodal, distinct base/tokenizer (not via _make_variant). diff --git a/tests-unit/comfy_test/gemma4_template_test.py b/tests-unit/comfy_test/gemma4_template_test.py new file mode 100644 index 000000000..77e274653 --- /dev/null +++ b/tests-unit/comfy_test/gemma4_template_test.py @@ -0,0 +1,61 @@ +"""Gemma4 chat template regression tests.""" + +import pytest +import torch + +from comfy.cli_args import args + +if not torch.cuda.is_available(): + args.cpu = True + +import comfy.text_encoders.gemma4 as gemma4 # noqa: E402 + +PROMPT = "describe a cute anime girl with fennec ears" +THOUGHT_BLOCK = "<|channel>thought\n" + +# E2B/E4B and 12B/31B ship different canonical chat templates: only the latter prime a +# closed thought block when thinking is off. +NO_PRIMING = [gemma4.Gemma4_E2B, gemma4.Gemma4_E4B] +PRIMING = [gemma4.Gemma4_31B, gemma4.Gemma4_12B] + + +class _CaptureTemplate: + """Stands in for SDTokenizer.tokenize_with_weights so the built template is checked without model files.""" + llama_text = "" + + def tokenize_with_weights(self, text, return_word_ids=False, **kwargs): + self.llama_text = text + return {} + + +def build_template(variant, **kwargs): + prime = variant.tokenizer.tokenizer_class.prime_empty_thought + probe = type("Probe", (gemma4.Gemma4_Tokenizer, _CaptureTemplate), {"prime_empty_thought": prime})() + probe.tokenize_with_weights(PROMPT, **kwargs) + return probe.llama_text + + +@pytest.mark.parametrize("variant", NO_PRIMING + PRIMING) +def test_thinking_enabled_only_asks_via_the_system_turn(variant): + template = build_template(variant, skip_template=False, thinking=True) + assert template == f"<|turn>system\n<|think|>\n\n<|turn>user\n{PROMPT}\n<|turn>model\n" + + +@pytest.mark.parametrize("variant", NO_PRIMING) +def test_thinking_disabled_does_not_prime_a_thought_channel(variant): + template = build_template(variant, skip_template=False, thinking=False) + assert template == f"<|turn>user\n{PROMPT}\n<|turn>model\n" + assert "channel" not in template + assert "<|think|>" not in template + + +@pytest.mark.parametrize("variant", PRIMING) +def test_thinking_disabled_primes_a_thought_channel(variant): + template = build_template(variant, skip_template=False, thinking=False) + assert template == f"<|turn>user\n{PROMPT}\n<|turn>model\n{THOUGHT_BLOCK}" + + +@pytest.mark.parametrize("variant", NO_PRIMING + PRIMING) +@pytest.mark.parametrize("thinking", [False, True]) +def test_skip_template_passes_text_through_unchanged(variant, thinking): + assert build_template(variant, skip_template=True, thinking=thinking) == PROMPT From b323a345bbbfb2f3a95b5b73b68eb7919a26515e Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 12 Aug 2026 20:01:31 -0700 Subject: [PATCH 34/76] Update comfy-kitchen package version to 0.2.31 (#15564) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index f461e3b76..4f505cc9c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.30 +comfy-kitchen==0.2.31 comfy-aimdo==0.4.13 requests simpleeval>=1.0.0 From 12666983cba9b43254ed993c2894dc727ca8ecfd Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Thu, 13 Aug 2026 17:46:08 +0300 Subject: [PATCH 35/76] [Partner Nodes] feat(MiniMax): add ContextIR and Regenerate nodes (#15471) Signed-off-by: bigcat88 --- comfy_api_nodes/apis/minimax.py | 21 +- comfy_api_nodes/nodes_minimax.py | 599 ++++++++++++++++++++++++++++++- 2 files changed, 616 insertions(+), 4 deletions(-) diff --git a/comfy_api_nodes/apis/minimax.py b/comfy_api_nodes/apis/minimax.py index bac4572d4..12a0853ac 100644 --- a/comfy_api_nodes/apis/minimax.py +++ b/comfy_api_nodes/apis/minimax.py @@ -161,12 +161,30 @@ class Hailuo03TaskCreationRequest(BaseModel): ..., min_length=1 ) resolution: str = Field(...) - duration: int = Field(..., ge=5, le=15) + duration: int = Field(..., ge=4, le=15) ratio: str | None = Field(None) seed: int | None = Field(None, ge=0, le=4294967295) aigc_watermark: bool | None = Field(None) +class Hailuo03ContextIRRequest(BaseModel): + model: str = Field(...) + content: list[Hailuo03TextContent | Hailuo03ImageContent | Hailuo03VideoContent | Hailuo03AudioContent] = Field( + ..., min_length=1 + ) + duration: int = Field(..., ge=4, le=15) + ratio: str | None = Field(None) + + +class Hailuo03RegenerationRequest(BaseModel): + model: str = Field(...) + content: list[Hailuo03TextContent | Hailuo03ImageContent | Hailuo03VideoContent | Hailuo03AudioContent] = Field( + ..., min_length=1 + ) + resolution: str = Field(...) + aigc_watermark: bool | None = Field(None) + + class Hailuo03TaskCreationResponse(BaseModel): task_id: str = Field(...) @@ -178,6 +196,7 @@ class Hailuo03TaskError(BaseModel): class Hailuo03TaskContent(BaseModel): url: str | None = Field(None) + prompt: str | None = Field(None) class Hailuo03TaskUsage(BaseModel): diff --git a/comfy_api_nodes/nodes_minimax.py b/comfy_api_nodes/nodes_minimax.py index 3c1d29257..de3895221 100644 --- a/comfy_api_nodes/nodes_minimax.py +++ b/comfy_api_nodes/nodes_minimax.py @@ -3,12 +3,14 @@ from typing import Optional import torch from typing_extensions import override -from comfy_api.latest import IO, ComfyExtension +from comfy_api.latest import IO, ComfyExtension, Input from comfy_api_nodes.apis.minimax import ( Hailuo03AudioContent, Hailuo03AudioContentUrl, + Hailuo03ContextIRRequest, Hailuo03ImageContent, Hailuo03ImageContentUrl, + Hailuo03RegenerationRequest, Hailuo03TaskCreationRequest, Hailuo03TaskCreationResponse, Hailuo03TaskQueryResponse, @@ -456,6 +458,9 @@ HAILUO_03_QUERY_ENDPOINT = "/proxy/minimax/v2/query/video_generation" # + /{tas HAILUO_03_MODELS = {"MiniMax H3": "MiniMax-H3"} HAILUO_03_FAILED_STATUSES = ["failed", "cancelled", "expired"] +HAILUO_03_CONTEXT_IR_ENDPOINT = "/proxy/minimax/v2/h3_context_ir" +HAILUO_03_REGENERATION_ENDPOINT = "/proxy/minimax/v2/video_regeneration" + def _hailuo03_model_inputs(include_ratio: bool = True, allow_adaptive: bool = True): inputs = [ @@ -487,10 +492,10 @@ def _hailuo03_model_inputs(include_ratio: bool = True, allow_adaptive: bool = Tr IO.Int.Input( "duration", default=5, - min=5, + min=4, max=15, step=1, - tooltip="Duration of the output video in seconds (5-15).", + tooltip="Duration of the output video in seconds (4-15).", display_mode=IO.NumberDisplay.slider, ) ) @@ -939,6 +944,592 @@ class MinimaxHailuo03ReferenceNode(IO.ComfyNode): ) +class MinimaxHailuo03ContextIRNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="MinimaxHailuo03ContextIRNode", + display_name="MiniMax H3 Context IR (Prompt Enhancer)", + category="partner/video/MiniMax", + description="Analyze text and media context with MiniMax H3 Context IR and produce an enhanced, " + "structured video prompt. Feed the output into the prompt of a MiniMax H3 video node and attach " + "the same media there in the same order, because the enhanced prompt refers to the attached " + "media by position.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "MiniMax H3", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Description of the video you intend to generate.", + ), + IO.Int.Input( + "duration", + default=5, + min=4, + max=15, + step=1, + tooltip="Duration of the video you intend to generate, in seconds (4-15).", + display_mode=IO.NumberDisplay.slider, + ), + IO.Combo.Input( + "ratio", + options=["adaptive", "16:9", "4:3", "1:1", "3:4", "9:16", "21:9"], + default="adaptive", + tooltip="Aspect ratio of the video you intend to generate. 'adaptive' " + "requires at least one image, video, or audio input.", + ), + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("reference_image"), + names=[ + "image_1", + "image_2", + "image_3", + "image_4", + "image_5", + "image_6", + "image_7", + "image_8", + "image_9", + ], + min=0, + ), + tooltip="Subject or style reference images, referred to in the prompt " + "as 'Image 1'..'Image 9' in connection order. Up to 9 images.", + ), + IO.Autogrow.Input( + "reference_videos", + template=IO.Autogrow.TemplateNames( + IO.Video.Input("reference_video"), + names=["video_1", "video_2", "video_3"], + min=0, + ), + tooltip="Motion or scene reference videos, referred to in the prompt " + "as 'Video 1'..'Video 3' in connection order. Up to 3 videos, " + "2-15 seconds each, 15 seconds in total.", + ), + IO.Autogrow.Input( + "reference_audios", + template=IO.Autogrow.TemplateNames( + IO.Audio.Input("reference_audio"), + names=["audio_1", "audio_2", "audio_3"], + min=0, + ), + tooltip="Audio references, referred to in the prompt as " + "'Audio 1'..'Audio 3' in connection order. Up to 3 clips, " + "2-15 seconds each, 15 seconds in total. Cannot be used without " + "a reference image or video.", + ), + ], + ) + ], + tooltip="Model to use for prompt enhancement.", + ), + IO.Image.Input( + "first_frame", + tooltip="First frame of the video you intend to generate. Cannot be combined with " + "reference media.", + optional=True, + ), + IO.Image.Input( + "last_frame", + tooltip="Last frame of the video you intend to generate. Cannot be combined with " + "reference media.", + optional=True, + ), + ], + outputs=[ + IO.String.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends( + inputs=["first_frame", "last_frame"], + input_groups=["model.reference_images", "model.reference_videos", "model.reference_audios"], + ), + expr=""" + ( + $imgsRaw := $lookup(inputGroups, "model.reference_images"); + $imgs := $imgsRaw ? $imgsRaw : 0; + $vidsRaw := $lookup(inputGroups, "model.reference_videos"); + $vids := $vidsRaw ? $vidsRaw : 0; + $audsRaw := $lookup(inputGroups, "model.reference_audios"); + $auds := $audsRaw ? $audsRaw : 0; + $frames := (inputs.first_frame.connected ? 1 : 0) + (inputs.last_frame.connected ? 1 : 0); + ($imgs + $vids + $auds) > 0 + ? {"type": "range_usd", "min_usd": 0.06, "max_usd": 0.11, "format": {"approximate": true}} + : $frames > 0 + ? {"type": "usd", "usd": 0.05, "format": {"approximate": true}} + : {"type": "usd", "usd": 0.02, "format": {"approximate": true}} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + first_frame: torch.Tensor | None = None, + last_frame: torch.Tensor | None = None, + ) -> IO.NodeOutput: + validate_string(model["prompt"], strip_whitespace=True, min_length=1) + + reference_images = {k: v for k, v in (model.get("reference_images") or {}).items() if v is not None} + reference_videos = {k: v for k, v in (model.get("reference_videos") or {}).items() if v is not None} + reference_audios = {k: v for k, v in (model.get("reference_audios") or {}).items() if v is not None} + has_frames = first_frame is not None or last_frame is not None + has_references = bool(reference_images) or bool(reference_videos) or bool(reference_audios) + if has_frames and has_references: + raise ValueError( + "First/last frame and reference media are mutually exclusive. Use frames for an " + "image-to-video prompt, or reference media for a reference-to-video prompt." + ) + if reference_audios and not reference_images and not reference_videos: + raise ValueError("Reference audio cannot be used without a reference image or video.") + if not has_frames and not has_references and model["ratio"] == "adaptive": + raise ValueError( + "Ratio 'adaptive' is not supported for text-only requests; select an explicit aspect ratio." + ) + + for frame in (first_frame, last_frame): + if frame is not None: + validate_image_aspect_ratio(frame, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(frame, min_width=256, min_height=256) + for image in reference_images.values(): + validate_image_aspect_ratio(image, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(image, min_width=256, min_height=256) + + total_video_duration = 0.0 + for i, video in enumerate(reference_videos.values(), 1): + try: + fps = float(video.get_frame_rate()) + except Exception: + fps = 0.0 + if fps and not (23.9 <= fps <= 60.5): + raise ValueError(f"Reference video {i} is {fps:.2f} FPS. Supported range is 23.976-60 FPS.") + try: + dur = video.get_duration() + except Exception: + continue + if dur < 1.8: + raise ValueError(f"Reference video {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_video_duration += dur + if total_video_duration > 15.1: + raise ValueError( + f"Total reference video duration is {total_video_duration:.1f}s. Maximum is 15 seconds." + ) + + total_audio_duration = 0.0 + for i, audio in enumerate(reference_audios.values(), 1): + dur = int(audio["waveform"].shape[-1]) / int(audio["sample_rate"]) + if dur < 1.8: + raise ValueError(f"Reference audio {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_audio_duration += dur + if total_audio_duration > 15.1: + raise ValueError( + f"Total reference audio duration is {total_audio_duration:.1f}s. Maximum is 15 seconds." + ) + + content: list = [Hailuo03TextContent(text=model["prompt"])] + if first_frame is not None: + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, first_frame, max_images=1, wait_label="Uploading first frame" + ) + )[0], + ), + role="first_frame", + ) + ) + if last_frame is not None: + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, last_frame, max_images=1, wait_label="Uploading last frame" + ) + )[0], + ), + role="last_frame", + ) + ) + for i, image in enumerate(reference_images.values(), 1): + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, image, max_images=1, wait_label=f"Uploading image {i}" + ) + )[0], + ), + role="reference_image", + ) + ) + for i, video in enumerate(reference_videos.values(), 1): + content.append( + Hailuo03VideoContent( + video_url=Hailuo03VideoContentUrl( + url=await upload_video_to_comfyapi(cls, video, wait_label=f"Uploading video {i}"), + ), + ) + ) + for audio in reference_audios.values(): + content.append( + Hailuo03AudioContent( + audio_url=Hailuo03AudioContentUrl( + url=await upload_audio_to_comfyapi( + cls, + audio, + container_format="mp3", + codec_name="libmp3lame", + mime_type="audio/mpeg", + ), + ), + ) + ) + + response = await sync_op( + cls, + ApiEndpoint(path=HAILUO_03_CONTEXT_IR_ENDPOINT, method="POST"), + response_model=Hailuo03TaskCreationResponse, + data=Hailuo03ContextIRRequest( + model=HAILUO_03_MODELS[model["model"]], + content=content, + duration=model["duration"], + ratio=None if model["ratio"] == "adaptive" else model["ratio"], + ), + ) + task_result = await poll_op( + cls, + ApiEndpoint(path=f"{HAILUO_03_QUERY_ENDPOINT}/{response.task_id}"), + response_model=Hailuo03TaskQueryResponse, + status_extractor=lambda r: r.task.status, + failed_statuses=HAILUO_03_FAILED_STATUSES, + poll_interval=5, + ) + prompt = task_result.task.content.prompt if task_result.task.content else None + if not prompt: + raise Exception(f"No enhanced prompt in the response: {task_result.model_dump()}") + return IO.NodeOutput(prompt) + + +class MinimaxHailuo03RegenerateNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="MinimaxHailuo03RegenerateNode", + display_name="MiniMax H3 Regenerate to 2K", + category="partner/video/MiniMax", + description="Re-render a MiniMax H3 768P output at 2K resolution. Connect the unmodified 768P " + "video and the exact prompt used to generate it; if the original generation used first/last " + "frames or reference media, attach the same inputs.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "MiniMax H3", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="The exact prompt used to generate the source video.", + ), + IO.Combo.Input( + "resolution", + options=["2K"], + tooltip="Resolution to re-render the source video at.", + ), + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("reference_image"), + names=[ + "image_1", + "image_2", + "image_3", + "image_4", + "image_5", + "image_6", + "image_7", + "image_8", + "image_9", + ], + min=0, + ), + tooltip="Reference images from the original generation, in the same " + "order. Up to 9 images.", + ), + IO.Autogrow.Input( + "reference_videos", + template=IO.Autogrow.TemplateNames( + IO.Video.Input("reference_video"), + names=["video_1", "video_2", "video_3"], + min=0, + ), + tooltip="Reference videos from the original generation, in the same " + "order. Up to 3 videos, 2-15 seconds each, 15 seconds in total.", + ), + IO.Autogrow.Input( + "reference_audios", + template=IO.Autogrow.TemplateNames( + IO.Audio.Input("reference_audio"), + names=["audio_1", "audio_2", "audio_3"], + min=0, + ), + tooltip="Audio references from the original generation, in the same " + "order. Up to 3 clips, 2-15 seconds each, 15 seconds in total. " + "Cannot be used without a reference image or video.", + ), + ], + ) + ], + tooltip="Model to use for video regeneration.", + ), + IO.Video.Input( + "video", + tooltip="The MiniMax H3 768P output video to re-render. Connect the unmodified output " + "of a MiniMax H3 video node (24 FPS, 4-15 seconds). 2K outputs cannot be used.", + ), + IO.Image.Input( + "first_frame", + tooltip="First frame image from the original generation, if one was used.", + optional=True, + ), + IO.Image.Input( + "last_frame", + tooltip="Last frame image from the original generation, if one was used.", + optional=True, + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AIGC watermark to the video.", + advanced=True, + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type": "usd", "usd": 0.0715, "format": {"suffix": "/second"}}""", + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + video: Input.Video, + watermark: bool, + first_frame: torch.Tensor | None = None, + last_frame: torch.Tensor | None = None, + ) -> IO.NodeOutput: + validate_string(model["prompt"], strip_whitespace=True, min_length=1) + + try: + fps = float(video.get_frame_rate()) + except Exception: + fps = 0.0 + if fps and not (23.9 <= fps <= 24.1): + raise ValueError( + f"The source video is {fps:.2f} FPS. Regeneration accepts unmodified MiniMax H3 768P " + "outputs, which are 24 FPS." + ) + try: + width, height = video.get_dimensions() + except Exception: + width = height = 0 + if width and height and (width % 32 or height % 32 or width * height > 1_032_192): + raise ValueError( + f"The source video is {width}x{height}. Regeneration accepts MiniMax H3 768P outputs " + "(width and height divisible by 32, at most 1,032,192 total pixels); 2K outputs cannot " + "be used as a source." + ) + try: + frame_count = video.get_frame_count() + except Exception: + frame_count = 0 + if frame_count and (frame_count < 107 or frame_count > 362 or (frame_count - 107) % 17): + raise ValueError( + f"The source video has {frame_count} frames. Regeneration accepts unmodified " + "MiniMax H3 outputs, whose length is 107 to 362 frames in steps of 17 " + "(4 to 15 seconds at 24 FPS)." + ) + + reference_images = {k: v for k, v in (model.get("reference_images") or {}).items() if v is not None} + reference_videos = {k: v for k, v in (model.get("reference_videos") or {}).items() if v is not None} + reference_audios = {k: v for k, v in (model.get("reference_audios") or {}).items() if v is not None} + if (first_frame is not None or last_frame is not None) and ( + reference_images or reference_videos or reference_audios + ): + raise ValueError( + "First/last frame and reference media are mutually exclusive. Use frames for an " + "image-to-video prompt, or reference media for a reference-to-video prompt." + ) + if reference_audios and not reference_images and not reference_videos: + raise ValueError("Reference audio cannot be used without a reference image or video.") + + for frame in (first_frame, last_frame): + if frame is not None: + validate_image_aspect_ratio(frame, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(frame, min_width=256, min_height=256) + for image in reference_images.values(): + validate_image_aspect_ratio(image, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(image, min_width=256, min_height=256) + + total_video_duration = 0.0 + for i, ref_video in enumerate(reference_videos.values(), 1): + try: + ref_fps = float(ref_video.get_frame_rate()) + except Exception: + ref_fps = 0.0 + if ref_fps and not (23.9 <= ref_fps <= 60.5): + raise ValueError(f"Reference video {i} is {ref_fps:.2f} FPS. Supported range is 23.976-60 FPS.") + try: + dur = ref_video.get_duration() + except Exception: + continue + if dur < 1.8: + raise ValueError(f"Reference video {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_video_duration += dur + if total_video_duration > 15.1: + raise ValueError( + f"Total reference video duration is {total_video_duration:.1f}s. Maximum is 15 seconds." + ) + + total_audio_duration = 0.0 + for i, audio in enumerate(reference_audios.values(), 1): + dur = int(audio["waveform"].shape[-1]) / int(audio["sample_rate"]) + if dur < 1.8: + raise ValueError(f"Reference audio {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_audio_duration += dur + if total_audio_duration > 15.1: + raise ValueError( + f"Total reference audio duration is {total_audio_duration:.1f}s. Maximum is 15 seconds." + ) + + content: list = [ + Hailuo03VideoContent( + video_url=Hailuo03VideoContentUrl( + url=await upload_video_to_comfyapi(cls, video, wait_label="Uploading source video"), + ), + role="base_video", + ), + Hailuo03TextContent(text=model["prompt"]), + ] + if first_frame is not None: + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, first_frame, max_images=1, wait_label="Uploading first frame" + ) + )[0], + ), + role="first_frame", + ) + ) + if last_frame is not None: + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, last_frame, max_images=1, wait_label="Uploading last frame" + ) + )[0], + ), + role="last_frame", + ) + ) + for i, image in enumerate(reference_images.values(), 1): + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, image, max_images=1, wait_label=f"Uploading image {i}" + ) + )[0], + ), + role="reference_image", + ) + ) + for i, ref_video in enumerate(reference_videos.values(), 1): + content.append( + Hailuo03VideoContent( + video_url=Hailuo03VideoContentUrl( + url=await upload_video_to_comfyapi(cls, ref_video, wait_label=f"Uploading video {i}"), + ), + ) + ) + for audio in reference_audios.values(): + content.append( + Hailuo03AudioContent( + audio_url=Hailuo03AudioContentUrl( + url=await upload_audio_to_comfyapi( + cls, + audio, + container_format="mp3", + codec_name="libmp3lame", + mime_type="audio/mpeg", + ), + ), + ) + ) + + response = await sync_op( + cls, + ApiEndpoint(path=HAILUO_03_REGENERATION_ENDPOINT, method="POST"), + response_model=Hailuo03TaskCreationResponse, + data=Hailuo03RegenerationRequest( + model=HAILUO_03_MODELS[model["model"]], + content=content, + resolution=model["resolution"], + aigc_watermark=watermark, + ), + ) + task_result = await poll_op( + cls, + ApiEndpoint(path=f"{HAILUO_03_QUERY_ENDPOINT}/{response.task_id}"), + response_model=Hailuo03TaskQueryResponse, + status_extractor=lambda r: r.task.status, + failed_statuses=HAILUO_03_FAILED_STATUSES, + poll_interval=10, + ) + video_url = task_result.task.content.url if task_result.task.content else None + if not video_url: + raise Exception(f"No video URL in the response: {task_result.model_dump()}") + return IO.NodeOutput(await download_url_to_video_output(video_url)) + + class MinimaxExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: @@ -950,6 +1541,8 @@ class MinimaxExtension(ComfyExtension): MinimaxHailuo03TextToVideoNode, MinimaxHailuo03FirstLastFrameNode, MinimaxHailuo03ReferenceNode, + MinimaxHailuo03ContextIRNode, + MinimaxHailuo03RegenerateNode, ] From efd4e951a00e85bd92e79f1d685427912b0dad5e Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Fri, 14 Aug 2026 02:10:08 +1000 Subject: [PATCH 36/76] Implement Minimax Music 3 + Core Support for Cuda Graphs (#15570) --- comfy/cli_args.py | 1 + comfy/latent_formats.py | 5 + comfy/ldm/minimax_music/__init__.py | 0 comfy/ldm/minimax_music/ar.py | 337 +++++++++++++++++++++++++++ comfy/ldm/minimax_music/dav.py | 137 +++++++++++ comfy/ldm/minimax_music/dit.py | 213 +++++++++++++++++ comfy/ldm/minimax_music/prompt.py | 70 ++++++ comfy/model_base.py | 13 ++ comfy/model_detection.py | 7 + comfy/model_management.py | 9 + comfy/model_patcher.py | 30 ++- comfy/model_prefetch.py | 96 +++++++- comfy/ops.py | 14 +- comfy/sd.py | 30 ++- comfy/supported_models.py | 21 ++ comfy/text_encoders/llama.py | 188 ++++++++++++--- comfy/text_encoders/minimax_music.py | 129 ++++++++++ comfy_extras/nodes_minimax_music.py | 77 ++++++ nodes.py | 6 +- 19 files changed, 1333 insertions(+), 50 deletions(-) create mode 100644 comfy/ldm/minimax_music/__init__.py create mode 100644 comfy/ldm/minimax_music/ar.py create mode 100644 comfy/ldm/minimax_music/dav.py create mode 100644 comfy/ldm/minimax_music/dit.py create mode 100644 comfy/ldm/minimax_music/prompt.py create mode 100644 comfy/text_encoders/minimax_music.py create mode 100644 comfy_extras/nodes_minimax_music.py diff --git a/comfy/cli_args.py b/comfy/cli_args.py index 9de244087..c6660846d 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -180,6 +180,7 @@ parser.add_argument("--disable-async-offload", action="store_true", help="Disabl parser.add_argument("--disable-dynamic-vram", action="store_true", help="Disable dynamic VRAM and use estimate based model loading.") parser.add_argument("--enable-dynamic-vram", action="store_true", help="Enable dynamic VRAM on systems where it's not enabled by default.") parser.add_argument("--fast-disk", action="store_true", help="Prefer disk-backed dynamic loading and offload over unpinned RAM. Can be faster for users with fast NVME disks.") +parser.add_argument("--disable-cuda-graphs", action="store_true", help="Disable CUDA graphs.") parser.add_argument("--force-non-blocking", action="store_true", help="Force ComfyUI to use non-blocking operations for all applicable tensors. This may improve performance on some non-Nvidia systems but can cause issues with some workflows.") diff --git a/comfy/latent_formats.py b/comfy/latent_formats.py index c4270022b..dc737fc7d 100644 --- a/comfy/latent_formats.py +++ b/comfy/latent_formats.py @@ -957,6 +957,11 @@ class ACEAudio15(LatentFormat): latent_dimensions = 1 temporal_downscale_ratio = 1764 +class MiniMaxMusic3(LatentFormat): + latent_channels = 128 + latent_dimensions = 1 + temporal_downscale_ratio = 512 + class ChromaRadiance(LatentFormat): latent_channels = 3 spacial_downscale_ratio = 1 diff --git a/comfy/ldm/minimax_music/__init__.py b/comfy/ldm/minimax_music/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/comfy/ldm/minimax_music/ar.py b/comfy/ldm/minimax_music/ar.py new file mode 100644 index 000000000..28215cb58 --- /dev/null +++ b/comfy/ldm/minimax_music/ar.py @@ -0,0 +1,337 @@ +import dataclasses +import hashlib + +import torch +from torch import nn + +import comfy.model_management +import comfy.model_prefetch +import comfy.ops +import comfy.utils +from comfy.ldm.modules.attention import optimized_attention_for_device +from comfy.text_encoders.llama import Llama2_, Qwen3_8BConfig + +from .prompt import AUDIO_CODE_OFFSET, SPECIAL_TOKEN_IDS + + +CFG_SCALE = 1.5 +CFG_TOP_K = 50 +C0_VOCAB_SIZE = 16384 +MAX_PROMPT_TOKENS = 5000 +MAX_AUDIO_FRAMES = 9000 +AUDIO_FRAMES_PER_SECOND = 25 + + +def derive_seed(seed, *parts): + digest = hashlib.blake2b(digest_size=8, person=b"minimax-ttm") + digest.update(int(seed).to_bytes(8, "little", signed=False)) + for part in parts: + value = str(part).encode("utf-8") + digest.update(len(value).to_bytes(4, "little")) + digest.update(value) + return int.from_bytes(digest.digest(), "little") & ((1 << 63) - 1) + + +def sample_topk(logits, top_k, generator): + values = torch.nan_to_num(logits.float(), nan=-1e9, posinf=1e9, neginf=-1e9) + top_k = min(top_k, values.shape[-1]) + threshold = torch.topk(values, top_k, dim=-1).values[..., -1, None] + values = values.masked_fill(values < threshold, -float("inf")) + probabilities = torch.nan_to_num(torch.softmax(values, dim=-1), nan=0.0) + probabilities = probabilities / probabilities.sum(dim=-1, keepdim=True).clamp_min(1e-12) + return torch.multinomial(probabilities, 1, generator=generator).squeeze(-1) + + +class RVQAttention(nn.Module): + def __init__(self, hidden_size, num_heads, dtype, device, operations): + super().__init__() + self.num_heads = num_heads + self.head_dim = hidden_size // num_heads + self.merged_qkv = None + self.qkv_proj = operations.Linear(hidden_size, hidden_size * 3, bias=False, dtype=dtype, device=device) + self.q_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.k_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.v_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.o_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + + def forward(self, x): + batch, length, hidden_size = x.shape + if self.merged_qkv: + q, k, v = self.qkv_proj(x).chunk(3, dim=-1) + else: + q = self.q_proj(x) + k = self.k_proj(x) + v = self.v_proj(x) + q = q.reshape(batch, length, self.num_heads, self.head_dim).transpose(1, 2) + k = k.reshape(batch, length, self.num_heads, self.head_dim).transpose(1, 2) + v = v.reshape(batch, length, self.num_heads, self.head_dim).transpose(1, 2) + mask = torch.full((length, length), torch.finfo(q.dtype).min, device=q.device, dtype=q.dtype).triu_(1) + attention = optimized_attention_for_device(q.device, mask=True, small_input=True) + out = attention(q, k, v, self.num_heads, mask=mask, skip_reshape=True) + return self.o_proj(out) + + +class RVQRMSNorm(nn.Module): + def __init__(self, hidden_size, dtype, device): + super().__init__() + self.weight = nn.Parameter(torch.empty(hidden_size, dtype=dtype, device=device)) + + def forward(self, x): + return torch.nn.functional.rms_norm(x, (x.shape[-1],), comfy.ops.cast_to_input(self.weight, x), 1e-6) + + +class RVQMLP(nn.Module): + def __init__(self, hidden_size, intermediate_size, dtype, device, operations): + super().__init__() + self.merged_mlp = None + self.gate_up_proj = operations.Linear(hidden_size, intermediate_size * 2, bias=False, dtype=dtype, device=device) + self.gate_proj = operations.Linear(hidden_size, intermediate_size, bias=False, dtype=dtype, device=device) + self.up_proj = operations.Linear(hidden_size, intermediate_size, bias=False, dtype=dtype, device=device) + self.down_proj = operations.Linear(intermediate_size, hidden_size, bias=False, dtype=dtype, device=device) + + def forward(self, x): + if self.merged_mlp: + return comfy.ops.linear_input_act(self.down_proj, self.gate_up_proj(x), "swiglu") + return self.down_proj(torch.nn.functional.silu(self.gate_proj(x)) * self.up_proj(x)) + + +class RVQDecoderBlock(nn.Module): + def __init__(self, hidden_size, num_heads, intermediate_size, dtype, device, operations): + super().__init__() + self.input_layernorm = RVQRMSNorm(hidden_size, dtype, device) + self.self_attn = RVQAttention(hidden_size, num_heads, dtype, device, operations) + self.post_attention_layernorm = RVQRMSNorm(hidden_size, dtype, device) + self.mlp = RVQMLP(hidden_size, intermediate_size, dtype, device, operations) + + def forward(self, x): + x = x + self.self_attn(self.input_layernorm(x)) + return x + self.mlp(self.post_attention_layernorm(x)) + + +class RVQDepthDecoder(nn.Module): + def __init__(self, config, dtype, device, operations): + super().__init__() + hidden_size = int(config["hidden_size"]) + audio_vocab_size = int(config["audio_vocab_size"]) + num_codebooks = int(config["audio_num_codebooks"]) + self.projection = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.pos_embedding = operations.Embedding(16, hidden_size, dtype=dtype, device=device) + self.audio_heads = nn.ModuleList([ + operations.Linear(hidden_size, audio_vocab_size, bias=False, dtype=dtype, device=device) + for _ in range(num_codebooks - 1) + ]) + self.layers = nn.ModuleList([ + RVQDecoderBlock( + hidden_size, + int(config["decoder_num_heads"]), + int(config["decoder_intermediate_size"]), + dtype, + device, + operations, + ) + for _ in range(int(config["decoder_num_layers"])) + ]) + self.norm = RVQRMSNorm(hidden_size, dtype, device) + + def forward(self, sequence): + positions = torch.arange(sequence.shape[1], device=sequence.device) + x = sequence + self.pos_embedding(positions, out_dtype=sequence.dtype).unsqueeze(0) + for layer in self.layers: + x = layer(x) + return self.norm(x) + + +class MiniMaxMusic3AR(nn.Module): + def __init__(self, config, dtype, device, operations): + super().__init__() + config_fields = {field.name for field in dataclasses.fields(Qwen3_8BConfig)} + qwen_config = Qwen3_8BConfig(**{key: value for key, value in config.items() if key in config_fields}) + qwen_config.lm_head = False + qwen_config.fixed_kv = True + qwen_config.merged_qkv = None + qwen_config.merged_mlp = None + self.model = Llama2_(qwen_config, device=device, dtype=dtype, ops=operations) + self.model.prefetch_dynamic_vbars = True + self.model.graph_dynamic_vbar_blocks = True + self.model.lm_head = operations.Linear(qwen_config.hidden_size, qwen_config.vocab_size, bias=False, dtype=dtype, device=device) + self.model.lm_head_pruned = operations.Linear(qwen_config.hidden_size, C0_VOCAB_SIZE + 1, bias=False, dtype=dtype, device=device) + self.model.embed_tokens_prefill = operations.Embedding(AUDIO_CODE_OFFSET, qwen_config.hidden_size, dtype=dtype, device=device) + self.model.embed_tokens_audio = operations.Embedding(C0_VOCAB_SIZE, qwen_config.hidden_size, dtype=dtype, device=device) + self.model.pruned_lm_head = None + self.model.pruned_embedding = None + self.model.audio_extra_embedding = operations.Embedding( + int(config["audio_vocab_size"]) * (int(config["audio_num_codebooks"]) - 1), + qwen_config.hidden_size, + dtype=dtype, + device=device, + ) + self.model.audio_decoder = RVQDepthDecoder(config, dtype, device, operations) + self.audio_vocab_size = int(config["audio_vocab_size"]) + self.num_codebooks = int(config["audio_num_codebooks"]) + self.embedding_scale = self.num_codebooks ** -0.5 + + def _guided_c0(self, logits, cfg_scale, top_k): + conditioned = logits[0:1].float() + unconditioned = logits[1:2].float() + guided = unconditioned + (conditioned - unconditioned) * cfg_scale + threshold = torch.topk(conditioned, top_k, dim=-1).values[..., -1, None] + return guided.masked_fill(conditioned < threshold, -float("inf")) + + def _depth_codes(self, hidden, c0, c0_embed, generator, execution_dtype, cfg_scale, top_k): + decoder = self.model.audio_decoder + sequence = [decoder.projection(hidden).unsqueeze(1)] + sequence.append(decoder.projection(c0_embed).unsqueeze(1)) + codes = [c0] + hidden_parts = [] + for index in range(1, self.num_codebooks): + out = decoder(torch.cat(sequence, dim=1))[:, -1] + hidden_parts.append(out[:1].detach()) + logits = decoder.audio_heads[index - 1](out) + conditioned = logits[:1].float() + unconditioned = logits[1:2].float() + code = sample_topk(unconditioned + (conditioned - unconditioned) * cfg_scale, top_k, generator).repeat(2) + codes.append(code) + if index < self.num_codebooks - 1: + embedding = self.model.audio_extra_embedding( + code + (index - 1) * self.audio_vocab_size, + out_dtype=execution_dtype, + ) + sequence.append(decoder.projection(embedding).unsqueeze(1)) + return torch.stack(codes, dim=1), torch.cat(hidden_parts, dim=-1) + + def _embed_c0(self, codes, execution_dtype): + if self.model.pruned_embedding: + return self.model.embed_tokens_audio(codes, out_dtype=execution_dtype) + return self.model.embed_tokens(codes + AUDIO_CODE_OFFSET, out_dtype=execution_dtype) + + def _embed_audio_frame(self, codes, execution_dtype): + c0 = self._embed_c0(codes[:, 0], execution_dtype) + offsets = torch.arange(self.num_codebooks - 1, device=codes.device) * self.audio_vocab_size + extra = self.model.audio_extra_embedding(codes[:, 1:] + offsets.unsqueeze(0), out_dtype=execution_dtype).sum(dim=1) + return ((c0 + extra) * self.embedding_scale).unsqueeze(1) + + def _sample_c0(self, hidden, cfg_scale, top_k, generator, vocab_mask): + if self.model.pruned_lm_head: + guided = self._guided_c0(self.model.lm_head_pruned(hidden).float(), cfg_scale, top_k) + code = sample_topk(guided, top_k, generator) + stop_token = 0 + offset = 1 + else: + logits = self.model.lm_head(hidden).float() + stop_token = SPECIAL_TOKEN_IDS["<|audio_end|>"] + logits = logits.masked_fill(vocab_mask, -float("inf")) + guided = self._guided_c0(logits, cfg_scale, top_k).masked_fill(vocab_mask, -float("inf")) + code = sample_topk(guided, top_k, generator) + offset = AUDIO_CODE_OFFSET + return torch.where(code == stop_token, 0, code - offset), code, stop_token + + def generate(self, input_ids, seed, max_audio_frames, device, cfg_scale=CFG_SCALE, top_k=CFG_TOP_K): + prompt_tokens = int(input_ids.shape[1]) + if prompt_tokens > MAX_PROMPT_TOKENS: + raise ValueError(f"MiniMax Music3 prompt has {prompt_tokens} tokens; maximum is {MAX_PROMPT_TOKENS}") + + input_ids = input_ids.to(device) + if comfy.model_management.should_use_bf16(device): + execution_dtype = torch.bfloat16 + else: + execution_dtype = torch.float32 + unconditioned = input_ids.clone() + unconditioned[:, 1:-2] = SPECIAL_TOKEN_IDS["<|audio_cfg|>"] + text_ids = torch.cat((input_ids, unconditioned), dim=0) + if self.model.pruned_embedding: + text_embeds = self.model.embed_tokens_prefill(text_ids, out_dtype=execution_dtype) + else: + text_embeds = self.model.embed_tokens(text_ids, out_dtype=execution_dtype) + decode_limit = min(int(max_audio_frames), MAX_AUDIO_FRAMES) + past = self.model.init_kv_cache(2, prompt_tokens + decode_limit + 1, device, execution_dtype) + output = self.model(None, embeds=text_embeds, past_key_values=past, dtype=execution_dtype) + last_hidden = output[0][:, -1] + past = output[2] + + generator = torch.Generator(device=device).manual_seed(derive_seed(seed, "ar")) + decoder = self.model.audio_decoder + depth_io = { + "hidden": torch.empty_like(last_hidden), + "c0": torch.empty((last_hidden.shape[0],), dtype=torch.long, device=device), + "c0_embed": torch.empty_like(last_hidden), + "codes": torch.empty((last_hidden.shape[0], self.num_codebooks), dtype=torch.long, device=device), + "depth_hidden": torch.empty((1, last_hidden.shape[-1] * (self.num_codebooks - 1)), dtype=execution_dtype, device=device), + } + decoder._comfy_cross_step_state = depth_io + comfy.model_management._register_cross_step(decoder) + hidden_frames = [] + pending_code = None + stop_token = None + pending_event = None + pending_hidden = None + progress = comfy.utils.ProgressBar(decode_limit) + cuda_device = torch.device(device).type == "cuda" + vocab_mask = None + if not self.model.pruned_lm_head: + vocab_mask = torch.ones(self.model.vocab_size, dtype=torch.bool, device=device) + vocab_mask[AUDIO_CODE_OFFSET:AUDIO_CODE_OFFSET + C0_VOCAB_SIZE] = False + vocab_mask[SPECIAL_TOKEN_IDS["<|audio_end|>"]] = False + + for frame_index in comfy.utils.model_trange(decode_limit + 1, desc="AR sampling"): + comfy.model_management.throw_exception_if_processing_interrupted() + if pending_code is not None: + if pending_event is not None: + pending_event.synchronize() + if int(pending_code.item()) == stop_token: + pending_hidden = None + break + if pending_hidden is not None: + hidden_frames.append(pending_hidden) + progress.update_absolute(len(hidden_frames)) + if len(hidden_frames) >= decode_limit: + break + + c0, code_or_stop, stop_token = self._sample_c0(last_hidden, cfg_scale, top_k, generator, vocab_mask) + if pending_code is None: + pending_code = torch.empty_like(code_or_stop, device="cpu", pin_memory=cuda_device) + if cuda_device: + pending_event = torch.cuda.Event() + pending_code.copy_(code_or_stop, non_blocking=cuda_device) + if pending_event is not None: + pending_event.record() + + c0 = c0.repeat(2) + c0_embed = self._embed_c0(c0, execution_dtype) + depth_io["hidden"].copy_(last_hidden) + depth_io["c0"].copy_(c0) + depth_io["c0_embed"].copy_(c0_embed) + + def depth_core(): + codes, depth_hidden = self._depth_codes( + depth_io["hidden"], depth_io["c0"], depth_io["c0_embed"], generator, execution_dtype, cfg_scale, top_k + ) + depth_io["codes"].copy_(codes) + depth_io["depth_hidden"].copy_(depth_hidden) + + depth_queue = comfy.model_prefetch.make_prefetch_queue( + [[decoder, self.model.audio_extra_embedding]], device, {"prefetch_dynamic_vbars": True} + ) + comfy.model_prefetch.prefetch_queue_pop( + depth_queue, device, decoder, execution_dtype, core=depth_core, enable_graph=True, generator=generator + ) + comfy.model_prefetch.prefetch_queue_pop(depth_queue, device, None) + feedback_codes = depth_io["codes"] + depth_hidden = depth_io["depth_hidden"] + frame_hidden = torch.cat((last_hidden[:1].detach(), depth_hidden), dim=-1) + if frame_index > 0: + pending_hidden = frame_hidden[0].clone() + + feedback = self._embed_audio_frame(feedback_codes, execution_dtype) + output = self.model(None, embeds=feedback, past_key_values=past, dtype=execution_dtype) + last_hidden = output[0][:, -1] + past = output[2] + + if pending_hidden is not None and len(hidden_frames) < decode_limit: + if pending_event is not None: + pending_event.synchronize() + if int(pending_code.item()) != stop_token: + hidden_frames.append(pending_hidden) + + if not hidden_frames: + raise ValueError("MiniMax Music3 generated zero audio frames") + return torch.stack(hidden_frames).to(device="cpu") diff --git a/comfy/ldm/minimax_music/dav.py b/comfy/ldm/minimax_music/dav.py new file mode 100644 index 000000000..d442559f4 --- /dev/null +++ b/comfy/ldm/minimax_music/dav.py @@ -0,0 +1,137 @@ +import math + +import torch +from torch import nn + +import comfy.ops + + +def snake(x, alpha): + shape = x.shape + flat = x.reshape(shape[0], shape[1], -1) + alpha = comfy.ops.cast_to_input(alpha, flat) + flat = flat + (alpha + 1e-9).reciprocal() * torch.sin(alpha * flat).pow(2) + return flat.reshape(shape) + + +class Snake1d(nn.Module): + def __init__(self, channels, dtype, device): + super().__init__() + self.alpha = nn.Parameter(torch.empty(1, channels, 1, dtype=dtype, device=device)) + + def forward(self, x): + return snake(x, self.alpha) + + +def _weight_norm_conv(operations, *args, **kwargs): + return nn.utils.parametrizations.weight_norm(operations.Conv1d(*args, **kwargs)) + + +def _weight_norm_conv_transpose(operations, *args, **kwargs): + return nn.utils.parametrizations.weight_norm(operations.ConvTranspose1d(*args, **kwargs)) + + +class ResidualUnit(nn.Module): + def __init__(self, dim, dilation, dtype, device, operations): + super().__init__() + padding = 3 * dilation + self.block = nn.Sequential( + Snake1d(dim, dtype, device), + _weight_norm_conv( + operations, + dim, + dim, + kernel_size=7, + dilation=dilation, + padding=padding, + dtype=dtype, + device=device, + ), + Snake1d(dim, dtype, device), + _weight_norm_conv(operations, dim, dim, kernel_size=1, dtype=dtype, device=device), + ) + + def forward(self, x): + residual = self.block(x) + if residual.shape[-1] != x.shape[-1]: + padding = (x.shape[-1] - residual.shape[-1]) // 2 + x = x[..., padding:x.shape[-1] - padding] + return x + residual + + +class DecoderBlock(nn.Module): + def __init__(self, input_dim, output_dim, stride, dtype, device, operations): + super().__init__() + self.block = nn.Sequential( + Snake1d(input_dim, dtype, device), + _weight_norm_conv_transpose( + operations, + input_dim, + output_dim, + kernel_size=2 * stride, + stride=stride, + padding=math.ceil(stride / 2), + dtype=dtype, + device=device, + ), + ResidualUnit(output_dim, 1, dtype, device, operations), + ResidualUnit(output_dim, 3, dtype, device, operations), + ResidualUnit(output_dim, 9, dtype, device, operations), + ) + + def forward(self, x): + return self.block(x) + + +class Decoder(nn.Module): + def __init__(self, dtype, device, operations): + super().__init__() + layers = [ + _weight_norm_conv( + operations, + 1024, + 1536, + kernel_size=7, + padding=3, + dtype=dtype, + device=device, + ) + ] + channels = 1536 + output_dim = channels + for index, stride in enumerate((8, 8, 4, 2)): + input_dim = channels // (2 ** index) + output_dim = channels // (2 ** (index + 1)) + layers.append(DecoderBlock(input_dim, output_dim, stride, dtype, device, operations)) + layers.extend(( + Snake1d(output_dim, dtype, device), + _weight_norm_conv( + operations, + output_dim, + 1, + kernel_size=7, + padding=3, + dtype=dtype, + device=device, + ), + nn.Tanh(), + )) + self.model = nn.Sequential(*layers) + + def forward(self, x): + return self.model(x) + + +class MiniMaxMusic3DAV(nn.Module): + def __init__(self, dtype=None, device=None, operations=None): + super().__init__() + self.dec_in_proj = operations.Conv1d(64, 1024, kernel_size=1, dtype=dtype, device=device) + self.decoder = Decoder(dtype, device, operations) + + def decode(self, latent): + batch, _, frames = latent.shape + folded = latent.reshape(batch * 2, 64, frames) + waveform = self.decoder(self.dec_in_proj(folded)) + return waveform.reshape(batch, 2, -1) + + forward = decode diff --git a/comfy/ldm/minimax_music/dit.py b/comfy/ldm/minimax_music/dit.py new file mode 100644 index 000000000..211e0d7db --- /dev/null +++ b/comfy/ldm/minimax_music/dit.py @@ -0,0 +1,213 @@ +import math + +import torch +from torch import nn + +import comfy.model_management +import comfy.ops +import comfy.quant_ops +from comfy.ldm.modules.attention import optimized_attention_for_device + + +MAX_CONDITION_FRAMES = 200 +CONDITION_HOP_FRAMES = 100 + + +def latent_length(audio_frames): + return max(1, int(audio_frames * 44100 / 24000 * 960 / 512)) + + +class FourierFeatures(nn.Module): + def __init__(self, in_features, out_features, dtype, device): + super().__init__() + self.weight = nn.Parameter(torch.empty(out_features // 2, in_features, dtype=dtype, device=device)) + + def forward(self, value): + weight = comfy.ops.cast_to_input(self.weight, value) + features = 2.0 * math.pi * value @ weight.T + return torch.cat((features.cos(), features.sin()), dim=-1) + + +class LayerNorm(nn.Module): + def __init__(self, dim, dtype, device): + super().__init__() + self.gamma = nn.Parameter(torch.empty(dim, dtype=dtype, device=device)) + self.register_buffer("beta", torch.empty(dim, dtype=dtype, device=device)) + + def forward(self, x): + return torch.nn.functional.layer_norm( + x, + (x.shape[-1],), + comfy.ops.cast_to_input(self.gamma, x), + comfy.ops.cast_to_input(self.beta, x), + ) + + +class RotaryEmbedding(nn.Module): + def __init__(self, dim, dtype, device): + super().__init__() + self.register_buffer("inv_freq", torch.empty(dim // 2, dtype=dtype, device=device)) + + def forward_from_seq_len(self, length, device, dtype): + positions = torch.arange(length, device=device, dtype=torch.float32) + frequencies = torch.outer(positions, comfy.ops.cast_to_input(self.inv_freq, positions)) + frequencies = frequencies.to(dtype) + cos, sin = frequencies.cos(), frequencies.sin() + return torch.stack((cos, -sin, sin, cos), dim=-1).reshape(1, 1, length, frequencies.shape[-1], 2, 2) + + +def _apply_rope(x, rotation_matrix): + x_dtype = x.dtype + x = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(rotation_matrix.dtype) + x = rotation_matrix[..., 0] * x[..., 0] + rotation_matrix[..., 1] * x[..., 1] + return x.movedim(-1, -2).flatten(-2).to(x_dtype) + + +class Attention(nn.Module): + def __init__(self, dim, dim_heads, dtype, device, operations): + super().__init__() + self.num_heads = dim // dim_heads + self.dim_heads = dim_heads + self.to_qkv = operations.Linear(dim, dim * 3, bias=False, dtype=dtype, device=device) + self.to_out = operations.Linear(dim, dim, bias=False, dtype=dtype, device=device) + + def forward(self, x, rotation_matrix): + batch, length, dim = x.shape + q, k, v = self.to_qkv(x).chunk(3, dim=-1) + q = q.reshape(batch, length, self.num_heads, self.dim_heads).transpose(1, 2) + k = k.reshape(batch, length, self.num_heads, self.dim_heads).transpose(1, 2) + v = v.reshape(batch, length, self.num_heads, self.dim_heads).transpose(1, 2) + rotary_dims = rotation_matrix.shape[-3] * 2 + if comfy.model_management.in_training: + q = torch.cat((_apply_rope(q[..., :rotary_dims], rotation_matrix), q[..., rotary_dims:]), dim=-1) + k = torch.cat((_apply_rope(k[..., :rotary_dims], rotation_matrix), k[..., rotary_dims:]), dim=-1) + else: + rotated_q, rotated_k = comfy.quant_ops.ck.apply_rope_split_half(q[..., :rotary_dims], k[..., :rotary_dims], rotation_matrix) + q = torch.cat((rotated_q, q[..., rotary_dims:]), dim=-1) + k = torch.cat((rotated_k, k[..., rotary_dims:]), dim=-1) + attention = optimized_attention_for_device(q.device) + out = attention(q, k, v, self.num_heads, skip_reshape=True) + return self.to_out(out) + + +class GLU(nn.Module): + def __init__(self, dim, inner_dim, dtype, device, operations): + super().__init__() + self.proj = operations.Linear(dim, inner_dim * 2, dtype=dtype, device=device) + + def forward(self, x): + value, gate = self.proj(x).chunk(2, dim=-1) + return value * torch.nn.functional.silu(gate) + + +class FeedForward(nn.Module): + def __init__(self, dim, inner_dim, dtype, device, operations): + super().__init__() + self.ff = nn.Sequential( + GLU(dim, inner_dim, dtype, device, operations), + nn.Identity(), + operations.Linear(inner_dim, dim, dtype=dtype, device=device), + ) + + def forward(self, x): + return self.ff(x) + + +class TransformerBlock(nn.Module): + def __init__(self, dim, dim_heads, inner_dim, dtype, device, operations): + super().__init__() + self.pre_norm = LayerNorm(dim, dtype, device) + self.self_attn = Attention(dim, dim_heads, dtype, device, operations) + self.ff_norm = LayerNorm(dim, dtype, device) + self.ff = FeedForward(dim, inner_dim, dtype, device, operations) + + def forward(self, x, rotation_matrix): + x = x + self.self_attn(self.pre_norm(x), rotation_matrix) + return x + self.ff(self.ff_norm(x)) + + +class ContinuousTransformer(nn.Module): + def __init__(self, dtype, device, operations): + super().__init__() + self.project_in = operations.Linear(2304, 2048, bias=False, dtype=dtype, device=device) + self.project_out = operations.Linear(2048, 128, bias=False, dtype=dtype, device=device) + self.rotary_pos_emb = RotaryEmbedding(32, dtype, device) + self.layers = nn.ModuleList([ + TransformerBlock(2048, 64, 8192, dtype, device, operations) + for _ in range(36) + ]) + + def forward(self, x, timestep_embedding): + x = self.project_in(x) + x = torch.cat((timestep_embedding.unsqueeze(1), x), dim=1) + rotation_matrix = self.rotary_pos_emb.forward_from_seq_len(x.shape[1], x.device, x.dtype) + for layer in self.layers: + x = layer(x, rotation_matrix) + return self.project_out(x[:, 1:]) + + +class DiffusionTransformer(nn.Module): + def __init__(self, dtype, device, operations): + super().__init__() + self.transformer = ContinuousTransformer(dtype, device, operations) + self.timestep_features = FourierFeatures(1, 256, dtype, device) + self.to_timestep_embed = nn.Sequential( + operations.Linear(256, 2048, dtype=dtype, device=device), + nn.SiLU(), + operations.Linear(2048, 2048, dtype=dtype, device=device), + ) + self.preprocess_conv = operations.Conv1d(2304, 2304, 1, bias=False, dtype=dtype, device=device) + self.postprocess_conv = operations.Conv1d(128, 128, 1, bias=False, dtype=dtype, device=device) + + def forward(self, x, timestep, condition): + full = torch.cat((x, torch.zeros_like(x), condition), dim=1) + full = self.preprocess_conv(full) + full + timestep_features = self.timestep_features(timestep[:, None]).to(dtype=x.dtype) + timestep_embedding = self.to_timestep_embed(timestep_features) + out = self.transformer(full.transpose(1, 2), timestep_embedding).transpose(1, 2) + return self.postprocess_conv(out) + out + + +class MiniMaxMusic3DiT(nn.Module): + def __init__(self, dtype=None, device=None, operations=None, **kwargs): + super().__init__() + self.dtype = dtype + self.latent_conditioners = nn.Sequential( + operations.Conv1d(4096, 2048, kernel_size=3, padding=1, dtype=dtype, device=device) + ) + self.diffusion_transformer = DiffusionTransformer(dtype, device, operations) + self.cond_layer_logits = nn.Parameter(torch.empty(8, dtype=dtype, device=device)) + self.cond_layer_scale = nn.Parameter(torch.empty(1, dtype=dtype, device=device)) + + def aligned_condition(self, hidden): + frames = hidden.shape[1] + hidden = hidden.transpose(1, 2).reshape(hidden.shape[0], 8, 4096, frames) + weights = torch.softmax(comfy.ops.cast_to_input(self.cond_layer_logits, hidden), dim=0) + hidden = torch.einsum("blht,l->bht", hidden, weights) + hidden = comfy.ops.cast_to_input(self.cond_layer_scale, hidden) * hidden + condition = self.latent_conditioners(hidden) + return torch.nn.functional.interpolate(condition, size=latent_length(frames), mode="nearest") + + def forward(self, x, timestep, context, conditioning_scale, **kwargs): + condition = self.aligned_condition(context) + condition = condition * conditioning_scale[:, :1, :1] + if condition.shape[-1] < x.shape[-1]: + condition = torch.nn.functional.pad(condition, (0, x.shape[-1] - condition.shape[-1])) + else: + condition = condition[..., :x.shape[-1]] + window = latent_length(MAX_CONDITION_FRAMES) + if x.shape[-1] <= window: + return -self.diffusion_transformer(x, timestep, condition) + + output = torch.zeros_like(x) + count = torch.zeros((1, 1, x.shape[-1]), device=x.device, dtype=x.dtype) + hop = latent_length(CONDITION_HOP_FRAMES) + start = 0 + while start < x.shape[-1]: + end = min(start + window, x.shape[-1]) + output[..., start:end] -= self.diffusion_transformer(x[..., start:end], timestep, condition[..., start:end]) + count[..., start:end] += 1 + if end == x.shape[-1]: + break + start += hop + return output / count diff --git a/comfy/ldm/minimax_music/prompt.py b/comfy/ldm/minimax_music/prompt.py new file mode 100644 index 000000000..5f197ee12 --- /dev/null +++ b/comfy/ldm/minimax_music/prompt.py @@ -0,0 +1,70 @@ +import re + + +SPECIAL_TOKEN_IDS = { + "<|im_start|>": 151644, + "<|im_end|>": 151645, + "<|audio_cfg|>": 151654, + "<|audio_start|>": 151669, + "<|audio_end|>": 151670, + "<|caption_start|>": 151671, + "<|caption_end|>": 151672, + "<|lyrics_start|>": 151673, + "<|lyrics_end|>": 151674, +} +AUDIO_CODE_OFFSET = 151675 + +_SPECIAL_TAG_RE = re.compile(r"<\|([^|]*)\|>") +_LYRIC_TAG_RE = re.compile(r"\s*(\[[^\]]+\])\s*") + + +def _remove_markdown_format(text): + lines = [] + for raw_line in text.splitlines(): + line = re.sub(r"^\s{0,3}#{1,6}\s+", "", raw_line) + line = re.sub(r"^\s*[*+-]\s+", "", line) + while "**" in line: + updated = re.sub(r"\*\*([^*]+)\*\*", r"\1", line) + if updated == line: + break + line = updated + line = re.sub(r"(?<|caption_start|>" + f"{clean_caption(caption)}" + "<|caption_end|><|lyrics_start|>" + f"{normalize_lyrics(lyrics)}" + "<|lyrics_end|><|im_end|><|audio_start|>" + ) + + +def validate_tokenizer(tokenizer): + for token, expected in SPECIAL_TOKEN_IDS.items(): + token_id = tokenizer.convert_tokens_to_ids(token) + if token_id != expected: + raise ValueError(f"MiniMax Music3 tokenizer mismatch for {token}: expected {expected}, got {token_id}") diff --git a/comfy/model_base.py b/comfy/model_base.py index 7d855f5a1..90cab7ac0 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -22,6 +22,7 @@ import torch import logging import comfy.ldm.lightricks.av_model import comfy.ldm.minimax.model +import comfy.ldm.minimax_music.dit import comfy.nested_tensor import comfy.ldm.lightricks.symmetric_patchifier import comfy.context_windows @@ -2337,6 +2338,18 @@ class ACEStep15(BaseModel): out['refer_audio'] = comfy.conds.CONDRegular(refer_audio) return out +class MiniMaxMusic3(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.minimax_music.dit.MiniMaxMusic3DiT) + + def process_timestep(self, timestep, **kwargs): + return 1.0 - timestep + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + out["conditioning_scale"] = comfy.conds.CONDRegular(kwargs["conditioning_scale"]) + return out + class Omnigen2(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.omnigen.omnigen2.OmniGen2Transformer2DModel) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index aec205290..e4bf30b78 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -44,6 +44,13 @@ def calculate_transformer_depth(prefix, state_dict_keys, state_dict): def detect_unet_config(state_dict, key_prefix, metadata=None): state_dict_keys = list(state_dict.keys()) + if ( + '{}cond_layer_logits'.format(key_prefix) in state_dict_keys + and '{}latent_conditioners.0.weight'.format(key_prefix) in state_dict_keys + and '{}diffusion_transformer.transformer.layers.0.self_attn.to_qkv.weight'.format(key_prefix) in state_dict_keys + ): + return {"audio_model": "minimax_music3"} + if '{}joint_blocks.0.context_block.attn.qkv.weight'.format(key_prefix) in state_dict_keys: #mmdit model unet_config = {} unet_config["in_channels"] = state_dict['{}x_embedder.proj.weight'.format(key_prefix)].shape[1] diff --git a/comfy/model_management.py b/comfy/model_management.py index 15c03dc77..ff963eb8e 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -1368,9 +1368,14 @@ STREAM_CAST_BUFFERS = {} LARGEST_CASTED_WEIGHT = (None, 0) STREAM_AIMDO_CAST_BUFFERS = {} LARGEST_AIMDO_CASTED_WEIGHT = (None, 0) +CROSS_STEP_STATE = weakref.WeakSet() DEFAULT_AIMDO_CAST_BUFFER_RESERVATION_SIZE = 16 * 1024 ** 3 +# NOTE: devs/agents: this is temporary and will be removed in a future comfy. Not supported for custom node use. +def _register_cross_step(module): + CROSS_STEP_STATE.add(module) + def get_cast_buffer(offload_stream, device, size, ref): global LARGEST_CASTED_WEIGHT @@ -1425,6 +1430,10 @@ def reset_cast_buffers(): mmap_obj.bounce() DIRTY_MMAPS.clear() + for module in CROSS_STEP_STATE: + del module._comfy_cross_step_state + CROSS_STEP_STATE.clear() + for loaded_model in current_loaded_models: model = loaded_model.model if model is not None and model.is_dynamic(): diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index cb44e7394..72942aa04 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -1887,8 +1887,29 @@ class ModelPatcherDynamic(ModelPatcher): loading = self._load_list(for_dynamic=True, default_device=device_to) loading.sort() + get_units = getattr(self.model, "get_dynamic_vram__units", None) + dynamic_units, last_dynamic_units = get_units() if get_units is not None else ([], []) + dynamic_units = list(dynamic_units) + last_dynamic_units = list(last_dynamic_units) + loading_by_module = {entry[-2]: entry for entry in loading} + loading = [] + for unit in dynamic_units: + unit_modules = unit if isinstance(unit, (list, tuple)) else (unit,) + modules = [module for root in unit_modules for module in root.modules() if module in loading_by_module] + for index, module in enumerate(modules): + loading.append((*loading_by_module.pop(module), unit if index == len(modules) - 1 else None)) + last_loading = [] + for unit in last_dynamic_units: + unit_modules = unit if isinstance(unit, (list, tuple)) else (unit,) + modules = [module for root in unit_modules for module in root.modules() if module in loading_by_module] + for index, module in enumerate(modules): + last_loading.append((*loading_by_module.pop(module), unit if index == len(modules) - 1 else None)) + loading.extend((*entry, None) for entry in loading_by_module.values()) + loading.extend(last_loading) + v_block = None + for x in loading: - *_, module_mem, n, m, params = x + *_, module_mem, n, m, params, end_of_block = x def set_dirty(item, dirty): if dirty or not hasattr(item, "_v_signature"): @@ -1981,6 +2002,13 @@ class ModelPatcherDynamic(ModelPatcher): move_weight_functions(m, device_to) + if hasattr(m, "_v"): + v_block = m._v if v_block is None else (v_block[0], v_block[1], max(v_block[2], m._v[1] + m._v[2] - v_block[1])) + if end_of_block is not None: + unit = end_of_block + (unit[0] if isinstance(unit, (list, tuple)) else unit)._v_block = v_block + v_block = None + for key, buf in self.model.named_buffers(recurse=True): if key not in self.backup_buffers: self.backup_buffers[key] = buf diff --git a/comfy/model_prefetch.py b/comfy/model_prefetch.py index aa6d22d77..2aad5eea7 100644 --- a/comfy/model_prefetch.py +++ b/comfy/model_prefetch.py @@ -1,11 +1,18 @@ +import torch +import weakref + import comfy_aimdo.model_vbar +from comfy.cli_args import args import comfy.memory_management import comfy.model_management import comfy.ops PREFETCH_QUEUES = [] +GRAPH_MODULES = weakref.WeakSet() +GRAPH_WARMED_MODULES = weakref.WeakSet() +GRAPH_CAPTURE_STREAMS = {} -def cleanup_prefetched_modules(comfy_modules): +def cleanup_prefetched_modules(module, comfy_modules): for s in comfy_modules: prefetch = getattr(s, "_prefetch", None) if prefetch is None: @@ -17,39 +24,74 @@ def cleanup_prefetched_modules(comfy_modules): if prefetch["signature"] is not None: comfy_aimdo.model_vbar.vbar_unpin(s._v) delattr(s, "_prefetch") + if getattr(module, "_v_block_faulted", False): + comfy_aimdo.model_vbar.vbar_unpin(module._v_block) + del module._v_block_faulted def cleanup_prefetch_queues(): - global PREFETCH_QUEUES + global PREFETCH_QUEUES, GRAPH_CAPTURE_STREAMS for queue in PREFETCH_QUEUES: for entry in queue: if entry is None or not isinstance(entry, tuple): continue _, prefetch_state = entry - comfy_modules = prefetch_state[1] + prefetched_module, comfy_modules = prefetch_state if comfy_modules is not None: - cleanup_prefetched_modules(comfy_modules) + cleanup_prefetched_modules(prefetched_module, comfy_modules) PREFETCH_QUEUES = [] + for module in GRAPH_MODULES: + del module._comfy_graph + GRAPH_MODULES.clear() + GRAPH_WARMED_MODULES.clear() + GRAPH_CAPTURE_STREAMS = {} -def prefetch_queue_pop(queue, device, module): +def prefetch_queue_pop(queue, device, module, dtype=None, core=None, enable_graph=False, generator=None): + enable_graph = enable_graph and not args.disable_cuda_graphs and comfy.model_management.is_device_cuda(device) if queue is None: + if core is not None: + core() return + capture_stream = None + if enable_graph: + capture_stream = GRAPH_CAPTURE_STREAMS.get(device) + if capture_stream is None: + capture_stream = torch.cuda.Stream(device=device) + GRAPH_CAPTURE_STREAMS[device] = capture_stream + + signature = None + graph_hit = False + graph = getattr(module, "_comfy_graph", None) if enable_graph else None + if graph is not None: + signature = comfy_aimdo.model_vbar.vbar_fault(module._v_block) + if signature is not None: + module._v_block_faulted = True + graph_hit = comfy_aimdo.model_vbar.vbar_signature_compare(signature, graph["signature"]) + consumed = queue.pop(0) if consumed is not None: offload_stream, prefetch_state = consumed if offload_stream is not None: offload_stream.wait_stream(comfy.model_management.current_stream(device)) - _, comfy_modules = prefetch_state + prefetched_module, comfy_modules = prefetch_state if comfy_modules is not None: - cleanup_prefetched_modules(comfy_modules) + cleanup_prefetched_modules(prefetched_module, comfy_modules) + if graph_hit: + queue[0] = (None, (module, [])) + graph["graph"].replay() + return + + fully_faulted = False prefetch = queue[0] if prefetch is not None: comfy_modules = [] - for s in prefetch.modules(): - if hasattr(s, "_v"): - comfy_modules.append(s) + prefetch_modules = prefetch if isinstance(prefetch, (list, tuple)) else (prefetch,) + for root in prefetch_modules: + for s in root.modules(): + if hasattr(s, "_v"): + comfy_modules.append(s) registerable_size = 0 for s in comfy_modules: @@ -59,11 +101,41 @@ def prefetch_queue_pop(queue, device, module): if lowvram_fn is not None: registerable_size += lowvram_fn.memory_required() - offload_stream = comfy.ops.cast_modules_with_vbar(comfy_modules, None, device, None, True) + offload_stream, fully_faulted = comfy.ops.cast_modules_with_vbar(comfy_modules, None, device, None, True, return_faulted=True) if not comfy.model_management.args.fast_disk: comfy.model_management.ensure_pin_registerable(registerable_size) comfy.model_management.sync_stream(device, offload_stream) - queue[0] = (offload_stream, (prefetch, comfy_modules)) + if fully_faulted and dtype is not None: + for comfy_module in comfy_modules: + comfy.ops.resolve_cast_module_with_vbar(comfy_module, dtype, device, dtype, None, False, return_weights=False) + queue[0] = (offload_stream, (module, comfy_modules)) + + if core is not None: + if enable_graph and fully_faulted and module in GRAPH_WARMED_MODULES: + if signature is None: + signature = comfy_aimdo.model_vbar.vbar_fault(module._v_block) + if signature is not None: + module._v_block_faulted = True + if signature is not None: + graph = torch.cuda.CUDAGraph() + if generator is not None: + graph.register_generator_state(generator) + capture_stream.wait_stream(comfy.model_management.current_stream(device)) + with torch.cuda.graph(graph, stream=capture_stream, capture_error_mode="thread_local"): + core() + comfy.model_management.current_stream(device).wait_stream(capture_stream) + graph.replay() + module._comfy_graph = {"graph": graph, "signature": signature} + GRAPH_MODULES.add(module) + return + if capture_stream is None: + core() + else: + capture_stream.wait_stream(comfy.model_management.current_stream(device)) + with torch.cuda.stream(capture_stream): + core() + comfy.model_management.current_stream(device).wait_stream(capture_stream) + GRAPH_WARMED_MODULES.add(module) def make_prefetch_queue(queue, device, transformer_options): if (not transformer_options.get("prefetch_dynamic_vbars", False) diff --git a/comfy/ops.py b/comfy/ops.py index 9ec44cfa2..73ae46674 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -123,10 +123,12 @@ def materialize_meta_param(s, param_keys): # FIXME: add n=1 cache hit fast path -def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blocking): +def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blocking, return_faulted=False): offload_stream = None cast_buffer = None cast_buffer_offset = 0 + if return_faulted: + fully_faulted = all(not getattr(s, param_key + "_function", []) for s in comfy_modules for param_key in ("weight", "bias")) def ensure_offload_stream(module, required_size, check_largest): nonlocal offload_stream @@ -163,6 +165,8 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin for s in comfy_modules: signature = comfy_aimdo.model_vbar.vbar_fault(s._v) resident = comfy_aimdo.model_vbar.vbar_signature_compare(signature, s._v_signature) + if return_faulted and (signature is None or not resident): + fully_faulted = False prefetch = { "signature": signature, "resident": resident, @@ -255,10 +259,12 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin prefetch["needs_cast"] = needs_cast s._prefetch = prefetch + if return_faulted: + return offload_stream, fully_faulted return offload_stream -def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, want_requant): +def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, want_requant, return_weights=True): prefetch = getattr(s, "_prefetch", None) @@ -298,7 +304,7 @@ def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, w tensor = tensor.dequantize() return tensor - if orig.dtype != dtype or len(fns) > 0: + if (return_weights and orig.dtype != dtype) or len(fns) > 0: x = to_dequant(x, dtype) if not resident and lowvram_fn is not None: x = to_dequant(x, dtype if compute_dtype is None else compute_dtype) @@ -325,7 +331,7 @@ def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, w if prefetch["signature"] is not None: prefetch["resident"] = True - return weight, bias + return (weight, bias) if return_weights else None def cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None, offloadable=False, compute_dtype=None, want_requant=False): diff --git a/comfy/sd.py b/comfy/sd.py index 46c9acba1..94f4f284f 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -25,6 +25,7 @@ import comfy.ldm.cogvideo.vae import comfy.ldm.hunyuan_video.vae import comfy.ldm.mmaudio.vae.autoencoder import comfy.ldm.audio.vae_sa3 +import comfy.ldm.minimax_music.dav import comfy.pixel_space_convert import comfy.weight_adapter import yaml @@ -32,6 +33,7 @@ import math import os import comfy.utils +import comfy.ops from . import clip_vision from . import gligen @@ -74,6 +76,7 @@ import comfy.text_encoders.longcat_image import comfy.text_encoders.qwen35 import comfy.text_encoders.qwen3vl import comfy.text_encoders.minimax +import comfy.text_encoders.minimax_music import comfy.ldm.minimax.vae import comfy.ldm.minimax.audio_vae import comfy.text_encoders.boogu @@ -515,7 +518,22 @@ class VAE: self.audio_sample_rate = 44100 if config is None: - if "decoder.mid.block_1.mix_factor" in sd: + if "dec_in_proj.weight" in sd and "decoder.model.0.weight_g" in sd: # MiniMax Music3 DAV + self.first_stage_model = comfy.ldm.minimax_music.dav.MiniMaxMusic3DAV(operations=comfy.ops.disable_weight_init) + self.latent_channels = 128 + self.output_channels = 2 + self.upscale_ratio = 512 + self.downscale_ratio = 512 + self.latent_dim = 1 + self.process_output = lambda audio: audio + self.process_input = lambda audio: audio + self.working_dtypes = [torch.float32] + self.disable_offload = True + self.memory_used_decode = lambda shape, dtype: (shape[-1] * 512 * 1400 + 800_000_000) * model_management.dtype_size(dtype) + def _no_encode(*args, **kwargs): + raise RuntimeError("MiniMax Music3 DAV cannot encode audio") + self.memory_used_encode = _no_encode + elif "decoder.mid.block_1.mix_factor" in sd: encoder_config = {'double_z': True, 'z_channels': 4, 'resolution': 256, 'in_channels': 3, 'out_ch': 3, 'ch': 128, 'ch_mult': [1, 2, 4, 4], 'num_res_blocks': 2, 'attn_resolutions': [], 'dropout': 0.0} decoder_config = encoder_config.copy() decoder_config["video_kernel_size"] = [3, 1, 1] @@ -1692,7 +1710,15 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.params = {} if len(clip_data) == 1: te_model = detect_te_model(clip_data[0]) - if te_model == TEModel.CLIP_G: + if clip_type == CLIPType.MINIMAX and "model.audio_decoder.projection.weight" in clip_data[0]: + tokenizer_data["tokenizer_json"] = clip_data[0].pop("tokenizer_json", None) + quant = comfy.utils.detect_layer_quantization(clip_data[0], "") + if quant is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = quant + clip_target.clip = comfy.text_encoders.minimax_music.MiniMaxMusic3TEModel + clip_target.tokenizer = comfy.text_encoders.minimax_music.MiniMaxMusic3Tokenizer + elif te_model == TEModel.CLIP_G: if clip_type == CLIPType.STABLE_CASCADE: clip_target.clip = sdxl_clip.StableCascadeClipModel clip_target.tokenizer = sdxl_clip.StableCascadeTokenizer diff --git a/comfy/supported_models.py b/comfy/supported_models.py index b9952db55..d6d3c857f 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -16,6 +16,7 @@ import comfy.text_encoders.genmo import comfy.text_encoders.lt import comfy.text_encoders.hunyuan_video import comfy.text_encoders.minimax +import comfy.text_encoders.minimax_music import comfy.text_encoders.cosmos import comfy.text_encoders.lumina2 import comfy.text_encoders.wan @@ -2200,6 +2201,25 @@ class ACEStep15(supported_models_base.BASE): return supported_models_base.ClipTarget(comfy.text_encoders.ace15.ACE15Tokenizer, comfy.text_encoders.ace15.te(**detect)) +class MiniMaxMusic3(supported_models_base.BASE): + unet_config = { + "audio_model": "minimax_music3", + } + + latent_format = comfy.latent_formats.MiniMaxMusic3 + memory_usage_factor = 2.0 + supported_inference_dtypes = [torch.float16, torch.bfloat16, torch.float32] + sampling_settings = {"multiplier": 1.0} + + def get_model(self, state_dict, prefix="", device=None): + return model_base.MiniMaxMusic3(self, device=device) + + def model_type(self, state_dict, prefix=""): + return model_base.ModelType.FLOW + + def clip_target(self, state_dict={}): + return supported_models_base.ClipTarget(comfy.text_encoders.minimax_music.MiniMaxMusic3Tokenizer, comfy.text_encoders.minimax_music.MiniMaxMusic3TEModel) + class LongCatImage(supported_models_base.BASE): unet_config = { @@ -2494,6 +2514,7 @@ models = [ ChromaRadiance, ACEStep, ACEStep15, + MiniMaxMusic3, Omnigen2, Boogu, MageFlow, diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index 371ec1bbc..4415d6e9f 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -5,15 +5,40 @@ from typing import Optional, Any, Tuple import math from tqdm import tqdm import comfy.utils +import comfy_kitchen from comfy.ldm.modules.attention import optimized_attention_for_device import comfy.model_management +import comfy.model_prefetch import comfy.ops import comfy.ldm.common_dit import comfy.clip_model from . import qwen_vl + +def detect_merged_config(state_dict, prefix="", layer_prefix="model.layers.0."): + return { + "merged_qkv": "{}{}self_attn.qkv_proj.weight".format(prefix, layer_prefix) in state_dict, + "merged_mlp": "{}{}mlp.gate_up_proj.weight".format(prefix, layer_prefix) in state_dict, + } + + +@dataclass +class FixedKV: + key: torch.Tensor + value: torch.Tensor + index: int + position: torch.Tensor + seqlen: torch.Tensor + + def prepare(self, num_tokens): + self.position.fill_(self.index) + self.seqlen.fill_(self.index + num_tokens) + + def advance(self, num_tokens): + self.index += num_tokens + @dataclass class Llama2Config: vocab_size: int = 128320 @@ -249,6 +274,9 @@ class Qwen3_8BConfig: rope_scale = None final_norm: bool = True lm_head: bool = True + fixed_kv: bool = False + merged_qkv: bool = False + merged_mlp: bool = False stop_tokens = [151643, 151645] @dataclass @@ -498,9 +526,14 @@ class Attention(nn.Module): self.inner_size = self.num_heads * self.head_dim ops = ops or nn - 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 * self.head_dim, bias=config.qkv_bias, device=device, dtype=dtype) - self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * self.head_dim, bias=config.qkv_bias, device=device, dtype=dtype) + self.kv_size = self.num_kv_heads * self.head_dim + self.merged_qkv = getattr(config, "merged_qkv", False) + if self.merged_qkv is not False: + self.qkv_proj = ops.Linear(config.hidden_size, self.inner_size + self.kv_size * 2, bias=config.qkv_bias, device=device, dtype=dtype) + if self.merged_qkv is not True: + 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.kv_size, bias=config.qkv_bias, device=device, dtype=dtype) + self.v_proj = ops.Linear(config.hidden_size, self.kv_size, bias=config.qkv_bias, device=device, dtype=dtype) self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype) self.q_norm = None @@ -522,9 +555,12 @@ class Attention(nn.Module): ): batch_size, seq_length, _ = hidden_states.shape - xq = self.q_proj(hidden_states) - xk = self.k_proj(hidden_states) - xv = self.v_proj(hidden_states) + if self.merged_qkv: + xq, xk, xv = self.qkv_proj(hidden_states).split((self.inner_size, self.kv_size, self.kv_size), dim=-1) + else: + xq = self.q_proj(hidden_states) + xk = self.k_proj(hidden_states) + xv = self.v_proj(hidden_states) xq = xq.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2) xk = xk.view(batch_size, seq_length, self.num_kv_heads, self.head_dim).transpose(1, 2) @@ -537,8 +573,29 @@ class Attention(nn.Module): xq, xk = apply_rope(xq, xk, freqs_cis=freqs_cis) - present_key_value = None - if past_key_value is not None: + fixed_cache = past_key_value if isinstance(past_key_value, FixedKV) else None + if fixed_cache is not None: + xq = xq.transpose(1, 2) + xk = xk.transpose(1, 2) + xv = xv.transpose(1, 2) + if seq_length == 1: + # CUDA-graphable decode path. + fixed_cache.key.index_copy_(1, fixed_cache.position, xk) + fixed_cache.value.index_copy_(1, fixed_cache.position, xv) + output = comfy_kitchen.flash_attention_decode(xq, fixed_cache.key, fixed_cache.value, fixed_cache.seqlen) + return self.o_proj(output.view(batch_size, seq_length, self.inner_size)), fixed_cache + + fixed_cache.key[:, fixed_cache.index:fixed_cache.index + seq_length].copy_(xk) + fixed_cache.value[:, fixed_cache.index:fixed_cache.index + seq_length].copy_(xv) + xk = fixed_cache.key[:, :fixed_cache.index + seq_length] + xv = fixed_cache.value[:, :fixed_cache.index + seq_length] + + xq = xq.transpose(1, 2) + xk = xk.transpose(1, 2) + xv = xv.transpose(1, 2) + + present_key_value = fixed_cache + if fixed_cache is None and past_key_value is not None: index = 0 num_tokens = xk.shape[2] if len(past_key_value) > 0: @@ -569,15 +626,27 @@ class MLP(nn.Module): def __init__(self, config: Llama2Config, device=None, dtype=None, ops: Any = None, intermediate_size=None): super().__init__() intermediate_size = intermediate_size or config.intermediate_size - self.gate_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) - self.up_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) + self.merged_mlp = getattr(config, "merged_mlp", False) + if self.merged_mlp is not False: + self.gate_up_proj = ops.Linear(config.hidden_size, intermediate_size * 2, bias=False, device=device, dtype=dtype) + if self.merged_mlp is not True: + self.gate_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) + self.up_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) self.down_proj = ops.Linear(intermediate_size, config.hidden_size, bias=False, device=device, dtype=dtype) if config.mlp_activation == "silu": self.activation = torch.nn.functional.silu + self.merged_input_act = "swiglu" elif config.mlp_activation == "gelu_pytorch_tanh": self.activation = lambda a: torch.nn.functional.gelu(a, approximate="tanh") + self.merged_input_act = None def forward(self, x): + if self.merged_mlp: + x = self.gate_up_proj(x) + if self.merged_input_act is not None: + return comfy.ops.linear_input_act(self.down_proj, x, self.merged_input_act) + gate, up = x.chunk(2, dim=-1) + return self.down_proj(self.activation(gate) * up) return self.down_proj(self.activation(self.gate_proj(x)) * self.up_proj(x)) class TransformerBlock(nn.Module): @@ -596,6 +665,7 @@ class TransformerBlock(nn.Module): optimized_attention=None, past_key_value: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, ): + output = x # Self Attention residual = x x = self.input_layernorm(x) @@ -612,7 +682,7 @@ class TransformerBlock(nn.Module): residual = x x = self.post_attention_layernorm(x) x = self.mlp(x) - x = residual + x + x = torch.add(residual, x, out=output) return x, present_key_value @@ -641,6 +711,7 @@ class TransformerBlockGemma2(nn.Module): optimized_attention=None, past_key_value: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, ): + output = x sliding_window = None if self.transformer_type == 'gemma3': if self.sliding_attention: @@ -676,7 +747,7 @@ class TransformerBlockGemma2(nn.Module): x = self.pre_feedforward_layernorm(x) x = self.mlp(x) x = self.post_feedforward_layernorm(x) - x = residual + x + x = torch.add(residual, x, out=output) return x, present_key_value @@ -688,9 +759,14 @@ def _make_scaled_embedding(ops, vocab_size, hidden_size, scale, device, dtype): class Llama2_(nn.Module): + fixed_kv = False + graph_dynamic_vbar_blocks = False + def __init__(self, config, device=None, dtype=None, ops=None): super().__init__() self.config = config + self.fixed_kv = getattr(config, "fixed_kv", False) + self.graph_dynamic_vbar_blocks = False self.vocab_size = config.vocab_size if self.config.transformer_type == "gemma2" or self.config.transformer_type == "gemma3": @@ -713,8 +789,27 @@ class Llama2_(nn.Module): if config.lm_head: self.lm_head = ops.Linear(config.hidden_size, config.vocab_size, bias=False, device=device, dtype=dtype) + def get_dynamic_vram__units(self): + return (list(self.layers), []) if self.graph_dynamic_vbar_blocks else ([], []) + def get_past_len(self, past_key_values): - return past_key_values[0][2] + first = past_key_values[0] + return first.index if isinstance(first, FixedKV) else first[2] + + def init_kv_cache(self, batch, capacity, device, dtype): + caches = [] + fixed_kv = self.fixed_kv and comfy_kitchen.flash_attention_decode_is_available(device) + for _ in range(self.config.num_hidden_layers): + if fixed_kv: + key = torch.empty((batch, capacity, self.config.num_key_value_heads, self.config.head_dim), device=device, dtype=dtype) + value = torch.empty_like(key) + position = torch.empty((1,), device=device, dtype=torch.int64) + seqlen = torch.empty((batch,), device=device, dtype=torch.int32) + caches.append(FixedKV(key, value, 0, position, seqlen)) + else: + key = torch.empty((batch, self.config.num_key_value_heads, capacity, self.config.head_dim), device=device, dtype=dtype) + caches.append((key, torch.empty_like(key), 0)) + return caches def compute_freqs_cis(self, position_ids, device): return precompute_freqs_cis(self.config.head_dim, @@ -756,6 +851,33 @@ class Llama2_(nn.Module): optimized_attention = optimized_attention_for_device(x.device, mask=mask is not None, small_input=True) + fixed_kv = past_key_values is not None and len(past_key_values) > 0 and isinstance(past_key_values[0], FixedKV) + enable_graph = self.graph_dynamic_vbar_blocks and fixed_kv and seq_len == 1 and mask is None + if enable_graph: + freqs_cis_groups = freqs_cis if isinstance(freqs_cis, list) else [freqs_cis] + cross_step_state_key = [(x.shape, x.stride(), x.dtype, x.device)] + for group in freqs_cis_groups: + for tensor in group: + cross_step_state_key.append((tensor.shape, tensor.stride(), tensor.dtype, tensor.device)) + cross_step_state_key = tuple(cross_step_state_key) + cross_step_state = getattr(self, "_comfy_cross_step_state", None) + if cross_step_state is None or cross_step_state["key"] != cross_step_state_key: + static_freqs_cis = [] + for group in freqs_cis_groups: + static_freqs_cis.append(tuple(torch.empty_like(tensor) for tensor in group)) + if not isinstance(freqs_cis, list): + static_freqs_cis = static_freqs_cis[0] + cross_step_state = {"key": cross_step_state_key, "x": torch.empty_like(x), "freqs_cis": static_freqs_cis} + self._comfy_cross_step_state = cross_step_state + comfy.model_management._register_cross_step(self) + cross_step_state["x"].copy_(x) + static_freqs_cis_groups = cross_step_state["freqs_cis"] if isinstance(freqs_cis, list) else [cross_step_state["freqs_cis"]] + for source_group, target_group in zip(freqs_cis_groups, static_freqs_cis_groups): + for source, target in zip(source_group, target_group): + target.copy_(source) + x = cross_step_state["x"] + freqs_cis = cross_step_state["freqs_cis"] + intermediate = None all_intermediate = None only_layers = None @@ -769,7 +891,8 @@ class Llama2_(nn.Module): elif intermediate_output < 0: intermediate_output = len(self.layers) + intermediate_output - next_key_values = [] + prefetch_queue = comfy.model_prefetch.make_prefetch_queue(list(self.layers), x.device, {"prefetch_dynamic_vbars": getattr(self, "prefetch_dynamic_vbars", False)}) + next_key_values = list(past_key_values) if past_key_values is not None else [] for i, layer in enumerate(self.layers): if all_intermediate is not None: if only_layers is None or (i in only_layers): @@ -779,16 +902,23 @@ class Llama2_(nn.Module): if past_key_values is not None: past_kv = past_key_values[i] if len(past_key_values) > 0 else [] - x, current_kv = layer( - x=x, - attention_mask=mask, - freqs_cis=freqs_cis, - optimized_attention=optimized_attention, - past_key_value=past_kv, - ) + if fixed_kv: + past_kv.prepare(seq_len) - if current_kv is not None: - next_key_values.append(current_kv) + def core(): + _, current_kv = layer( + x=x, + attention_mask=mask, + freqs_cis=freqs_cis, + optimized_attention=optimized_attention, + past_key_value=past_kv, + ) + if next_key_values: + next_key_values[i] = current_kv + + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, x.device, layer, x.dtype, core=core, enable_graph=enable_graph) + if fixed_kv: + next_key_values[i].advance(seq_len) # DeepStack: add per-layer visual features into the first len() decoder layers at image positions (Qwen3-VL) if deepstack_embeds is not None and i < len(deepstack_embeds): @@ -797,6 +927,9 @@ class Llama2_(nn.Module): if i == intermediate_output: intermediate = x.clone() + if prefetch_queue is not None: + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, x.device, None) + if self.norm is not None: x = self.norm(x) @@ -810,7 +943,7 @@ class Llama2_(nn.Module): if intermediate is not None and final_layer_norm_intermediate and self.norm is not None: intermediate = self.norm(intermediate) - if len(next_key_values) > 0: + if next_key_values: return x, intermediate, next_key_values else: return x, intermediate @@ -874,12 +1007,7 @@ class BaseGenerate: return torch.nn.functional.linear(input, weight, None) def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): - model_config = self.model.config - past_key_values = [] - for x in range(model_config.num_hidden_layers): - past_key_values.append((torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), - 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 self.model.init_kv_cache(batch, max_cache_len, device, execution_dtype) 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 diff --git a/comfy/text_encoders/minimax_music.py b/comfy/text_encoders/minimax_music.py new file mode 100644 index 000000000..37d072664 --- /dev/null +++ b/comfy/text_encoders/minimax_music.py @@ -0,0 +1,129 @@ +import torch +from tokenizers import Tokenizer + +import comfy.ops +import comfy.text_encoders.llama +from comfy.ldm.minimax_music.ar import CFG_SCALE, CFG_TOP_K, MAX_AUDIO_FRAMES, MiniMaxMusic3AR +from comfy.ldm.minimax_music.prompt import SPECIAL_TOKEN_IDS, build_prompt + + +MODEL_CONFIG = { + "vocab_size": 200000, + "hidden_size": 4096, + "intermediate_size": 12288, + "num_hidden_layers": 36, + "num_attention_heads": 32, + "num_key_value_heads": 8, + "max_position_embeddings": 10240, + "rms_norm_eps": 1e-6, + "rope_theta": 1000000.0, + "head_dim": 128, + "audio_vocab_size": 1024, + "audio_num_codebooks": 8, + "decoder_num_heads": 16, + "decoder_intermediate_size": 6144, + "decoder_num_layers": 4, +} + + +class MiniMaxMusic3Tokenizer: + def __init__(self, embedding_directory=None, tokenizer_data={}): + tokenizer_json = tokenizer_data.get("tokenizer_json") + if tokenizer_json is None: + raise ValueError("MiniMax Music3 text encoder checkpoint is missing tokenizer_json") + if torch.is_tensor(tokenizer_json): + tokenizer_json = tokenizer_json.detach().cpu().numpy().tobytes() + self.tokenizer_json = tokenizer_json + self.tokenizer = Tokenizer.from_str(tokenizer_json.decode("utf-8")) + for token, expected in SPECIAL_TOKEN_IDS.items(): + if self.tokenizer.token_to_id(token) != expected: + raise ValueError(f"MiniMax Music3 tokenizer mismatch for {token}") + + def tokenize_with_weights(self, text, return_word_ids=False, **kwargs): + prompt = build_prompt(text, kwargs.get("lyrics", "")) + token_ids = self.tokenizer.encode(prompt, add_special_tokens=False).ids + return { + "minimax_music3": [[(token, 1.0) for token in token_ids]], + "seed": int(kwargs.get("seed", 0)), + "max_audio_frames": int(kwargs.get("max_audio_frames", MAX_AUDIO_FRAMES)), + "cfg_scale": float(kwargs.get("cfg_scale", CFG_SCALE)), + "top_k": int(kwargs.get("top_k", CFG_TOP_K)), + } + + def state_dict(self): + return {"tokenizer_json": torch.frombuffer(bytearray(self.tokenizer_json), dtype=torch.uint8)} + + def decode(self, token_ids, skip_special_tokens=True): + return self.tokenizer.decode(token_ids, skip_special_tokens=skip_special_tokens) + + +class MiniMaxMusic3TEModel(MiniMaxMusic3AR): + def __init__(self, device="cpu", dtype=None, model_options={}): + dtype = torch.bfloat16 + quant_config = model_options.get("quantization_metadata", None) + operations = model_options.get("custom_operations", None) + if operations is None: + operations = comfy.ops.mixed_precision_ops(quant_config, dtype) if quant_config is not None else comfy.ops.manual_cast + super().__init__(MODEL_CONFIG, dtype, device, operations) + self.dtypes = {dtype} + self.execution_device = device + + def set_clip_options(self, options): + self.execution_device = options.get("execution_device", self.execution_device) + + def reset_clip_options(self): + pass + + def get_dynamic_vram__units(self): + units, last_units = self.model.get_dynamic_vram__units() + if self.model.pruned_embedding: + last_units = [*last_units, self.model.embed_tokens_prefill] + return [(self.model.audio_decoder, self.model.audio_extra_embedding), *units], last_units + + def encode_token_weights(self, token_weight_pairs): + token_ids = [token for token, _ in token_weight_pairs["minimax_music3"][0]] + input_ids = torch.tensor([token_ids], dtype=torch.long) + seed = token_weight_pairs["seed"] + max_audio_frames = token_weight_pairs["max_audio_frames"] + cfg_scale = token_weight_pairs["cfg_scale"] + top_k = token_weight_pairs["top_k"] + hidden = self.generate(input_ids, seed, max_audio_frames, self.execution_device, cfg_scale, top_k) + return hidden.unsqueeze(0), None, {} + + def load_state_dict(self, state_dict, strict=True, assign=False): + def select_projections(layers, config): + for layer in layers: + if layer.self_attn.merged_qkv is None: + if config["merged_qkv"]: + del layer.self_attn.q_proj, layer.self_attn.k_proj, layer.self_attn.v_proj + else: + del layer.self_attn.qkv_proj + layer.self_attn.merged_qkv = config["merged_qkv"] + if layer.mlp.merged_mlp is None: + if config["merged_mlp"]: + del layer.mlp.gate_proj, layer.mlp.up_proj + else: + del layer.mlp.gate_up_proj + layer.mlp.merged_mlp = config["merged_mlp"] + + select_projections(self.model.layers, comfy.text_encoders.llama.detect_merged_config(state_dict)) + select_projections( + self.model.audio_decoder.layers, + comfy.text_encoders.llama.detect_merged_config(state_dict, layer_prefix="model.audio_decoder.layers.0."), + ) + if self.model.pruned_embedding is None: + self.model.pruned_embedding = "model.embed_tokens_prefill.weight" in state_dict + if self.model.pruned_embedding: + del self.model.embed_tokens + else: + del self.model.embed_tokens_prefill, self.model.embed_tokens_audio + if self.model.pruned_lm_head is None: + self.model.pruned_lm_head = "model.lm_head_pruned.weight" in state_dict + if self.model.pruned_lm_head: + del self.model.lm_head + else: + del self.model.lm_head_pruned + return super().load_state_dict(state_dict, strict=strict, assign=assign) + + def load_sd(self, state_dict): + return self.load_state_dict(state_dict, strict=False, assign=getattr(self, "can_assign_sd", False)) diff --git a/comfy_extras/nodes_minimax_music.py b/comfy_extras/nodes_minimax_music.py new file mode 100644 index 000000000..e22103b08 --- /dev/null +++ b/comfy_extras/nodes_minimax_music.py @@ -0,0 +1,77 @@ +import torch +from typing_extensions import override + +import comfy.model_management +from comfy.ldm.minimax_music.ar import AUDIO_FRAMES_PER_SECOND, CFG_SCALE, CFG_TOP_K, C0_VOCAB_SIZE, MAX_AUDIO_FRAMES +from comfy.ldm.minimax_music.dit import latent_length +from comfy_api.latest import ComfyExtension, io + + +class MiniMaxMusic3TextEncode(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="MiniMaxMusic3TextEncode", + display_name="MiniMax Music3 Text Encode", + category="model/conditioning/minimax music", + description="Uses a MiniMax Music3 CLIP model to generate the acoustic conditioning sequence.", + inputs=[ + io.Clip.Input("clip"), + io.String.Input("caption", multiline=True, dynamic_prompts=True), + io.String.Input("lyrics", multiline=True, dynamic_prompts=True), + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff, control_after_generate=True), + io.Float.Input("max_duration", default=120.0, min=0.04, max=MAX_AUDIO_FRAMES / AUDIO_FRAMES_PER_SECOND, step=0.04, tooltip="Maximum duration in seconds; the model can end the song earlier."), + io.Float.Input("cfg_scale", default=CFG_SCALE, min=0.0, max=100.0, step=0.1, round=0.01, advanced=True), + io.Int.Input("top_k", default=CFG_TOP_K, min=1, max=C0_VOCAB_SIZE, advanced=True), + ], + outputs=[ + io.Conditioning.Output(), + io.Float.Output(display_name="seconds"), + ], + ) + + @classmethod + def execute(cls, clip, caption, lyrics, seed, max_duration, cfg_scale, top_k): + max_audio_frames = min(MAX_AUDIO_FRAMES, max(1, round(max_duration * AUDIO_FRAMES_PER_SECOND))) + tokens = clip.tokenize(caption, lyrics=lyrics, seed=seed, max_audio_frames=max_audio_frames, cfg_scale=cfg_scale, top_k=top_k) + conditioning = clip.encode_from_tokens_scheduled(tokens) + for cond in conditioning: + hidden = cond[0] + cond[1]["conditioning_scale"] = torch.ones((hidden.shape[0], 1, 1), device=hidden.device, dtype=hidden.dtype) + return io.NodeOutput(conditioning, conditioning[0][0].shape[1] / AUDIO_FRAMES_PER_SECOND) + + +class EmptyMiniMaxMusic3LatentAudio(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="EmptyMiniMaxMusic3LatentAudio", + display_name="Empty MiniMax Music3 Latent Audio", + category="model/latent/minimax music", + description="Creates an empty MiniMax Music3 audio latent for the requested duration.", + inputs=[ + io.Float.Input("seconds", default=120.0, min=0.04, max=MAX_AUDIO_FRAMES / AUDIO_FRAMES_PER_SECOND, step=0.04), + io.Int.Input("batch_size", default=1, min=1, max=4096), + ], + outputs=[io.Latent.Output()], + ) + + @classmethod + def execute(cls, seconds, batch_size): + audio_frames = min(MAX_AUDIO_FRAMES, max(1, round(seconds * AUDIO_FRAMES_PER_SECOND))) + latent = torch.zeros( + (batch_size, 128, latent_length(audio_frames)), + device=comfy.model_management.intermediate_device(), + dtype=comfy.model_management.intermediate_dtype(), + ) + return io.NodeOutput({"samples": latent, "type": "audio", "downscale_ratio_temporal": 512}) + + +class MiniMaxMusic3Extension(ComfyExtension): + @override + async def get_node_list(self): + return [MiniMaxMusic3TextEncode, EmptyMiniMaxMusic3LatentAudio] + + +async def comfy_entrypoint(): + return MiniMaxMusic3Extension() diff --git a/nodes.py b/nodes.py index ec298e1de..1a3dd3f48 100644 --- a/nodes.py +++ b/nodes.py @@ -290,6 +290,9 @@ class ConditioningZeroOut: conditioning_lyrics = d.get("conditioning_lyrics", None) if conditioning_lyrics is not None: d["conditioning_lyrics"] = torch.zeros_like(conditioning_lyrics) + conditioning_scale = d.get("conditioning_scale", None) + if conditioning_scale is not None: + d["conditioning_scale"] = torch.zeros_like(conditioning_scale) n = [torch.zeros_like(t[0]), d] c.append(n) return (c, ) @@ -1015,7 +1018,7 @@ class CLIPLoader: CATEGORY = "model/loaders" - DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\njoyimage: qwen3-vl 8B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm" + DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\njoyimage: qwen3-vl 8B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm\nminimax: MiniMax H3 Qwen3-VL or Music3 Qwen/RVQ" def load_clip(self, clip_name, type="stable_diffusion", device="default"): clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION) @@ -2449,6 +2452,7 @@ async def init_builtin_extra_nodes(): "nodes_mahiro.py", "nodes_lt_upsampler.py", "nodes_lt_audio.py", + "nodes_minimax_music.py", "nodes_minimax_h3.py", "nodes_lt.py", "nodes_hooks.py", From e535e59e133b0921b2e396678d1fd060322e594b Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Thu, 13 Aug 2026 19:18:12 +0300 Subject: [PATCH 37/76] [Partner Nodes] feat(Bria): add GenFill, Eraser, Expand and Increase Resolution nodes (#15572) Signed-off-by: Alexander Piskun --- comfy_api_nodes/apis/bria.py | 95 ++++++ comfy_api_nodes/nodes_bria.py | 524 ++++++++++++++++++++++++++++++++++ 2 files changed, 619 insertions(+) diff --git a/comfy_api_nodes/apis/bria.py b/comfy_api_nodes/apis/bria.py index 7a98428c3..f55de74bc 100644 --- a/comfy_api_nodes/apis/bria.py +++ b/comfy_api_nodes/apis/bria.py @@ -57,6 +57,81 @@ class BriaRemoveBackgroundRequest(BaseModel): seed: int = Field(...) +class BriaGenFillRequest(BaseModel): + image: str = Field(...) + mask: str = Field( + ..., + description="Binary mask defining the region to fill: white (255) pixels are generated, " + "black (0) pixels are preserved. Must have the same aspect ratio as the image.", + ) + prompt: str = Field(...) + negative_prompt: str | None = Field(None) + refine_prompt: bool = Field(True) + seed: int = Field(...) + prompt_content_moderation: bool = Field(False, description="If true, returns 422 on prompt moderation failure.") + visual_input_content_moderation: bool = Field( + False, description="If true, returns 422 on image or mask moderation failure." + ) + visual_output_content_moderation: bool = Field( + False, description="If true, returns 422 on visual output moderation failure." + ) + + +class BriaEraseRequest(BaseModel): + image: str = Field(...) + mask: str = Field( + ..., + description="Binary mask defining the region to erase: white (255) pixels are removed, " + "black (0) pixels are preserved. Must have the same aspect ratio as the image.", + ) + mask_type: str = Field("manual", description="'manual' for hand-drawn masks, 'automatic' for segmentation masks.") + visual_input_content_moderation: bool = Field( + False, description="If true, returns 422 on image or mask moderation failure." + ) + visual_output_content_moderation: bool = Field( + False, description="If true, returns 422 on visual output moderation failure." + ) + + +class BriaExpandRequest(BaseModel): + image: str = Field(...) + aspect_ratio: str | float | None = Field( + None, + description="Target ratio: a preset string (1:1, 2:3, 3:2, 3:4, 4:3, 4:5, 5:4, 9:16, 16:9) " + "or a float between 0.5 and 3.0. When set, the canvas/placement fields are ignored.", + ) + canvas_size: list[int] | None = Field(None, description="Output canvas [width, height]; area up to 5000x5000.") + original_image_size: list[int] | None = Field( + None, description="Size [width, height] of the original image inside the canvas." + ) + original_image_location: list[int] | None = Field( + None, + description="Top-left corner [x, y] of the original image inside the canvas; " + "values may fall outside the canvas, cropping the image.", + ) + prompt: str | None = Field(None, description="If omitted, Bria auto-generates a prompt from the image.") + negative_prompt: str | None = Field(None) + seed: int = Field(...) + prompt_content_moderation: bool = Field(False, description="If true, returns 422 on prompt moderation failure.") + visual_input_content_moderation: bool = Field( + False, description="If true, returns 422 on image moderation failure." + ) + visual_output_content_moderation: bool = Field( + False, description="If true, returns 422 on visual output moderation failure." + ) + + +class BriaIncreaseResolutionRequest(BaseModel): + image: str = Field(...) + desired_increase: int = Field(..., description="Resolution multiplier, 2 or 4.") + visual_input_content_moderation: bool = Field( + False, description="If true, returns 422 on image moderation failure." + ) + visual_output_content_moderation: bool = Field( + False, description="If true, returns 422 on visual output moderation failure." + ) + + class BriaStatusResponse(BaseModel): request_id: str = Field(...) status_url: str = Field(...) @@ -72,6 +147,26 @@ class BriaRemoveBackgroundResponse(BaseModel): result: BriaRemoveBackgroundResult | None = Field(None) +class BriaImageResult(BaseModel): + image_url: str = Field(...) + + +class BriaImageResultResponse(BaseModel): + status: str = Field(...) + result: BriaImageResult | None = Field(None) + + +class BriaExpandResult(BaseModel): + image_url: str = Field(...) + prompt: str | None = Field(None) + seed: int | None = Field(None) + + +class BriaExpandResponse(BaseModel): + status: str = Field(...) + result: BriaExpandResult | None = Field(None) + + class BriaImageEditResult(BaseModel): structured_prompt: str = Field(...) image_url: str = Field(...) diff --git a/comfy_api_nodes/nodes_bria.py b/comfy_api_nodes/nodes_bria.py index 77f780a3b..90cade2d0 100644 --- a/comfy_api_nodes/nodes_bria.py +++ b/comfy_api_nodes/nodes_bria.py @@ -6,7 +6,13 @@ from typing_extensions import override from comfy_api.latest import IO, ComfyExtension, Input from comfy_api_nodes.apis.bria import ( BriaEditImageRequest, + BriaEraseRequest, + BriaExpandRequest, + BriaExpandResponse, + BriaGenFillRequest, BriaImageEditResponse, + BriaImageResultResponse, + BriaIncreaseResolutionRequest, BriaRemoveBackgroundRequest, BriaRemoveBackgroundResponse, BriaRemoveVideoBackgroundRequest, @@ -21,13 +27,30 @@ from comfy_api_nodes.util import ( convert_mask_to_image, download_url_to_image_tensor, download_url_to_video_output, + downscale_image_tensor_by_max_side, + get_image_dimensions, poll_op, sync_op, upload_image_to_comfyapi, upload_video_to_comfyapi, + validate_string, validate_video_duration, ) +BRIA_MAX_OUTPUT_SIDE = 8192 +BRIA_MIN_RATIO = 0.5 +BRIA_MAX_RATIO = 3.0 +BRIA_MIN_SHORT_SIDE = 224 + + +def _upscaled_output_side(height: int, width: int, multiplier: int) -> int: + prescale = max(1.0, BRIA_MIN_SHORT_SIDE / min(height, width)) + return round(max(height, width) * prescale * multiplier) + + +def _smallest_output_side(height: int, width: int, multiplier: int) -> int: + return round(max(height, width) / min(height, width) * BRIA_MIN_SHORT_SIDE * multiplier) + class BriaImageEditNode(IO.ComfyNode): @@ -243,6 +266,503 @@ class BriaRemoveImageBackground(IO.ComfyNode): return IO.NodeOutput(await download_url_to_image_tensor(response.result.image_url)) +def _mask_to_binary_image(mask: Input.Image, action: str) -> torch.Tensor: + binary = (mask > 0.5).float() + if not binary.any(): + raise ValueError( + f"The mask is empty, so there is nothing to {action}. Masks are binarized at 50%: " + f"areas painted at less than half opacity are ignored." + ) + return convert_mask_to_image(binary) + + +def _validate_mask_aspect_ratio(image: Input.Image, mask: Input.Image) -> None: + ih, iw = image.shape[1], image.shape[2] + mh, mw = mask.shape[-2], mask.shape[-1] + if abs(iw * mh - ih * mw) > 0.01 * ih * mw: + raise ValueError(f"Mask must have the same aspect ratio as the image: image is {iw}x{ih}, mask is {mw}x{mh}.") + + +class BriaGenFill(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="BriaGenFill", + display_name="Bria Generative Fill", + category="partner/image/Bria", + description="Generate objects or scenery inside a masked region of an image using Bria.", + inputs=[ + IO.Image.Input("image"), + IO.Mask.Input( + "mask", + tooltip="White areas are filled with generated content, black areas are preserved. " + "The mask is binarized before sending, so partially painted areas count as white. " + "Must have the same aspect ratio as the image.", + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Description of what to generate inside the masked region.", + ), + IO.String.Input("negative_prompt", multiline=True, default=""), + IO.Boolean.Input( + "refine_prompt", + default=True, + tooltip="Automatically adjust the prompt for better results; " + "disable to use the prompt exactly as written.", + ), + IO.Int.Input( + "seed", + default=42, + min=1, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + ), + IO.DynamicCombo.Input( + "moderation", + options=[ + IO.DynamicCombo.Option("false", []), + IO.DynamicCombo.Option( + "true", + [ + IO.Boolean.Input("prompt_content_moderation", default=False), + IO.Boolean.Input("visual_input_moderation", default=False), + IO.Boolean.Input("visual_output_moderation", default=False), + ], + ), + ], + tooltip="Moderation settings", + ), + ], + outputs=[IO.Image.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.0429}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + mask: Input.Image, + prompt: str, + negative_prompt: str, + refine_prompt: bool, + seed: int, + moderation: InputModerationSettings, + ) -> IO.NodeOutput: + validate_string(prompt, min_length=1) + _validate_mask_aspect_ratio(image, mask) + mask_image = _mask_to_binary_image(mask, "fill") + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/bria/v2/image/edit/gen_fill", method="POST"), + data=BriaGenFillRequest( + image=await upload_image_to_comfyapi(cls, image, total_pixels=None, wait_label="Uploading image"), + mask=await upload_image_to_comfyapi( + cls, mask_image, total_pixels=None, wait_label="Uploading mask" + ), + prompt=prompt, + negative_prompt=negative_prompt if negative_prompt else None, + refine_prompt=refine_prompt, + seed=seed, + prompt_content_moderation=moderation.get("prompt_content_moderation", False), + visual_input_content_moderation=moderation.get("visual_input_moderation", False), + visual_output_content_moderation=moderation.get("visual_output_moderation", False), + ), + response_model=BriaStatusResponse, + ) + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/bria/v2/status/{response.request_id}"), + status_extractor=lambda r: r.status, + response_model=BriaImageResultResponse, + ) + return IO.NodeOutput(await download_url_to_image_tensor(response.result.image_url)) + + +class BriaEraser(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="BriaEraser", + display_name="Bria Eraser", + category="partner/image/Bria", + description="Remove objects or areas outlined by a mask from an image using Bria.", + inputs=[ + IO.Image.Input("image"), + IO.Mask.Input( + "mask", + tooltip="White areas are erased, black areas are preserved. " + "The mask is binarized before sending, so partially painted areas count as white. " + "Must have the same aspect ratio as the image.", + ), + IO.Combo.Input( + "mask_type", + options=["manual", "automatic"], + tooltip="manual for hand-drawn or brush masks, " + "automatic for masks produced by segmentation models such as SAM.", + ), + IO.DynamicCombo.Input( + "moderation", + options=[ + IO.DynamicCombo.Option("false", []), + IO.DynamicCombo.Option( + "true", + [ + IO.Boolean.Input("visual_input_moderation", default=False), + IO.Boolean.Input("visual_output_moderation", default=False), + ], + ), + ], + tooltip="Moderation settings", + ), + ], + outputs=[IO.Image.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.0286}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + mask: Input.Image, + mask_type: str, + moderation: dict, + ) -> IO.NodeOutput: + _validate_mask_aspect_ratio(image, mask) + mask_image = _mask_to_binary_image(mask, "erase") + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/bria/v2/image/edit/erase", method="POST"), + data=BriaEraseRequest( + image=await upload_image_to_comfyapi(cls, image, total_pixels=None, wait_label="Uploading image"), + mask=await upload_image_to_comfyapi( + cls, mask_image, total_pixels=None, wait_label="Uploading mask" + ), + mask_type=mask_type, + visual_input_content_moderation=moderation.get("visual_input_moderation", False), + visual_output_content_moderation=moderation.get("visual_output_moderation", False), + ), + response_model=BriaStatusResponse, + ) + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/bria/v2/status/{response.request_id}"), + status_extractor=lambda r: r.status, + response_model=BriaImageResultResponse, + ) + return IO.NodeOutput(await download_url_to_image_tensor(response.result.image_url)) + + +class BriaExpandImage(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="BriaExpandImage", + display_name="Bria Expand Image", + category="partner/image/Bria", + description="Expand an image beyond its borders with generated content using Bria.", + inputs=[ + IO.Image.Input("image"), + IO.DynamicCombo.Input( + "expand_mode", + options=[ + *[IO.DynamicCombo.Option(ratio, []) for ratio in + ["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"]], + IO.DynamicCombo.Option( + "custom_ratio", + [ + IO.Int.Input( + "ratio_width", + default=21, + min=1, + max=100, + tooltip="Width side of the target ratio: 21 and 9 give 21:9.", + ), + IO.Int.Input( + "ratio_height", + default=9, + min=1, + max=100, + tooltip="Height side of the target ratio: 21 and 9 give 21:9. " + f"Bria only accepts width/height between {BRIA_MIN_RATIO} and " + f"{BRIA_MAX_RATIO}, so anything taller than 1:2 needs the manual mode.", + ), + ], + ), + IO.DynamicCombo.Option( + "manual", + [ + IO.Int.Input("canvas_width", default=1000, min=64, max=5000), + IO.Int.Input("canvas_height", default=1000, min=64, max=5000), + IO.Int.Input( + "image_width", + default=500, + min=1, + max=5000, + tooltip="Width of the original image inside the canvas.", + ), + IO.Int.Input( + "image_height", + default=500, + min=1, + max=5000, + tooltip="Height of the original image inside the canvas.", + ), + IO.Int.Input( + "image_x", + default=250, + min=-5000, + max=5000, + tooltip="X position of the image's top-left corner inside the canvas; " + "may fall outside the canvas, cropping the image.", + ), + IO.Int.Input( + "image_y", + default=250, + min=-5000, + max=5000, + tooltip="Y position of the image's top-left corner inside the canvas; " + "may fall outside the canvas, cropping the image.", + ), + ], + ), + ], + tooltip="Target shape of the expanded image: a preset aspect ratio, a custom ratio, " + "or manual placement of the original image on a canvas. " + "Manual is the only mode that can reach a canvas taller than 1:2.", + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Optional description of the expanded scene; " + "when empty, Bria generates one from the image.", + ), + IO.String.Input("negative_prompt", multiline=True, default=""), + IO.Int.Input( + "seed", + default=42, + min=1, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + ), + IO.DynamicCombo.Input( + "moderation", + options=[ + IO.DynamicCombo.Option("false", []), + IO.DynamicCombo.Option( + "true", + [ + IO.Boolean.Input("prompt_content_moderation", default=False), + IO.Boolean.Input("visual_input_moderation", default=False), + IO.Boolean.Input("visual_output_moderation", default=False), + ], + ), + ], + tooltip="Moderation settings", + ), + ], + outputs=[ + IO.Image.Output(), + IO.String.Output(display_name="prompt", tooltip="The prompt used for the expansion; " + "auto-generated by Bria when the prompt input is empty."), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.0286}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + expand_mode: dict, + prompt: str, + negative_prompt: str, + seed: int, + moderation: InputModerationSettings, + ) -> IO.NodeOutput: + mode = expand_mode["expand_mode"] + aspect_ratio = canvas_size = original_image_size = original_image_location = None + if mode == "manual": + canvas_size = [expand_mode["canvas_width"], expand_mode["canvas_height"]] + original_image_size = [expand_mode["image_width"], expand_mode["image_height"]] + original_image_location = [expand_mode["image_x"], expand_mode["image_y"]] + elif mode == "custom_ratio": + ratio_width, ratio_height = expand_mode["ratio_width"], expand_mode["ratio_height"] + aspect_ratio = ratio_width / ratio_height + if not BRIA_MIN_RATIO <= aspect_ratio <= BRIA_MAX_RATIO: + raise ValueError( + f"Bria accepts a width-to-height ratio between {BRIA_MIN_RATIO} and {BRIA_MAX_RATIO}: " + f"{ratio_width}:{ratio_height} is {aspect_ratio:.4f}. " + f"Use the manual expand mode to reach a canvas of any shape." + ) + else: + aspect_ratio = mode + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/bria/v2/image/edit/expand", method="POST"), + data=BriaExpandRequest( + image=await upload_image_to_comfyapi(cls, image, total_pixels=None, wait_label="Uploading image"), + aspect_ratio=aspect_ratio, + canvas_size=canvas_size, + original_image_size=original_image_size, + original_image_location=original_image_location, + prompt=prompt if prompt else None, + negative_prompt=negative_prompt if negative_prompt else None, + seed=seed, + prompt_content_moderation=moderation.get("prompt_content_moderation", False), + visual_input_content_moderation=moderation.get("visual_input_moderation", False), + visual_output_content_moderation=moderation.get("visual_output_moderation", False), + ), + response_model=BriaStatusResponse, + ) + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/bria/v2/status/{response.request_id}"), + status_extractor=lambda r: r.status, + response_model=BriaExpandResponse, + ) + return IO.NodeOutput( + await download_url_to_image_tensor(response.result.image_url), + response.result.prompt or "", + ) + + +class BriaIncreaseResolution(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="BriaIncreaseResolution", + display_name="Bria Increase Resolution", + category="partner/image/Bria", + description="Upscale an image by 2x or 4x using Bria, preserving the original content.", + inputs=[ + IO.Image.Input("image"), + IO.Combo.Input( + "desired_increase", + options=["2", "4"], + tooltip="Resolution multiplier. The output must fit within 8192 pixels on each side.", + ), + IO.Boolean.Input( + "auto_downscale", + default=False, + tooltip="Automatically lower the multiplier, and downscale the input image if that is " + "still not enough, when the output would exceed the limit.", + ), + IO.DynamicCombo.Input( + "moderation", + options=[ + IO.DynamicCombo.Option("false", []), + IO.DynamicCombo.Option( + "true", + [ + IO.Boolean.Input("visual_input_moderation", default=False), + IO.Boolean.Input("visual_output_moderation", default=False), + ], + ), + ], + tooltip="Moderation settings", + ), + ], + outputs=[IO.Image.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.0286}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + desired_increase: str, + auto_downscale: bool, + moderation: dict, + ) -> IO.NodeOutput: + multiplier = int(desired_increase) + height, width = get_image_dimensions(image) + if _upscaled_output_side(height, width, multiplier) > BRIA_MAX_OUTPUT_SIDE: + candidates = [c for c in (4, 2) if c <= multiplier] + if not auto_downscale: + predicted = _upscaled_output_side(height, width, multiplier) + raise ValueError( + f"Bria can upscale up to a maximum output dimension of {BRIA_MAX_OUTPUT_SIDE} pixels: " + f"input is {width}x{height}, x{multiplier} would be {predicted} pixels on the long side. " + f"Enable auto_downscale, or use a smaller input image or a lower multiplier." + ) + fitted = next( + (c for c in candidates if _upscaled_output_side(height, width, c) <= BRIA_MAX_OUTPUT_SIDE), None + ) + if fitted is not None: + multiplier = fitted + else: + shrinkable = next((c for c in sorted(candidates) if _smallest_output_side(height, width, c) + <= BRIA_MAX_OUTPUT_SIDE), None) + if shrinkable is None: + raise ValueError( + f"This image cannot be upscaled by Bria at any multiplier: it is {width}x{height}, and " + f"Bria first enlarges the short side to {BRIA_MIN_SHORT_SIDE} pixels, which pushes the " + f"long side past the {BRIA_MAX_OUTPUT_SIDE} pixel limit. Crop it to a squarer shape first." + ) + multiplier = shrinkable + image = downscale_image_tensor_by_max_side(image, max_side=BRIA_MAX_OUTPUT_SIDE // multiplier) + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/bria/v2/image/edit/increase_resolution", method="POST"), + data=BriaIncreaseResolutionRequest( + image=await upload_image_to_comfyapi(cls, image, total_pixels=None, wait_label="Uploading image"), + desired_increase=multiplier, + visual_input_content_moderation=moderation.get("visual_input_moderation", False), + visual_output_content_moderation=moderation.get("visual_output_moderation", False), + ), + response_model=BriaStatusResponse, + ) + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/bria/v2/status/{response.request_id}"), + status_extractor=lambda r: r.status, + response_model=BriaImageResultResponse, + ) + return IO.NodeOutput(await download_url_to_image_tensor(response.result.image_url)) + + class BriaRemoveVideoBackground(IO.ComfyNode): @classmethod @@ -572,6 +1092,10 @@ class BriaExtension(ComfyExtension): return [ BriaImageEditNode, BriaRemoveImageBackground, + BriaGenFill, + BriaEraser, + BriaExpandImage, + BriaIncreaseResolution, BriaRemoveVideoBackground, BriaVideoGreenScreen, BriaVideoReplaceBackground, From af3d2153a73f8be48719f8c77f752a454de75902 Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Fri, 14 Aug 2026 02:21:26 +1000 Subject: [PATCH 38/76] llama: fix non-local x path (#15580) --- comfy/text_encoders/llama.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index 4415d6e9f..ddac547e0 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -906,7 +906,8 @@ class Llama2_(nn.Module): past_kv.prepare(seq_len) def core(): - _, current_kv = layer( + nonlocal x + x, current_kv = layer( x=x, attention_mask=mask, freqs_cis=freqs_cis, From 86aedfd943d36d485e5ed3cb9d962f21f73d1741 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Fri, 14 Aug 2026 00:23:02 +0800 Subject: [PATCH 39/76] chore: update workflow templates to v0.11.41 (#15578) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 4f505cc9c..e180e7884 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.48.7 -comfyui-workflow-templates==0.11.40 +comfyui-workflow-templates==0.11.41 comfyui-embedded-docs==0.5.9 torch torchsde From ddbaa8752874c275290d054ee4fddd6e004f5fdf Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Fri, 14 Aug 2026 03:08:08 +1000 Subject: [PATCH 40/76] minimax: early detect qkv vs q,k,v (#15581) avoid a commit charge surge on non-dynamic windows due to double linear creation. --- comfy/ldm/minimax_music/ar.py | 38 ++++++++++++++++------------ comfy/sd.py | 1 + comfy/supported_models.py | 5 +++- comfy/text_encoders/llama.py | 15 +++-------- comfy/text_encoders/minimax_music.py | 34 ++++++++----------------- 5 files changed, 42 insertions(+), 51 deletions(-) diff --git a/comfy/ldm/minimax_music/ar.py b/comfy/ldm/minimax_music/ar.py index 28215cb58..2a8935318 100644 --- a/comfy/ldm/minimax_music/ar.py +++ b/comfy/ldm/minimax_music/ar.py @@ -43,15 +43,17 @@ def sample_topk(logits, top_k, generator): class RVQAttention(nn.Module): - def __init__(self, hidden_size, num_heads, dtype, device, operations): + def __init__(self, hidden_size, num_heads, merged_qkv, dtype, device, operations): super().__init__() self.num_heads = num_heads self.head_dim = hidden_size // num_heads - self.merged_qkv = None - self.qkv_proj = operations.Linear(hidden_size, hidden_size * 3, bias=False, dtype=dtype, device=device) - self.q_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) - self.k_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) - self.v_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.merged_qkv = merged_qkv + if merged_qkv: + self.qkv_proj = operations.Linear(hidden_size, hidden_size * 3, bias=False, dtype=dtype, device=device) + else: + self.q_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.k_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.v_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) self.o_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) def forward(self, x): @@ -81,12 +83,14 @@ class RVQRMSNorm(nn.Module): class RVQMLP(nn.Module): - def __init__(self, hidden_size, intermediate_size, dtype, device, operations): + def __init__(self, hidden_size, intermediate_size, merged_mlp, dtype, device, operations): super().__init__() - self.merged_mlp = None - self.gate_up_proj = operations.Linear(hidden_size, intermediate_size * 2, bias=False, dtype=dtype, device=device) - self.gate_proj = operations.Linear(hidden_size, intermediate_size, bias=False, dtype=dtype, device=device) - self.up_proj = operations.Linear(hidden_size, intermediate_size, bias=False, dtype=dtype, device=device) + self.merged_mlp = merged_mlp + if merged_mlp: + self.gate_up_proj = operations.Linear(hidden_size, intermediate_size * 2, bias=False, dtype=dtype, device=device) + else: + self.gate_proj = operations.Linear(hidden_size, intermediate_size, bias=False, dtype=dtype, device=device) + self.up_proj = operations.Linear(hidden_size, intermediate_size, bias=False, dtype=dtype, device=device) self.down_proj = operations.Linear(intermediate_size, hidden_size, bias=False, dtype=dtype, device=device) def forward(self, x): @@ -96,12 +100,12 @@ class RVQMLP(nn.Module): class RVQDecoderBlock(nn.Module): - def __init__(self, hidden_size, num_heads, intermediate_size, dtype, device, operations): + def __init__(self, hidden_size, num_heads, intermediate_size, merged_qkv, merged_mlp, dtype, device, operations): super().__init__() self.input_layernorm = RVQRMSNorm(hidden_size, dtype, device) - self.self_attn = RVQAttention(hidden_size, num_heads, dtype, device, operations) + self.self_attn = RVQAttention(hidden_size, num_heads, merged_qkv, dtype, device, operations) self.post_attention_layernorm = RVQRMSNorm(hidden_size, dtype, device) - self.mlp = RVQMLP(hidden_size, intermediate_size, dtype, device, operations) + self.mlp = RVQMLP(hidden_size, intermediate_size, merged_mlp, dtype, device, operations) def forward(self, x): x = x + self.self_attn(self.input_layernorm(x)) @@ -113,6 +117,8 @@ class RVQDepthDecoder(nn.Module): super().__init__() hidden_size = int(config["hidden_size"]) audio_vocab_size = int(config["audio_vocab_size"]) + merged_qkv = config.get("decoder_merged_qkv", False) + merged_mlp = config.get("decoder_merged_mlp", False) num_codebooks = int(config["audio_num_codebooks"]) self.projection = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) self.pos_embedding = operations.Embedding(16, hidden_size, dtype=dtype, device=device) @@ -125,6 +131,8 @@ class RVQDepthDecoder(nn.Module): hidden_size, int(config["decoder_num_heads"]), int(config["decoder_intermediate_size"]), + merged_qkv, + merged_mlp, dtype, device, operations, @@ -148,8 +156,6 @@ class MiniMaxMusic3AR(nn.Module): qwen_config = Qwen3_8BConfig(**{key: value for key, value in config.items() if key in config_fields}) qwen_config.lm_head = False qwen_config.fixed_kv = True - qwen_config.merged_qkv = None - qwen_config.merged_mlp = None self.model = Llama2_(qwen_config, device=device, dtype=dtype, ops=operations) self.model.prefetch_dynamic_vbars = True self.model.graph_dynamic_vbar_blocks = True diff --git a/comfy/sd.py b/comfy/sd.py index 94f4f284f..4bdaa978c 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -1716,6 +1716,7 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip if quant is not None: model_options = model_options.copy() model_options["quantization_metadata"] = quant + clip_target.params["projection_config"] = comfy.text_encoders.minimax_music.detect_merged_config(clip_data[0]) clip_target.clip = comfy.text_encoders.minimax_music.MiniMaxMusic3TEModel clip_target.tokenizer = comfy.text_encoders.minimax_music.MiniMaxMusic3Tokenizer elif te_model == TEModel.CLIP_G: diff --git a/comfy/supported_models.py b/comfy/supported_models.py index d6d3c857f..33c378435 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -2218,7 +2218,10 @@ class MiniMaxMusic3(supported_models_base.BASE): return model_base.ModelType.FLOW def clip_target(self, state_dict={}): - return supported_models_base.ClipTarget(comfy.text_encoders.minimax_music.MiniMaxMusic3Tokenizer, comfy.text_encoders.minimax_music.MiniMaxMusic3TEModel) + detect = comfy.text_encoders.minimax_music.detect_merged_config(state_dict, self.text_encoder_key_prefix[0]) + target = supported_models_base.ClipTarget(comfy.text_encoders.minimax_music.MiniMaxMusic3Tokenizer, comfy.text_encoders.minimax_music.MiniMaxMusic3TEModel) + target.params["projection_config"] = detect + return target class LongCatImage(supported_models_base.BASE): diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index ddac547e0..f182e5147 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -17,13 +17,6 @@ import comfy.clip_model from . import qwen_vl -def detect_merged_config(state_dict, prefix="", layer_prefix="model.layers.0."): - return { - "merged_qkv": "{}{}self_attn.qkv_proj.weight".format(prefix, layer_prefix) in state_dict, - "merged_mlp": "{}{}mlp.gate_up_proj.weight".format(prefix, layer_prefix) in state_dict, - } - - @dataclass class FixedKV: key: torch.Tensor @@ -528,9 +521,9 @@ class Attention(nn.Module): ops = ops or nn self.kv_size = self.num_kv_heads * self.head_dim self.merged_qkv = getattr(config, "merged_qkv", False) - if self.merged_qkv is not False: + if self.merged_qkv: self.qkv_proj = ops.Linear(config.hidden_size, self.inner_size + self.kv_size * 2, bias=config.qkv_bias, device=device, dtype=dtype) - if self.merged_qkv is not True: + else: 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.kv_size, bias=config.qkv_bias, device=device, dtype=dtype) self.v_proj = ops.Linear(config.hidden_size, self.kv_size, bias=config.qkv_bias, device=device, dtype=dtype) @@ -627,9 +620,9 @@ class MLP(nn.Module): super().__init__() intermediate_size = intermediate_size or config.intermediate_size self.merged_mlp = getattr(config, "merged_mlp", False) - if self.merged_mlp is not False: + if self.merged_mlp: self.gate_up_proj = ops.Linear(config.hidden_size, intermediate_size * 2, bias=False, device=device, dtype=dtype) - if self.merged_mlp is not True: + else: self.gate_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) self.up_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) self.down_proj = ops.Linear(intermediate_size, config.hidden_size, bias=False, device=device, dtype=dtype) diff --git a/comfy/text_encoders/minimax_music.py b/comfy/text_encoders/minimax_music.py index 37d072664..c88d463cc 100644 --- a/comfy/text_encoders/minimax_music.py +++ b/comfy/text_encoders/minimax_music.py @@ -2,7 +2,6 @@ import torch from tokenizers import Tokenizer import comfy.ops -import comfy.text_encoders.llama from comfy.ldm.minimax_music.ar import CFG_SCALE, CFG_TOP_K, MAX_AUDIO_FRAMES, MiniMaxMusic3AR from comfy.ldm.minimax_music.prompt import SPECIAL_TOKEN_IDS, build_prompt @@ -26,6 +25,15 @@ MODEL_CONFIG = { } +def detect_merged_config(state_dict, prefix=""): + return { + "merged_qkv": "{}model.layers.0.self_attn.qkv_proj.weight".format(prefix) in state_dict, + "merged_mlp": "{}model.layers.0.mlp.gate_up_proj.weight".format(prefix) in state_dict, + "decoder_merged_qkv": "{}model.audio_decoder.layers.0.self_attn.qkv_proj.weight".format(prefix) in state_dict, + "decoder_merged_mlp": "{}model.audio_decoder.layers.0.mlp.gate_up_proj.weight".format(prefix) in state_dict, + } + + class MiniMaxMusic3Tokenizer: def __init__(self, embedding_directory=None, tokenizer_data={}): tokenizer_json = tokenizer_data.get("tokenizer_json") @@ -58,13 +66,13 @@ class MiniMaxMusic3Tokenizer: class MiniMaxMusic3TEModel(MiniMaxMusic3AR): - def __init__(self, device="cpu", dtype=None, model_options={}): + def __init__(self, device="cpu", dtype=None, model_options={}, projection_config=None): dtype = torch.bfloat16 quant_config = model_options.get("quantization_metadata", None) operations = model_options.get("custom_operations", None) if operations is None: operations = comfy.ops.mixed_precision_ops(quant_config, dtype) if quant_config is not None else comfy.ops.manual_cast - super().__init__(MODEL_CONFIG, dtype, device, operations) + super().__init__({**MODEL_CONFIG, **(projection_config or {})}, dtype, device, operations) self.dtypes = {dtype} self.execution_device = device @@ -91,26 +99,6 @@ class MiniMaxMusic3TEModel(MiniMaxMusic3AR): return hidden.unsqueeze(0), None, {} def load_state_dict(self, state_dict, strict=True, assign=False): - def select_projections(layers, config): - for layer in layers: - if layer.self_attn.merged_qkv is None: - if config["merged_qkv"]: - del layer.self_attn.q_proj, layer.self_attn.k_proj, layer.self_attn.v_proj - else: - del layer.self_attn.qkv_proj - layer.self_attn.merged_qkv = config["merged_qkv"] - if layer.mlp.merged_mlp is None: - if config["merged_mlp"]: - del layer.mlp.gate_proj, layer.mlp.up_proj - else: - del layer.mlp.gate_up_proj - layer.mlp.merged_mlp = config["merged_mlp"] - - select_projections(self.model.layers, comfy.text_encoders.llama.detect_merged_config(state_dict)) - select_projections( - self.model.audio_decoder.layers, - comfy.text_encoders.llama.detect_merged_config(state_dict, layer_prefix="model.audio_decoder.layers.0."), - ) if self.model.pruned_embedding is None: self.model.pruned_embedding = "model.embed_tokens_prefill.weight" in state_dict if self.model.pruned_embedding: From 2f35f4a08176d993cded35dac3332be4f7287f41 Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Thu, 13 Aug 2026 13:20:22 -0400 Subject: [PATCH 41/76] ComfyUI v0.33.0 --- comfyui_version.py | 2 +- pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/comfyui_version.py b/comfyui_version.py index 568e75b88..42f52f6b8 100644 --- a/comfyui_version.py +++ b/comfyui_version.py @@ -1,3 +1,3 @@ # This file is automatically generated by the build process when version is # updated in pyproject.toml. -__version__ = "0.32.0" +__version__ = "0.33.0" diff --git a/pyproject.toml b/pyproject.toml index 941b04d45..e7feada84 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ComfyUI" -version = "0.32.0" +version = "0.33.0" readme = "README.md" license = { file = "LICENSE" } requires-python = ">=3.10" From 03fa4e48ba524173736bf299ee0f981fc57c7414 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 13 Aug 2026 12:47:44 -0700 Subject: [PATCH 42/76] Fix minimax music not working on non dynamic vram. (#15588) --- comfy/model_prefetch.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy/model_prefetch.py b/comfy/model_prefetch.py index 2aad5eea7..7aedab530 100644 --- a/comfy/model_prefetch.py +++ b/comfy/model_prefetch.py @@ -47,7 +47,7 @@ def cleanup_prefetch_queues(): GRAPH_CAPTURE_STREAMS = {} def prefetch_queue_pop(queue, device, module, dtype=None, core=None, enable_graph=False, generator=None): - enable_graph = enable_graph and not args.disable_cuda_graphs and comfy.model_management.is_device_cuda(device) + enable_graph = enable_graph and not args.disable_cuda_graphs and comfy.model_management.is_device_cuda(device) and getattr(module, "_v_block", None) is not None if queue is None: if core is not None: core() From e01fb4c56b7a88149d469b99cbbfe3223d715054 Mon Sep 17 00:00:00 2001 From: Barish Ozbay <17261091+drozbay@users.noreply.github.com> Date: Thu, 13 Aug 2026 15:55:36 -0400 Subject: [PATCH 43/76] Add MiniMaxH3AddGuide for anchoring image and audio guides at any frame (#15439) --- comfy/ldm/minimax/model.py | 73 +++++++++++++-------- comfy/model_base.py | 10 +-- comfy_extras/nodes_minimax_h3.py | 108 ++++++++++++++++++++++++++----- 3 files changed, 143 insertions(+), 48 deletions(-) diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py index f745db884..b6feb8860 100644 --- a/comfy/ldm/minimax/model.py +++ b/comfy/ldm/minimax/model.py @@ -91,6 +91,18 @@ def _video_t_grid(n, origin): return float(origin) + torch.cat([torch.zeros(1, dtype=torch.float64), spans[:-1].cumsum(0)]) +def _ref_t_span(blk): + # time-axis span a reference block occupies ahead of the target streams + kind = blk["kind"] + if kind == "image": + return 1.0 + if kind == "audio": + return float(blk["ref_audio_t"]) + if kind in ("video", "video_audio"): + return max(float(blk["ref_audio_t"]), sum(_video_t_spans(blk["latent_t"]))) + return 0.0 + + def _audio_grid(cursor, t, w_low, w_high): # channel-major stereo rows: t advances per latent frame, w pinned to the grid extremes per stereo channel, h stays 0 g = torch.zeros(t * 2, 3, dtype=torch.float64) @@ -288,7 +300,7 @@ class FinalLayer(nn.Module): class PackedLayout: """Static packed-sequence structure for one shape/conditioning signature.""" - def __init__(self, text_len, latent_t, latent_h, latent_w, audio_t, keyframes=None, refs=None, frame_count=None): + def __init__(self, text_len, latent_t, latent_h, latent_w, audio_t, keyframes=None, refs=None): frame, w_grid = _frame_grid(latent_h, latent_w) frame_rows = frame.shape[0] @@ -299,29 +311,37 @@ class PackedLayout: img_pos, img_update = [], [] audio_pos, audio_update = [], [] - cursor = text_len row = text_len - if keyframes: - # fl2va: keyframe cond rows right after text, sharing the target spatial grid - for kf in keyframes: - pixel_index = kf["resolved_frame_index"] - if pixel_index == 0: - cond_t = float(text_len) - elif frame_count is not None and pixel_index == frame_count - 1: - cond_t = float(text_len) + sum(_video_t_spans(latent_t)) - FRAME_RESCALE - else: - raise ValueError("only first/last keyframe anchors are supported") - g = torch.empty(frame_rows, 3, dtype=torch.float64) - g[:, 0] = cond_t - g[:, 1:] = frame - segments.append(("cond", frame_rows)) - pos.append(g) - img_pos.append(torch.arange(row, row + frame_rows)) - img_update.append(torch.zeros(frame_rows, dtype=torch.bool)) - row += frame_rows - target_audio_w = (float(w_grid[0]), float(w_grid[-1])) + # refs pack between text and the targets, so the target timeline starts after their spans + cursor = float(text_len) + for blk in refs or (): + cursor += _ref_t_span(blk) + + if keyframes: + # fl2va: keyframe cond rows right after text, sharing the target spatial grid; + # anchors count from the target timeline origin, FRAME_RESCALE per pixel frame, 1.0 per audio latent frame + for kf in keyframes: + cond_t = cursor + FRAME_RESCALE * kf["resolved_frame_index"] + video_latent = kf.get("latent") + if video_latent is not None: + vt = video_latent.shape[2] + n = vt * frame_rows + segments.append(("cond", n)) + pos.append(_video_grid(vt, frame, cond_t)) + img_pos.append(torch.arange(row, row + n)) + img_update.append(torch.zeros(n, dtype=torch.bool)) + row += n + audio_latent = kf.get("audio_latent") + if audio_latent is not None: + rt = audio_latent.shape[-1] + segments.append(("cond_audio", rt * 2)) + pos.append(_audio_grid(cond_t, rt, *target_audio_w)) + audio_pos.append(torch.arange(row, row + rt * 2)) + audio_update.append(torch.zeros(rt * 2, dtype=torch.bool)) + row += rt * 2 + if refs: cursor = float(text_len) for blk in refs: @@ -389,7 +409,7 @@ class PackedLayout: self.audio_update = torch.cat(audio_update) self.signature = (text_len, latent_t, latent_h, latent_w, audio_t) # contiguous segment table (start, stop, kind) - # kinds: text / cond / ref_img / ref_audio / audio / video + # kinds: text / cond / cond_audio / ref_img / ref_audio / audio / video # the packed sequence is uniform per segment in (modality tag, timestep class), # except the text span (tag runs resolved at forward time from the presentation tags) seg_abs = [] @@ -529,8 +549,7 @@ class MiniMaxH3Model(nn.Module): if layout is None or layout.signature != (text_len, latent_t, lat_h, lat_w, audio_t): layout = PackedLayout(text_len, latent_t, lat_h, lat_w, audio_t, keyframes=payload.get("keyframes"), - refs=payload.get("refs"), - frame_count=payload.get("frame_count")) + refs=payload.get("refs")) # model_base passes model_sampling.timestep(sigma) = sigma * 1000 shift_v = float(transformer_options.get("minimax_h3_sigma_shift_video", self.sigma_shift_video)) @@ -543,14 +562,14 @@ class MiniMaxH3Model(nn.Module): vis_aug = float(payload.get("visual_cond_noise_aug", VISUAL_COND_TIMESTEP)) aud_aug = float(payload.get("audio_cond_noise_aug", AUDIO_COND_TIMESTEP)) has_vis_cond = any(k in ("cond", "ref_img") for _, _, k in layout.segments) - has_aud_cond = any(k == "ref_audio" for _, _, k in layout.segments) + has_aud_cond = any(k in ("cond_audio", "ref_audio") for _, _, k in layout.segments) seg_t = {"text": t_v, "video": t_v, "audio": t_a, "cond": max(t_v, vis_aug), "ref_img": max(t_v, vis_aug), - "ref_audio": max(t_a, aud_aug)} + "cond_audio": max(t_a, aud_aug), "ref_audio": max(t_a, aud_aug)} unique_t = sorted({t_v, t_a} | ({seg_t["cond"]} if has_vis_cond else set()) | ({seg_t["ref_audio"]} if has_aud_cond else set())) t_row = {t: i for i, t in enumerate(unique_t)} - seg_tag = {"text": 1, "video": 0, "audio": 2, "cond": 0, "ref_img": 0, "ref_audio": 2} + seg_tag = {"text": 1, "video": 0, "audio": 2, "cond": 0, "ref_img": 0, "cond_audio": 2, "ref_audio": 2} text_tags = payload.get("text_token_tags") mod_segments = [] diff --git a/comfy/model_base.py b/comfy/model_base.py index 90cab7ac0..6705eb6c3 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -2165,13 +2165,13 @@ class MiniMaxH3(BaseModel): keyframes = kwargs.get("minimax_keyframes", None) if keyframes is not None: payload["keyframes"] = keyframes - payload["frame_count"] = kwargs.get("minimax_frame_count", None) - payload["cond_video_latents"] = [kf["latent"] for kf in keyframes] + payload["cond_video_latents"] = [kf["latent"] for kf in keyframes if kf.get("latent") is not None] + payload["cond_audio_latents"] = [kf["audio_latent"] for kf in keyframes if kf.get("audio_latent") is not None] refs = kwargs.get("minimax_refs", None) if refs is not None: payload["refs"] = refs - payload["cond_video_latents"] = [r["latent"] for r in refs if "latent" in r] - payload["cond_audio_latents"] = [r["audio_latent"] for r in refs if r.get("audio_latent") is not None] + payload["cond_video_latents"] = payload.get("cond_video_latents", []) + [r["latent"] for r in refs if "latent" in r] + payload["cond_audio_latents"] = payload.get("cond_audio_latents", []) + [r["audio_latent"] for r in refs if r.get("audio_latent") is not None] if kwargs.get("minimax_visual_cond_noise_aug", None) is not None: payload["visual_cond_noise_aug"] = kwargs["minimax_visual_cond_noise_aug"] if kwargs.get("minimax_audio_cond_noise_aug", None) is not None: @@ -2185,7 +2185,7 @@ class MiniMaxH3(BaseModel): payload["layout"] = comfy.ldm.minimax.model.PackedLayout( cross_attn.shape[1], vs[2], (vs[3] + 1) // 2 * 2, (vs[4] + 1) // 2 * 2, latent_shapes[1][-1], keyframes=payload.get("keyframes"), - refs=payload.get("refs"), frame_count=payload.get("frame_count")) + refs=payload.get("refs")) out['minimax_payload'] = comfy.conds.CONDConstant(payload) return out diff --git a/comfy_extras/nodes_minimax_h3.py b/comfy_extras/nodes_minimax_h3.py index 0b1840e85..0a08f185f 100644 --- a/comfy_extras/nodes_minimax_h3.py +++ b/comfy_extras/nodes_minimax_h3.py @@ -20,6 +20,7 @@ import comfy.model_sampling import comfy.nested_tensor import comfy.utils import node_helpers +from comfy.ldm.minimax.model import FRAME_PER_TOKEN, FRAME_RESCALE from comfy_api.latest import ComfyExtension, io CANVAS_MULTIPLE = 32 @@ -67,6 +68,16 @@ def _resize(image, width, height, crop): return samples.movedim(1, -1) +def _encode_ref_audio(audio_vae, audio): + waveform = audio["waveform"] # [B, C, L] + sr = audio["sample_rate"] + vae_sr = getattr(audio_vae, "audio_sample_rate", 32000) + if sr != vae_sr: + waveform = torchaudio.functional.resample(waveform, sr, vae_sr) + z = audio_vae.encode(waveform[:1].movedim(1, -1)) # [1, 32, 2, T] + return z, z.shape[-1] + + def _empty_av_latent(width, height, length, batch_size=1): frame_count, latent_t, audio_t = temporal_shape(length) video = torch.zeros([batch_size, 24, latent_t, height // 16, width // 16], @@ -144,13 +155,87 @@ class MiniMaxH3ImageToVideo(io.ComfyNode): if keyframes: for kf in keyframes: kf["latent"] = vae.encode(kf.pop("image")) - cond = node_helpers.conditioning_set_values(cond, { - "minimax_keyframes": keyframes, - "minimax_frame_count": frame_count, - }) + cond = node_helpers.conditioning_set_values(cond, {"minimax_keyframes": keyframes}) return io.NodeOutput(cond, latent) +class MiniMaxH3AddGuide(io.ComfyNode): + """Anchor image and/or audio guides at an arbitrary pixel frame of the target video.""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="MiniMaxH3AddGuide", + display_name="Add Guide for MiniMax H3", + category="model/conditioning/minimax", + description="Anchor an image, a short clip, audio, or a clip with its soundtrack at any frame of a MiniMax H3 video. Chain several nodes to anchor several frames.", + inputs=[ + io.Conditioning.Input("positive"), + io.Vae.Input("vae", optional=True, tooltip="Video VAE, needed when an image is connected."), + io.Vae.Input("audio_vae", optional=True, tooltip="Audio VAE, needed when an audio is connected."), + io.Latent.Input("latent"), + io.Image.Input("image", optional=True, tooltip="Image or video frames to anchor. Multi-frame batches are anchored as a clip and cropped down to the model's valid clip lengths: 5, 22, 39... (17k + 5) frames. Batches shorter than 5 frames use only the first image."), + io.Audio.Input("audio", optional=True, + tooltip="Soundtrack to anchor starting at the same frame index, cropped to the video's remaining duration."), + io.Int.Input("frame_idx", default=0, min=-9999, max=9999, + tooltip="Frame index to anchor the image or the clip's first frame at. Negative values are counted from the end of the video."), + ], + outputs=[io.Conditioning.Output(display_name="positive")], + ) + + @classmethod + def execute(cls, positive, latent, frame_idx, vae=None, audio_vae=None, image=None, audio=None) -> io.NodeOutput: + samples = latent["samples"] + if not samples.is_nested or len(samples.tensors) != 2 or samples.tensors[0].ndim != 5 or samples.tensors[0].shape[1] != 24: + raise ValueError("MiniMaxH3AddGuide expects a MiniMax H3 AV latent") + if image is None and audio is None: + raise ValueError("MiniMaxH3AddGuide needs an image or an audio to anchor") + video = samples.tensors[0] + height = video.shape[3] * 16 + width = video.shape[4] * 16 + frame_count = sum(FRAME_PER_TOKEN[k % 5] for k in range(video.shape[2])) + + guide_frames = 1 + if image is not None: + if vae is None: + raise ValueError("anchoring guide frames needs the vae input") + guide_frames = image.shape[0] + if guide_frames < 5: + guide_frames = 1 + else: + while guide_frames % 17 != 5: + guide_frames -= 1 + + resolved_frame_index = frame_idx if frame_idx >= 0 else frame_count + frame_idx + if resolved_frame_index < 0 or resolved_frame_index + guide_frames > frame_count: + if guide_frames == 1: + raise ValueError("frame_idx {} is outside the video's {} frames".format(frame_idx, frame_count)) + raise ValueError("a {} frame guide clip at frame_idx {} does not fit in the video's {} frames".format( + guide_frames, frame_idx, frame_count)) + + keyframe = {"resolved_frame_index": resolved_frame_index} + if image is not None: + frames = _resize(image[:guide_frames], width, height, "center") + keyframe["latent"] = vae.encode(frames) + + if audio is not None: + if audio_vae is None: + raise ValueError("anchoring guide audio needs the audio_vae input") + audio_latent, audio_rt = _encode_ref_audio(audio_vae, audio) + # the streams share one time axis: FRAME_RESCALE per pixel frame, 1.0 per audio latent frame + max_rt = math.floor(samples.tensors[1].shape[-1] - FRAME_RESCALE * resolved_frame_index) + if max_rt < 1: + raise ValueError("frame_idx {} is past the end of the video's audio track".format(frame_idx)) + if audio_rt > max_rt: + audio_latent = audio_latent[..., :max_rt].clone() + keyframe["audio_latent"] = audio_latent + + keyframes = list(positive[0][1].get("minimax_keyframes", [])) + keyframes.append(keyframe) + positive = node_helpers.conditioning_set_values(positive, {"minimax_keyframes": keyframes}) + return io.NodeOutput(positive) + + class MiniMaxH3ReferenceToVideo(io.ComfyNode): """ref2va: prompt + reference images / videos / audio -> conditioning + AV latent. @@ -197,16 +282,6 @@ class MiniMaxH3ReferenceToVideo(io.ComfyNode): outputs=[io.Conditioning.Output(display_name="positive"), io.Latent.Output()], ) - @staticmethod - def _encode_ref_audio(audio_vae, audio): - waveform = audio["waveform"] # [B, C, L] - sr = audio["sample_rate"] - vae_sr = getattr(audio_vae, "audio_sample_rate", 32000) - if sr != vae_sr: - waveform = torchaudio.functional.resample(waveform, sr, vae_sr) - z = audio_vae.encode(waveform[:1].movedim(1, -1)) # [1, 32, 2, T] - return z, z.shape[-1] - @classmethod def execute(cls, clip, vae, audio_vae, prompt, width, height, length, ref_image_size="match", ref_images=None, ref_videos=None, ref_video_audios=None, ref_audios=None) -> io.NodeOutput: @@ -254,7 +329,7 @@ class MiniMaxH3ReferenceToVideo(io.ComfyNode): z = vae.encode(frames) audio_latent, ref_audio_t = (None, 0) if soundtrack is not None: - audio_latent, ref_audio_t = cls._encode_ref_audio(audio_vae, soundtrack) + audio_latent, ref_audio_t = _encode_ref_audio(audio_vae, soundtrack) # the soundtrack gets its own