diff --git a/__init__.py b/__init__.py index 7e7b855..8aae7a9 100644 --- a/__init__.py +++ b/__init__.py @@ -7,62 +7,76 @@ __version__ = "1.0.0" __author__ = "marcoags" __package_name__ = "AnotherUtils" +import logging + +# Basic Image Processing from .image_processing.custom_crop import CustomCropNode from .image_processing.smart_resize import SmartResizeNode from .image_processing.nearest_upscale import NearestUpscaleNode from .image_processing.image_grid_slicer import ImageGridSlicer +from .image_processing.remove_alpha import RemoveAlphaNode +from .image_processing.interactive_crop import InteractiveCropNode +from .image_processing.image_composite_masked import AnotherImageCompositeMasked +from .image_processing.segs_adapter import SEGStoBBox, SEGStoSAM2Points, GetFirstFrame, ManualPointToSAM2, RefineMask +from .image_processing.point_collector import PointCollectorSAM2 + +# Loaders from .loaders.load_images import LoadImagesOriginalSize -from .pixel_art.pixel_normalizer import PixelArtNormalizerNode -from .characters.fighting_game_character import FightingGameCharacter -from .characters.walking_pose import WalkingPoseGenerator from .loaders.load_remove_alpha import LoadImageRemoveAlpha +from .loaders.last_image import LastImage +from .loaders.csv_prompt_loader import CSVPromptLoader +from .loaders.trello_prompt_loader import TrelloPromptLoader +from .loaders.trello_browser import TrelloBrowser +from .loaders.caption_image_loader import CaptionImageLoader +from .loaders.load_image_metadata import LoadImageAndExtractPrompt +from .loaders.folder_image_metadata import FolderImageAndExtractPrompt +from .loaders.folder_image_metadata_by_name import FolderImageMetadataByName +from .loaders.load_gif_frames import LoadGifFrames, RemapGifFrames +from .loaders.batch_image_list import BatchToImageList +from .loaders.folder_image_loader import FolderImageLoader + +# Pixel Art +from .pixel_art.pixel_normalizer import PixelArtNormalizerNode from .pixel_art.pixel_art_converter import PixelArtConverterNode from .pixel_art.pixel_art_converter_parallel import PixelArtConverterNodeParallel -from .loaders.last_image import LastImage + +# Characters +from .characters.fighting_game_character import FightingGameCharacter +from .characters.walking_pose import WalkingPoseGenerator from .characters.character_constructor import CharacterConstructor from .characters.character_generator import CharacterRandomizer -from .image_processing.remove_alpha import RemoveAlphaNode + +# GIMP / GEGL Like from .gimp_nodes.adaptive_noise import AdaptiveNoise from .gimp_nodes.cie_lch_noise_gegl_like import CIELChNoiseGEGLLike from .gimp_nodes.image_type_detector import ImageTypeDetector from .gimp_nodes.mean_curvature_blur_gegl_like import MeanCurvatureBlurGEGLLike from .gimp_nodes.rgb_noise_gegl_like import RGBNoiseGEGLLike -from .loaders.csv_prompt_loader import CSVPromptLoader -from .loaders.trello_prompt_loader import TrelloPromptLoader -from .loaders.trello_browser import TrelloBrowser + +# Video General from .video.comparison_swipe import ComparisonSwipeNode from .video.folder_video_concatenator import FolderVideoConcatenator -from .image_processing.interactive_crop import InteractiveCropNode -from .image_processing.image_composite_masked import AnotherImageCompositeMasked -from .loaders.caption_image_loader import CaptionImageLoader +from .video.animated_composite import AnotherTransformKeyframes, AnotherAnimatedCompositeMasked, AnotherTransformOrchestrator +from .video.camera_switcher import AnotherCameraSwitcher +from .video.video_auto_sync_hstack import VideoAutoSyncHStack from .video.video_audio_combiner import ( VideoAudioCombiner, VideoAudioCombinerSimple, HAS_NEW_VIDEO_API, ) -from .video.animated_composite import AnotherTransformKeyframes, AnotherAnimatedCompositeMasked, AnotherTransformOrchestrator -from .video.camera_switcher import AnotherCameraSwitcher -from .loaders.load_image_metadata import LoadImageAndExtractPrompt -from .loaders.folder_image_metadata import FolderImageAndExtractPrompt -from .loaders.folder_image_metadata_by_name import FolderImageMetadataByName -if HAS_NEW_VIDEO_API: - from .video.video_audio_combiner import VideoAudioCombinerV3 +# Audio from .audio.audio_waveform_slicer import AudioWaveformSlicer from .audio.audio_slice_selector import AudioSliceSelector from .audio.audio_concatenate import AudioConcatenate -from .loaders.load_gif_frames import LoadGifFrames, RemapGifFrames -from .loaders.batch_image_list import BatchToImageList -from .video.video_auto_sync_hstack import VideoAutoSyncHStack -from .video.ltxv_multi_guide import LTXVMultiGuide -from .video.ltxv_multi_concat import LTXVMultiConcat -from .video.ltxv_multi_concat_beta import LTXVMultiConcatBeta -from .video.ltxv_vid2vid import LTXVVid2Vid -from .loaders.folder_image_loader import FolderImageLoader + +# Logic & Management +from .logic_management.image_list_to_batch import ImageListToBatch +from .logic_management.indices_list_to_50 import IndicesListTo50 from .logic_management.dataset_loader import DatasetLoader from .logic_management.image_list_sampler import ImageListSampler -from .image_processing.segs_adapter import SEGStoBBox, SEGStoSAM2Points, GetFirstFrame, ManualPointToSAM2, RefineMask -from .image_processing.point_collector import PointCollectorSAM2 + +# Inference from .inference_nodes import ( AnotherLoadYOLO, AnotherLoadSAM2, @@ -79,8 +93,10 @@ from .inference_nodes import ( AnotherMaskMath, AnotherMaskBlur ) + from .core import server_routes # Register Custom API Routes +# Initial Mappings NODE_CLASS_MAPPINGS = { "CustomCrop": CustomCropNode, "SmartResize": SmartResizeNode, @@ -119,10 +135,6 @@ NODE_CLASS_MAPPINGS = { "FolderImageLoader": FolderImageLoader, "DatasetLoader": DatasetLoader, "ImageListSampler": ImageListSampler, - "LTXVMultiGuide": LTXVMultiGuide, - "LTXVMultiConcat": LTXVMultiConcat, - "LTXVMultiConcatBeta": LTXVMultiConcatBeta, - "LTXVVid2Vid": LTXVVid2Vid, "SEGStoBBox": SEGStoBBox, "SEGStoSAM2Points": SEGStoSAM2Points, "GetFirstFrame": GetFirstFrame, @@ -153,12 +165,10 @@ NODE_CLASS_MAPPINGS = { "LoadImageAndExtractPrompt": LoadImageAndExtractPrompt, "FolderImageAndExtractPrompt": FolderImageAndExtractPrompt, "FolderImageMetadataByName": FolderImageMetadataByName, + "ImageListToBatch": ImageListToBatch, + "IndicesListTo50": IndicesListTo50, } -# Add V3 nodes if the new API is available -if HAS_NEW_VIDEO_API: - NODE_CLASS_MAPPINGS["VideoAudioCombinerV3"] = VideoAudioCombinerV3 - NODE_DISPLAY_NAME_MAPPINGS = { "CustomCrop": "Custom Crop", "SmartResize": "Smart Resize with Border Fill", @@ -196,10 +206,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "FolderImageLoader": "Folder Image Loader", "DatasetLoader": "Dataset Loader (Images + Captions)", "ImageListSampler": "Image List Sampler", - "LTXVMultiGuide": "LTXV Multi Guide (N Frames)", - "LTXVMultiConcat": "LTXV Multi Concat (N Frames)", - "LTXVMultiConcatBeta": "LTXV Multi Concat (N Frames) (beta)", - "LTXVVid2Vid": "LTXV Vid2Vid Encode", "SEGStoBBox": "SEGS to BBox", "SEGStoSAM2Points": "SEGS to SAM2 Points (JSON)", "GetFirstFrame": "Get First Frame (Batch to Single)", @@ -230,11 +236,42 @@ NODE_DISPLAY_NAME_MAPPINGS = { "LoadImageAndExtractPrompt": "Load Image and Extract Prompt", "FolderImageAndExtractPrompt": "Folder Image and Extract Prompt", "FolderImageMetadataByName": "Folder Metadata by Node Name", + "ImageListToBatch": "Image List To Multi Batch", + "IndicesListTo 50": "Indices List To 50 Inputs", } -# Add V3 display names if available +# LTX Video Specific - Conditional Loading if HAS_NEW_VIDEO_API: - NODE_DISPLAY_NAME_MAPPINGS["VideoAudioCombinerV3"] = "Video + Audio Combiner (V3)" + try: + from .video.video_audio_combiner import VideoAudioCombinerV3 + from .video.ltxv_multi_guide import LTXVMultiGuide + from .video.another_ltx_sequencer import AnotherLTXSequencer + from .video.ltxv_multi_concat import LTXVMultiConcat + from .video.ltxv_multi_concat_beta import LTXVMultiConcatBeta + from .video.ltxv_vid2vid import LTXVVid2Vid + from .video.ltxv_diagnostic import LTXVDiagnosticNode + + NODE_CLASS_MAPPINGS.update({ + "VideoAudioCombinerV3": VideoAudioCombinerV3, + "LTXVMultiGuide": LTXVMultiGuide, + "AnotherLTXSequencer": AnotherLTXSequencer, + "LTXVMultiConcat": LTXVMultiConcat, + "LTXVMultiConcatBeta": LTXVMultiConcatBeta, + "LTXVVid2Vid": LTXVVid2Vid, + "LTXVDiagnosticNode": LTXVDiagnosticNode, + }) + + NODE_DISPLAY_NAME_MAPPINGS.update({ + "VideoAudioCombinerV3": "Video + Audio Combiner (V3)", + "LTXVMultiGuide": "LTXV Multi Guide (N Frames)", + "AnotherLTXSequencer": "LTX Sequencer (Automated)", + "LTXVMultiConcat": "LTXV Multi Concat (N Frames)", + "LTXVMultiConcatBeta": "LTXV Multi Concat (N Frames) (beta)", + "LTXVVid2Vid": "LTXV Vid2Vid Encode", + "LTXVDiagnosticNode": "LTXV Diagnostic (Shape Checker)", + }) + except Exception as e: + logging.error(f"Failed to load LTX Video extension nodes: {e}") WEB_DIRECTORY = "./js" diff --git a/video/another_ltx_sequencer.py b/video/another_ltx_sequencer.py new file mode 100644 index 0000000..1a4940d --- /dev/null +++ b/video/another_ltx_sequencer.py @@ -0,0 +1,171 @@ +import torch +import logging + +logger = logging.getLogger(__name__) + + +class AnotherLTXSequencer: + """ + Automated version of LTXSequencer (Guide mode). + Takes 'multi_input' (batched images) and 'indices' (list of frame positions). + + IMPORTANT: This node uses LTXVAddGuide.append_keyframe which EXTENDS the video + latent along the time dimension. For LTX 2.3 AV workflows, this node should be + placed BEFORE LTXVConcatAVLatent (i.e., it should receive a plain video latent, + NOT a NestedTensor). The audio latent is handled separately by ConcatAVLatent. + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING",), + "vae": ("VAE",), + "latent": ("LATENT",), + "multi_input": ("IMAGE",), + "indices": ("INT",), + "num_images": ("INT", {"default": 1, "min": 0, "max": 50, "step": 1, "tooltip": "Number of images to process from the batch."}), + "insert_mode": (["frames", "seconds"], {"default": "frames", "tooltip": "Select the method for determining insertion points."}), + "frame_rate": ("INT", {"default": 24, "min": 1, "max": 120, "step": 1, "tooltip": "Video FPS (used for calculating second insertions)."}), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Global strength for all guide images."}), + } + } + + INPUT_IS_LIST = True + RETURN_TYPES = ("CONDITIONING", "CONDITIONING", "LATENT") + RETURN_NAMES = ("positive", "negative", "latent") + FUNCTION = "execute" + CATEGORY = "AnotherUtils/video" + + def execute(self, positive, negative, vae, latent, multi_input, indices, num_images, insert_mode, frame_rate, strength): + from comfy_extras.nodes_lt import LTXVAddGuide, get_noise_mask, _append_guide_attention_entry + + # Unwrap list inputs because INPUT_IS_LIST = True + positive = positive[0] if isinstance(positive, list) else positive + negative = negative[0] if isinstance(negative, list) else negative + vae = vae[0] if isinstance(vae, list) else vae + latent = latent[0] if isinstance(latent, list) else latent + multi_input = multi_input[0] if isinstance(multi_input, list) else multi_input + num_images = num_images[0] if isinstance(num_images, list) else num_images + insert_mode = insert_mode[0] if isinstance(insert_mode, list) else insert_mode + frame_rate = frame_rate[0] if isinstance(frame_rate, list) else frame_rate + strength = strength[0] if isinstance(strength, list) else strength + + # Keep indices as a list + idx_list = indices if isinstance(indices, list) else [indices] + if len(idx_list) > 0 and isinstance(idx_list[0], list): + idx_list = idx_list[0] + + scale_factors = vae.downscale_index_formula + + # --- Extract video latent, separating from audio if nested --- + try: + from comfy.nested_tensor import NestedTensor + is_nested = isinstance(latent["samples"], NestedTensor) + except ImportError: + is_nested = False + + if is_nested: + # Extract video-only portion for guide processing + latent_samples = latent["samples"].tensors[0].clone() + audio_samples = latent["samples"].tensors[1] + + # Extract noise masks + noise_mask_raw = latent.get("noise_mask", None) + is_mask_nested = isinstance(noise_mask_raw, NestedTensor) if noise_mask_raw is not None else False + + if is_mask_nested: + noise_mask = noise_mask_raw.tensors[0].clone() + audio_mask = noise_mask_raw.tensors[1] + else: + # Create default video noise mask + batch, _, lat_t, _, _ = latent_samples.shape + noise_mask = torch.ones((batch, 1, lat_t, 1, 1), dtype=torch.float32, device=latent_samples.device) + audio_mask = None + + logger.info(f"[AnotherLTXSequencer] Nested input - Video: {latent_samples.shape}, Audio: {audio_samples.shape}") + else: + latent_samples = latent["samples"].clone() + noise_mask = get_noise_mask(latent).clone() + audio_samples = None + audio_mask = None + is_mask_nested = False + + _, _, latent_length, latent_height, latent_width = latent_samples.shape + batch_size = multi_input.shape[0] if multi_input is not None else 0 + + cur_pos = positive + cur_neg = negative + + # Process guide images + for i in range(1, num_images + 1): + if i > batch_size: + continue + + img = multi_input[i-1:i] + if img is None: + continue + + if (i-1) >= len(idx_list): + continue + + # Calculate frame index + f_idx = idx_list[i-1] + if insert_mode == "seconds": + f_idx = int(f_idx * frame_rate) + + # Encode and get latent index + image_1, t = LTXVAddGuide.encode(vae, latent_width, latent_height, img, scale_factors) + frame_idx, latent_idx = LTXVAddGuide.get_latent_index(cur_pos, latent_length, len(image_1), f_idx, scale_factors) + + delta_t = t.shape[2] + + if latent_idx + delta_t > latent_length: + logger.warning(f"[AnotherLTXSequencer] Guide at frame {f_idx} exceeds latent length. Skipping.") + continue + + # append_keyframe concatenates guide frames onto the video latent (dim=2) + cur_pos, cur_neg, latent_samples, noise_mask = LTXVAddGuide.append_keyframe( + cur_pos, cur_neg, + frame_idx, + latent_samples, + noise_mask, + t, + strength, + scale_factors, + ) + + # Add attention entry for the guide + try: + pre_filter_count = t.shape[2] * t.shape[3] * t.shape[4] + guide_latent_shape = list(t.shape[2:]) + cur_pos, cur_neg = _append_guide_attention_entry( + cur_pos, cur_neg, pre_filter_count, guide_latent_shape, strength=strength + ) + except Exception as e: + logger.error(f"[AnotherLTXSequencer] Failed to append attention entry: {e}") + + # --- Build output latent --- + new_latent = latent.copy() + + if is_nested: + # Re-wrap as NestedTensor. The audio is NOT padded because + # the sampler's pack/unpack handles mismatched time dimensions. + # CropGuides will trim the video back after sampling. + from comfy.nested_tensor import NestedTensor + new_latent["samples"] = NestedTensor((latent_samples, audio_samples)) + if is_mask_nested and audio_mask is not None: + new_latent["noise_mask"] = NestedTensor((noise_mask, audio_mask)) + else: + new_latent["noise_mask"] = noise_mask + logger.info(f"[AnotherLTXSequencer] Nested output - Video: {latent_samples.shape}, Audio: {audio_samples.shape}") + else: + new_latent["samples"] = latent_samples + new_latent["noise_mask"] = noise_mask + + return (cur_pos, cur_neg, new_latent) + + + + diff --git a/video/ltxv_diagnostic.py b/video/ltxv_diagnostic.py new file mode 100644 index 0000000..7087d76 --- /dev/null +++ b/video/ltxv_diagnostic.py @@ -0,0 +1,85 @@ +import torch +import logging + +logger = logging.getLogger(__name__) + +class LTXVDiagnosticNode: + """ + Diagnostic node to inspect LTX Video latent structures (NestedTensors, shapes, etc.) + Place this before Audio Decode or Sampler to see what's happening inside. + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "latent": ("LATENT",), + "node_name": ("STRING", {"default": "LTXV Diagnostic"}), + "print_to_console": ("BOOLEAN", {"default": True}), + } + } + + RETURN_TYPES = ("LATENT", "STRING") + RETURN_NAMES = ("latent", "info") + FUNCTION = "inspect" + CATEGORY = "AnotherUtils/video" + + def inspect(self, latent, node_name, print_to_console): + info_lines = [f"--- {node_name} Diagnostic ---"] + + samples = latent.get("samples", None) + if samples is None: + info_lines.append("ERROR: No 'samples' found in latent.") + else: + info_lines.append(f"Samples Type: {type(samples)}") + + # Check for ComfyUI NestedTensor wrapper + is_wrapper = False + try: + from comfy.nested_tensor import NestedTensor as ComfyNested + if isinstance(samples, ComfyNested): + is_wrapper = True + info_lines.append("Format: ComfyUI NestedTensor Wrapper") + for i, t in enumerate(samples.tensors): + info_lines.append(f" Inner Tensor {i} Shape: {list(t.shape)} | dtype: {t.dtype}") + except ImportError: + pass + + # Check for Native PyTorch Nested Tensor + if not is_wrapper and hasattr(samples, "is_nested"): + if samples.is_nested: + info_lines.append("Format: Native PyTorch Nested Tensor") + try: + # Try to unbind to see components + unbound = samples.unbind() + for i, t in enumerate(unbound): + info_lines.append(f" Unbound Component {i} Shape: {list(t.shape)} | dtype: {t.dtype}") + except Exception as e: + info_lines.append(f" Error unbinding native nested tensor: {e}") + else: + info_lines.append(f"Format: Standard Tensor | Shape: {list(samples.shape)} | dtype: {samples.dtype}") + elif not is_wrapper: + info_lines.append(f"Format: Standard Tensor | Shape: {list(samples.shape)} | dtype: {samples.dtype}") + + # Check for noise_mask + mask = latent.get("noise_mask", None) + if mask is not None: + info_lines.append(f"Noise Mask Type: {type(mask)}") + if hasattr(mask, "shape"): + info_lines.append(f" Mask Shape: {list(mask.shape)}") + elif is_wrapper: + info_lines.append(" Mask is likely also a Nested Wrapper (checking components...)") + # Many LTX implementations wrap mask too + else: + info_lines.append("Noise Mask: Not present") + + # Other keys + other_keys = [k for k in latent.keys() if k not in ["samples", "noise_mask"]] + if other_keys: + info_lines.append(f"Other Keys: {other_keys}") + + info_text = "\n".join(info_lines) + if print_to_console: + print(info_text) + + return (latent, info_text) diff --git a/video/ltxv_multi_concat.py b/video/ltxv_multi_concat.py index e5cc1da..3e6c25f 100644 --- a/video/ltxv_multi_concat.py +++ b/video/ltxv_multi_concat.py @@ -1,5 +1,7 @@ import torch from comfy_extras.nodes_lt import LTXVAddGuide, get_noise_mask +from typing import Dict, Tuple, Any +import logging from .ltxv_utils import resolve_frame_indices, parse_strengths, flatten_images class LTXVMultiConcat: @@ -83,7 +85,20 @@ class LTXVMultiConcat: scale_factors = vae.downscale_index_formula time_scale_factor = scale_factors[0] - latent_samples = latent["samples"].clone() + + # Handle NestedTensor from audio-enabled latents + try: + from comfy.nested_tensor import NestedTensor + is_nested = isinstance(latent["samples"], NestedTensor) + except ImportError: + is_nested = False + + if is_nested: + latent_samples = latent["samples"].tensors[0].clone() + audio_samples = latent["samples"].tensors[1] + else: + latent_samples = latent["samples"].clone() + latent_length = latent_samples.shape[2] total_pixel_frames = (latent_length - 1) * time_scale_factor + 1 @@ -94,7 +109,17 @@ class LTXVMultiConcat: n = min(num_images, len(frame_indices)) - noise_mask = get_noise_mask(latent).clone() + has_nested_mask = False + noise_mask_obj = get_noise_mask(latent) + if is_nested: + from comfy.nested_tensor import NestedTensor + if isinstance(noise_mask_obj, NestedTensor): + noise_mask = noise_mask_obj.tensors[0].clone() + else: + noise_mask = noise_mask_obj.clone() + else: + noise_mask = noise_mask_obj.clone() + _, _, lat_len, lat_h, lat_w = latent_samples.shape for i in range(n): @@ -118,6 +143,23 @@ class LTXVMultiConcat: # noise_mask is [B, 1, T, 1, 1] typically noise_mask[:, :, latent_idx:latent_idx+t.shape[2], :, :] = 1.0 - strength + new_latent = latent.copy() + + if is_nested: + from comfy.nested_tensor import NestedTensor + new_latent["samples"] = NestedTensor((latent_samples, audio_samples)) + # Re-wrap noise_mask as NestedTensor to match samples structure. + # The audio mask stays unchanged (all ones = fully denoised). + noise_mask_raw = latent.get("noise_mask", None) + if noise_mask_raw is not None and isinstance(noise_mask_raw, NestedTensor): + audio_noise_mask = noise_mask_raw.tensors[1] + else: + audio_noise_mask = torch.ones_like(audio_samples) + new_latent["noise_mask"] = NestedTensor((noise_mask, audio_noise_mask)) + else: + new_latent["samples"] = latent_samples + new_latent["noise_mask"] = noise_mask + # In Concat mode, we DON'T modify conditioning with attention entries. # We just return the modified latent with frames and mask. - return (positive, negative, {"samples": latent_samples, "noise_mask": noise_mask}) + return (positive, negative, new_latent) diff --git a/video/ltxv_multi_concat_beta.py b/video/ltxv_multi_concat_beta.py index b17f8cc..0e04c40 100644 --- a/video/ltxv_multi_concat_beta.py +++ b/video/ltxv_multi_concat_beta.py @@ -1,5 +1,7 @@ import torch from comfy_extras.nodes_lt import LTXVAddGuide, get_noise_mask +from typing import Dict, Tuple, Any +import logging from .ltxv_utils import resolve_frame_indices, parse_strengths, flatten_images class LTXVMultiConcatBeta: @@ -83,7 +85,20 @@ class LTXVMultiConcatBeta: scale_factors = vae.downscale_index_formula time_scale_factor = scale_factors[0] - latent_samples = latent["samples"].clone() + + # Handle NestedTensor from audio-enabled latents + try: + from comfy.nested_tensor import NestedTensor + is_nested = isinstance(latent["samples"], NestedTensor) + except ImportError: + is_nested = False + + if is_nested: + latent_samples = latent["samples"].tensors[0].clone() + audio_samples = latent["samples"].tensors[1] + else: + latent_samples = latent["samples"].clone() + latent_length = latent_samples.shape[2] total_pixel_frames = (latent_length - 1) * time_scale_factor + 1 @@ -93,9 +108,22 @@ class LTXVMultiConcatBeta: strength_list = parse_strengths(strengths_str, num_images, strength) n = min(num_images, len(frame_indices)) - noise_mask = get_noise_mask(latent).clone() + + has_nested_mask = False + noise_mask_obj = get_noise_mask(latent) + if is_nested: + from comfy.nested_tensor import NestedTensor + if isinstance(noise_mask_obj, NestedTensor): + noise_mask = noise_mask_obj.tensors[0].clone() + else: + noise_mask = noise_mask_obj.clone() + else: + noise_mask = noise_mask_obj.clone() + _, _, lat_len, lat_h, lat_w = latent_samples.shape + new_latent = latent.copy() + if not smooth_strength: # Original behavior for i in range(n): @@ -122,7 +150,14 @@ class LTXVMultiConcatBeta: anchors[latent_idx] = (t, strength_list[i]) if not anchors: - return (positive, negative, {"samples": latent_samples, "noise_mask": noise_mask}) + if is_nested: + from comfy.nested_tensor import NestedTensor + new_latent["samples"] = NestedTensor((latent_samples, audio_samples)) + new_latent["noise_mask"] = noise_mask + else: + new_latent["samples"] = latent_samples + new_latent["noise_mask"] = noise_mask + return (positive, negative, new_latent) sorted_idx = sorted(anchors.keys()) @@ -182,4 +217,18 @@ class LTXVMultiConcatBeta: latent_samples[:, :, i:i+1, :, :] = t_interp noise_mask[:, :, i:i+1, :, :] = 1.0 - s_interp - return (positive, negative, {"samples": latent_samples, "noise_mask": noise_mask}) + if is_nested: + from comfy.nested_tensor import NestedTensor + new_latent["samples"] = NestedTensor((latent_samples, audio_samples)) + # Re-wrap noise_mask as NestedTensor to match samples structure. + noise_mask_raw = latent.get("noise_mask", None) + if noise_mask_raw is not None and isinstance(noise_mask_raw, NestedTensor): + audio_noise_mask = noise_mask_raw.tensors[1] + else: + audio_noise_mask = torch.ones_like(audio_samples) + new_latent["noise_mask"] = NestedTensor((noise_mask, audio_noise_mask)) + else: + new_latent["samples"] = latent_samples + new_latent["noise_mask"] = noise_mask + + return (positive, negative, new_latent)