from __future__ import annotations import logging from typing import Any import torch LOGGER = logging.getLogger("easy_media.h3_motion_context") FRAME_PER_TOKEN = (1, 4, 4, 4, 4) FPS = 24 FRAME_RESCALE = 5.0 / 3.0 AUDIO_HZ = 40.0 VIDEO_RUN_GRID = (124, 107, 90, 73, 56, 39, 22, 5, 1) CONTEXT_SWAP_NOISE_ALPHA = 0.45 CONTEXT_SWAP_NOISE_ALPHA_END = 0.10 CONTEXT_SWAP_NOISE_RAMP_STEPS = 2 REENCODED_ANCHOR_FRAMES = 5 def _pixel_frames(latent_steps: int) -> int: return sum(FRAME_PER_TOKEN[index % 5] for index in range(latent_steps)) def _step_offsets(latent_steps: int) -> list[int]: offsets: list[int] = [] frame = 0 for index in range(latent_steps): offsets.append(frame) frame += FRAME_PER_TOKEN[index % 5] return offsets def _streams_from_latent(latent: dict[str, Any]) -> list[torch.Tensor]: if not isinstance(latent, dict) or "samples" not in latent: raise ValueError("easy h3 motion context: expected an H3 AV latent") samples = latent["samples"] if isinstance(samples, (tuple, list)): parts = list(samples) elif isinstance(getattr(samples, "tensors", None), (tuple, list)): parts = list(samples.tensors) elif getattr(samples, "is_nested", False) and hasattr(samples, "unbind"): parts = list(samples.unbind()) else: raise ValueError( "easy h3 motion context: expected a nested video/audio latent, " f"got {type(samples)!r}" ) if not parts: raise ValueError("easy h3 motion context: AV latent contains no streams") return parts def _official_nested_tensor(parts: tuple[torch.Tensor, ...]) -> Any: """Rebuild H3 AV data with ComfyUI's supported NestedTensor wrapper.""" try: import comfy.nested_tensor except ImportError as error: raise RuntimeError( "easy h3 motion context: ComfyUI nested tensor support is unavailable" ) from error return comfy.nested_tensor.NestedTensor(parts) def _noise_mask_streams( latent: dict[str, Any], ) -> tuple[torch.Tensor | None, torch.Tensor | None]: """Return optional video/audio masks, including legacy video-only masks.""" mask = latent.get("noise_mask") if mask is None: return None, None if isinstance(mask, torch.Tensor): return mask, None if isinstance(mask, (tuple, list)): parts = tuple(mask) elif isinstance(getattr(mask, "tensors", None), (tuple, list)): parts = tuple(mask.tensors) elif getattr(mask, "is_nested", False) and hasattr(mask, "unbind"): parts = tuple(mask.unbind()) else: raise ValueError( "easy h3 context: noise_mask must be a tensor or nested streams" ) if not parts or len(parts) > 2 or not all( isinstance(part, torch.Tensor) for part in parts ): raise ValueError("easy h3 context: noise_mask contains invalid streams") return parts[0], parts[1] if len(parts) > 1 else None def _merge_noise_mask( generated: torch.Tensor, existing: torch.Tensor | None, stream_name: str, ) -> torch.Tensor: """Preserve existing locks while adding a context release mask.""" if existing is None: return generated existing = existing.to(device=generated.device, dtype=generated.dtype) try: return torch.minimum(generated, existing).contiguous() except RuntimeError as error: raise ValueError( f"easy h3 context: existing {stream_name} noise_mask shape " f"{tuple(existing.shape)} is incompatible with {tuple(generated.shape)}" ) from error def _video_from_latent(latent: dict[str, Any]) -> torch.Tensor: if not isinstance(latent, dict) or "samples" not in latent: raise ValueError("easy h3 motion context: expected a video latent") samples = latent["samples"] # A directly VAE-encoded anchor is a normal ComfyUI video LATENT, while # sampled H3 data stores video and audio in a nested tensor. if isinstance(samples, torch.Tensor) and not getattr(samples, "is_nested", False): video = samples else: video = _streams_from_latent(latent)[0] if video.ndim == 4: video = video.unsqueeze(0) if video.ndim != 5: raise ValueError( "easy h3 motion context: expected video latent [B,C,T,H,W], " f"got shape {tuple(video.shape)}" ) return video def _steps_for_frames(frame_count: int) -> int | None: steps = 0 covered = 0 while covered < frame_count: covered += FRAME_PER_TOKEN[steps % 5] steps += 1 return steps if covered == frame_count else None def _video_tail_from_latent( latent: dict[str, Any], frame_count: int ) -> tuple[list[torch.Tensor], list[int], int]: video, start, steps, covered = _video_tail_bounds(latent, frame_count) blocks = [ video[:1, :, start + index : start + index + 1].clone() for index in range(steps) ] return blocks, _step_offsets(steps), covered def _video_tail_bounds( latent: dict[str, Any], frame_count: int ) -> tuple[torch.Tensor, int, int, int]: """Validate an H3 tail and return its source tensor and temporal bounds.""" video = _video_from_latent(latent) total_steps = int(video.shape[2]) steps = _steps_for_frames(frame_count) if steps is None: raise ValueError( "easy h3 motion context: context length cannot be represented by " "whole H3 latent steps" ) if steps > total_steps: raise ValueError( f"easy h3 motion context: requested {steps} latent steps, but the " f"context latent contains {total_steps}" ) start = total_steps - steps if start % 5 != 0: raise RuntimeError( "easy h3 motion context: the context tail starts at an invalid " "H3 temporal cycle position" ) covered = _pixel_frames(steps) return video, start, steps, covered def _audio_tail_from_latent( latent: dict[str, Any], audio_frames: int ) -> tuple[torch.Tensor, int, float]: parts = _streams_from_latent(latent) if len(parts) < 2: raise ValueError( "easy h3 motion context: context_latent has no audio stream" ) video, audio = parts[0], parts[1] if video.ndim == 4: video = video.unsqueeze(0) if audio.ndim == 3: audio = audio.unsqueeze(0) if audio.ndim != 4: raise ValueError( "easy h3 motion context: expected audio latent [B,C,2,T], " f"got shape {tuple(audio.shape)}" ) total_steps = int(audio.shape[-1]) video_frames = _pixel_frames(int(video.shape[2])) overhang = total_steps - FRAME_RESCALE * video_frames if not -0.5 < overhang < 0.5: LOGGER.warning( "Unexpected H3 audio grid (%d steps for %d frames); ignoring overhang", total_steps, video_frames, ) overhang = 0.0 requested_steps = int(round(audio_frames / FPS * AUDIO_HZ)) if requested_steps > total_steps: LOGGER.warning( "Requested %d audio steps but only %d are available", requested_steps, total_steps, ) requested_steps = total_steps if requested_steps < 1: raise ValueError("easy h3 motion context: audio context window is empty") return ( audio[:1, ..., total_steps - requested_steps :].clone(), requested_steps, float(overhang), ) def trim_motion_context_latent( latent: dict[str, Any], context_length: int | str = "22", ) -> dict[str, Any]: """Copy only the H3 AV tail needed by the next context segment to CPU.""" frame_count = int(context_length) video, start, steps, covered = _video_tail_bounds(latent, frame_count) audio, _, _ = _audio_tail_from_latent(latent, covered) video = video[:1, :, start : start + steps].detach() streams = tuple( stream.detach().to(device="cpu", copy=True).contiguous() for stream in (video, audio) ) output = {"samples": _official_nested_tensor(streams)} anchor_samples = latent.get("anchor_samples") if anchor_samples is not None: if not isinstance(anchor_samples, torch.Tensor) or anchor_samples.ndim != 5: raise ValueError( "easy h3 motion context: anchor_samples must be a video latent" ) output["anchor_samples"] = ( anchor_samples.detach().to(device="cpu", copy=True).contiguous() ) return output def apply_context_swap_noise( latent: dict[str, Any], context_length: int | str = "22", seed: int = 0, alpha: float = CONTEXT_SWAP_NOISE_ALPHA, alpha_end: float = CONTEXT_SWAP_NOISE_ALPHA_END, ramp_steps: int = CONTEXT_SWAP_NOISE_RAMP_STEPS, ) -> dict[str, Any]: """Return a disposable context copy with tapered video-latent noise.""" alpha = float(alpha) alpha_end = float(alpha_end) if not 0.0 <= alpha_end <= alpha <= 1.0: raise ValueError( "easy h3 context swap: expected 0 <= alpha_end <= alpha <= 1" ) parts = _streams_from_latent(latent) source_video, start, steps, covered = _video_tail_bounds( latent, int(context_length), ) ramp = min(max(1, int(ramp_steps)), steps) output_video = source_video.clone() tail = output_video[:, :, start : start + steps] scale = float(tail.detach().float().std().item()) if scale != scale or scale in {float("inf"), float("-inf")} or scale <= 0.0: LOGGER.warning( "Context swap latent has a degenerate standard deviation; using unit noise" ) scale = 1.0 generator = torch.Generator(device="cpu").manual_seed( int(seed) & 0xFFFFFFFFFFFFFFFF ) schedule: list[float] = [] for index in range(steps): from_end = steps - 1 - index amount = ( alpha if from_end >= ramp else alpha + (alpha_end - alpha) * (ramp - from_end) / ramp ) schedule.append(round(amount, 4)) block = output_video[:, :, start + index] noise = torch.randn( tuple(block.shape), generator=generator, dtype=torch.float32, ).mul_(scale) noise = noise.to(device=block.device, dtype=block.dtype) output_video[:, :, start + index] = ( block * (1.0 - amount) + noise * amount ) if parts[0].ndim == 4: output_video = output_video[0] output_parts = list(parts) output_parts[0] = output_video output = latent.copy() output["samples"] = _official_nested_tensor(tuple(output_parts)) LOGGER.info( "Context swap noise: video=%d steps/%d frames alpha=%s; audio untouched", steps, covered, schedule, ) return output def _resize_frames( images: torch.Tensor, width: int, height: int ) -> torch.Tensor: try: import comfy.utils except ImportError as error: raise RuntimeError("ComfyUI image resize utilities are unavailable") from error samples = images[..., :3].movedim(-1, 1) samples = comfy.utils.common_upscale( samples, width, height, "lanczos", "disabled", ) return samples.movedim(1, -1) def _encode_tail_audio( audio_vae: Any, audio: dict[str, Any], seconds: float, ) -> tuple[torch.Tensor, int]: waveform = audio.get("waveform") sample_rate = int(audio.get("sample_rate", 0)) if not isinstance(waveform, torch.Tensor) or sample_rate <= 0: raise ValueError("easy h3 motion context: invalid context_audio") target_rate = int(getattr(audio_vae, "audio_sample_rate", 32000)) if sample_rate != target_rate: try: import torchaudio except ImportError as error: raise RuntimeError( "torchaudio is required to resample H3 context audio" ) from error waveform = torchaudio.functional.resample( waveform, sample_rate, target_rate, ) wanted = int(round(seconds * target_rate)) available = int(waveform.shape[-1]) if available > wanted: waveform = waveform[..., available - wanted :] encoded = audio_vae.encode(waveform[:1].movedim(1, -1)) return encoded, int(encoded.shape[-1]) def _release_mask_inside_prefix( mask: torch.Tensor, prefix_steps: int, transition_steps: int, ) -> tuple[int, list[float]]: """Lock a copied prefix, then release it before generated latent begins.""" prefix_steps = int(prefix_steps) ramp_steps = max(0, min(int(transition_steps), prefix_steps)) locked_steps = prefix_steps - ramp_steps if locked_steps > 0: mask[..., :locked_steps] = 0.0 ramp_values: list[float] = [] if ramp_steps > 0: values = torch.arange( 1, ramp_steps + 1, device=mask.device, dtype=mask.dtype, ) / float(ramp_steps + 1) shape = [1] * mask.ndim shape[-1] = ramp_steps mask[..., locked_steps:prefix_steps] = values.view(*shape) ramp_values = [round(float(value), 4) for value in values.detach().cpu()] return locked_steps, ramp_values def _native_keyframe_stats(conditioning: Any) -> tuple[int, int]: """Count 0.4-style video/audio keyframes for compatibility checks.""" max_video = 0 max_audio = 0 for _embedding, metadata in conditioning: video = 0 audio = 0 for keyframe in metadata.get("minimax_keyframes") or []: if keyframe.get("latent") is not None: video += 1 if keyframe.get("audio_latent") is not None: audio += 1 max_video = max(max_video, video) max_audio = max(max_audio, audio) return max_video, max_audio def _hard_av_latent( latent: dict[str, Any], context_latent: dict[str, Any], video_frames: int, video_transition_steps: int, audio_transition_steps: int, ) -> dict[str, Any]: """Copy previous video/audio tails into the current H3 sampling seed.""" existing_video_mask, existing_audio_mask = _noise_mask_streams(latent) target_parts = _streams_from_latent(latent) if len(target_parts) < 2: raise ValueError("easy h3 hard context: target latent has no audio stream") target_video, target_audio = target_parts[:2] if target_video.ndim == 4: target_video = target_video.unsqueeze(0) if target_audio.ndim == 3: target_audio = target_audio.unsqueeze(0) if ( target_video.ndim != 5 or target_audio.ndim != 4 or int(target_audio.shape[2]) != 2 ): raise ValueError("easy h3 hard context: expected H3 video/audio latent streams") if int(target_video.shape[0]) != int(target_audio.shape[0]): raise ValueError( "easy h3 hard context: target video/audio batch sizes differ" ) video_blocks, _, covered = _video_tail_from_latent( context_latent, int(video_frames), ) copied_video = torch.cat(video_blocks, dim=2) video_steps = int(copied_video.shape[2]) if covered != int(video_frames): raise RuntimeError("easy h3 hard context: video context span changed") if video_steps < 1 or video_steps >= int(target_video.shape[2]): raise ValueError( "easy h3 hard context: copied video prefix must be shorter than target" ) if ( copied_video.shape[0] != target_video.shape[0] or copied_video.shape[1] != target_video.shape[1] or copied_video.shape[3:] != target_video.shape[3:] ): raise ValueError( "easy h3 hard context: context and target video latent shapes differ" ) output_video = target_video.clone() output_video[:, :, :video_steps] = copied_video.to( device=output_video.device, dtype=output_video.dtype, ) video_mask = torch.ones_like(output_video[:, :1], dtype=torch.float32) temporal_mask = video_mask.permute(0, 1, 3, 4, 2) video_locked, video_ramp = _release_mask_inside_prefix( temporal_mask, video_steps, video_transition_steps, ) video_mask = temporal_mask.permute(0, 1, 4, 2, 3).contiguous() copied_audio, audio_steps, overhang = _audio_tail_from_latent( context_latent, int(video_frames), ) if audio_steps < 1 or audio_steps >= int(target_audio.shape[-1]): raise ValueError( "easy h3 hard context: copied audio prefix must be shorter than target" ) if copied_audio.shape[:3] != target_audio.shape[:3]: raise ValueError( "easy h3 hard context: context and target audio latent shapes differ" ) output_audio = target_audio.clone() output_audio[..., :audio_steps] = copied_audio.to( device=output_audio.device, dtype=output_audio.dtype, ) audio_mask = torch.ones_like(output_audio[:, :1], dtype=torch.float32) audio_locked, audio_ramp = _release_mask_inside_prefix( audio_mask, audio_steps, audio_transition_steps, ) video_mask = _merge_noise_mask( video_mask, existing_video_mask, "video", ) audio_mask = _merge_noise_mask( audio_mask, existing_audio_mask, "audio", ) output = latent.copy() output["samples"] = _official_nested_tensor((output_video, output_audio)) output["noise_mask"] = _official_nested_tensor((video_mask, audio_mask)) LOGGER.info( "Hard AV context: video=%d steps/%d frames (lock=%d ramp=%s), " "audio=%d steps/%d frames (overhang=%.3f lock=%d ramp=%s)", video_steps, covered, video_locked, video_ramp or "off", audio_steps, video_frames, overhang, audio_locked, audio_ramp or "off", ) return output def _reencoded_anchor_av_latent( latent: dict[str, Any], anchor_latent: dict[str, Any], context_frames: int, anchor_frames: int = REENCODED_ANCHOR_FRAMES, ) -> dict[str, Any]: """Hard-pin a short re-encoded video anchor at the end of a soft prefix.""" context_steps = _steps_for_frames(int(context_frames)) anchor_video_steps = _steps_for_frames(int(anchor_frames)) if context_steps is None or anchor_video_steps is None: raise ValueError( "easy h3 anchor context: context and anchor lengths must align " "to whole H3 latent steps" ) anchor_start = context_steps - anchor_video_steps if anchor_start < 0 or anchor_start % len(FRAME_PER_TOKEN) != 0: raise ValueError( "easy h3 anchor context: anchor must begin on an H3 temporal cycle" ) existing_video_mask, existing_audio_mask = _noise_mask_streams(latent) target_parts = _streams_from_latent(latent) if len(target_parts) < 2: raise ValueError("easy h3 anchor context: target latent has no audio stream") target_video, target_audio = target_parts[:2] if target_video.ndim == 4: target_video = target_video.unsqueeze(0) if target_audio.ndim == 3: target_audio = target_audio.unsqueeze(0) if target_video.ndim != 5 or target_audio.ndim != 4: raise ValueError("easy h3 anchor context: invalid target AV latent dimensions") if context_steps >= int(target_video.shape[2]): raise ValueError( "easy h3 anchor context: context prefix must be shorter than target" ) blocks, _, covered = _video_tail_from_latent(anchor_latent, int(anchor_frames)) anchor_video = torch.cat(blocks, dim=2) if covered != int(anchor_frames): raise RuntimeError("easy h3 anchor context: video anchor span changed") if ( anchor_video.shape[0] != target_video.shape[0] or anchor_video.shape[1] != target_video.shape[1] or anchor_video.shape[3:] != target_video.shape[3:] ): raise ValueError( "easy h3 anchor context: anchor and target video latent shapes differ" ) output_video = target_video.clone() output_video[:, :, anchor_start:context_steps] = anchor_video.to( device=output_video.device, dtype=output_video.dtype, ) video_mask = torch.ones_like(output_video[:, :1], dtype=torch.float32) video_mask[:, :, anchor_start:context_steps] = 0.0 output_audio = target_audio audio_mask = torch.ones_like(output_audio[:, :1], dtype=torch.float32) output = latent.copy() output["samples"] = _official_nested_tensor((output_video, output_audio)) output["noise_mask"] = _official_nested_tensor( ( _merge_noise_mask(video_mask, existing_video_mask, "video"), _merge_noise_mask(audio_mask, existing_audio_mask, "audio"), ) ) LOGGER.info( "Re-encoded video anchor: soft context=%d frames, hard anchor=%d frames " "(video steps=%d:%d); audio remains generative", context_frames, anchor_frames, anchor_start, context_steps, ) return output def _video_anchor_from_context_latent( context_latent: dict[str, Any], anchor_frames: int = REENCODED_ANCHOR_FRAMES, ) -> dict[str, torch.Tensor]: """Use an embedded pixel anchor or the matching context-video tail.""" anchor_samples = context_latent.get("anchor_samples") if anchor_samples is not None: if not isinstance(anchor_samples, torch.Tensor) or anchor_samples.ndim != 5: raise ValueError( "easy h3 anchor context: anchor_samples must be [B,C,T,H,W]" ) return {"samples": anchor_samples} blocks, _, covered = _video_tail_from_latent(context_latent, int(anchor_frames)) if covered != int(anchor_frames): raise RuntimeError("easy h3 anchor context: video anchor span changed") return {"samples": torch.cat(blocks, dim=2)} def apply_reencoded_anchor_motion_context( conditioning: Any, vae: Any, latent: dict[str, Any], context_latent: dict[str, Any], context_length: int | str = "22", anchor_length: int | str = str(REENCODED_ANCHOR_FRAMES), ) -> tuple[Any, int, dict[str, Any]]: """Use long native Motion Context with a short immutable video anchor.""" output, trim_frames = apply_motion_context( conditioning=conditioning, vae=vae, latent=latent, context_length=context_length, audio_context_length=0, context_latent=context_latent, ) video_keyframes, audio_keyframes = _native_keyframe_stats(output) if video_keyframes < 1 or audio_keyframes < 1: raise RuntimeError( "easy h3 anchor context requires Motion Context 0.4.0+ native " "video/audio keyframes and ComfyUI 0.34.0+; got " f"video_keyframes={video_keyframes}, audio_keyframes={audio_keyframes}" ) anchored = _reencoded_anchor_av_latent( latent, _video_anchor_from_context_latent(context_latent, int(anchor_length)), context_frames=int(trim_frames), anchor_frames=int(anchor_length), ) return output, trim_frames, anchored def apply_hard_motion_context( conditioning: Any, vae: Any, latent: dict[str, Any], context_latent: dict[str, Any], context_length: int | str = "22", video_transition_steps: int = 4, audio_transition_steps: int = 4, ) -> tuple[Any, int, dict[str, Any]]: """Apply normal H3 layout conditioning plus hard AV latent continuity.""" output, trim_frames = apply_motion_context( conditioning=conditioning, vae=vae, latent=latent, context_length=context_length, audio_context_length=0, context_latent=context_latent, ) return build_hard_motion_context( conditioning=output, trim_frames=trim_frames, latent=latent, context_latent=context_latent, video_transition_steps=video_transition_steps, audio_transition_steps=audio_transition_steps, ) def build_hard_motion_context( conditioning: Any, trim_frames: int, latent: dict[str, Any], context_latent: dict[str, Any], video_transition_steps: int = 4, audio_transition_steps: int = 4, ) -> tuple[Any, int, dict[str, Any]]: """Validate native conditioning and add hard continuity to its target.""" video_keyframes, audio_keyframes = _native_keyframe_stats(conditioning) if video_keyframes < 1 or audio_keyframes < 1: raise RuntimeError( "easy h3 hard context requires Motion Context 0.4.0+ native " "video/audio keyframes and ComfyUI 0.34.0+; got " f"video_keyframes={video_keyframes}, audio_keyframes={audio_keyframes}" ) hard_latent = _hard_av_latent( latent, context_latent, int(trim_frames), video_transition_steps, audio_transition_steps, ) LOGGER.info( "Hard AV kept native conditioning (%d video keyframes, %d audio " "keyframes); minimax_refs untouched", video_keyframes, audio_keyframes, ) return conditioning, trim_frames, hard_latent def apply_hires_continuity( current_hires_latent: dict[str, Any], previous_hires_latent: dict[str, Any], context_length: int | str = "22", video_transition_steps: int = 4, ) -> tuple[dict[str, Any], int]: """Build the masked high-resolution seed for an H3 second pass.""" current_parts = _streams_from_latent(current_hires_latent) previous_parts = _streams_from_latent(previous_hires_latent) if len(current_parts) < 2 or not previous_parts: raise ValueError("easy h3 hires continuity: missing H3 AV latent streams") current_video, current_audio = current_parts[:2] previous_video = previous_parts[0] if current_video.ndim == 4: current_video = current_video.unsqueeze(0) if previous_video.ndim == 4: previous_video = previous_video.unsqueeze(0) if current_audio.ndim == 3: current_audio = current_audio.unsqueeze(0) if current_video.ndim != 5 or previous_video.ndim != 5 or current_audio.ndim != 4: raise ValueError("easy h3 hires continuity: invalid H3 latent dimensions") if ( previous_video.shape[0] != current_video.shape[0] or previous_video.shape[1] != current_video.shape[1] or previous_video.shape[3:] != current_video.shape[3:] ): raise ValueError( "easy h3 hires continuity: previous and current video resolutions differ" ) blocks, _, covered = _video_tail_from_latent( previous_hires_latent, int(context_length), ) copied_video = torch.cat(blocks, dim=2) copied_steps = int(copied_video.shape[2]) if copied_steps < 1 or copied_steps >= int(current_video.shape[2]): raise ValueError( "easy h3 hires continuity: copied prefix must be shorter than target" ) output_video = current_video.clone() output_video[:, :, :copied_steps] = copied_video.to( device=output_video.device, dtype=output_video.dtype, ) video_mask = torch.ones_like(output_video[:, :1], dtype=torch.float32) temporal_mask = video_mask.permute(0, 1, 3, 4, 2) locked_steps, ramp_values = _release_mask_inside_prefix( temporal_mask, copied_steps, video_transition_steps, ) video_mask = temporal_mask.permute(0, 1, 4, 2, 3).contiguous() audio_mask = torch.zeros_like(current_audio[:, :1], dtype=torch.float32) output = current_hires_latent.copy() if "noise_mask" in output: LOGGER.info("HiRes continuity is replacing the incoming noise_mask") output["samples"] = _official_nested_tensor((output_video, current_audio)) output["noise_mask"] = _official_nested_tensor((video_mask, audio_mask)) LOGGER.info( "HiRes continuity: video=%d steps/%d frames (lock=%d ramp=%s); " "second-pass audio is frozen", copied_steps, covered, locked_steps, ramp_values or "off", ) return output, int(covered) def apply_hires_anchor_continuity( current_hires_latent: dict[str, Any], previous_hires_latent: dict[str, Any], context_length: int | str = "22", anchor_length: int | str = str(REENCODED_ANCHOR_FRAMES), ) -> tuple[dict[str, Any], int]: """Hard-pin a re-encoded boundary anchor during ordinary hi-res refine.""" output = _reencoded_anchor_av_latent( current_hires_latent, _video_anchor_from_context_latent( previous_hires_latent, int(anchor_length), ), context_frames=int(context_length), anchor_frames=int(anchor_length), ) masks = _streams_from_latent({"samples": output["noise_mask"]}) samples = _streams_from_latent(output) if len(masks) < 2 or len(samples) < 2: raise ValueError("easy h3 hires anchor: missing H3 AV streams") masks[1] = torch.zeros_like(samples[1][:, :1], dtype=torch.float32) output["noise_mask"] = _official_nested_tensor(tuple(masks)) LOGGER.info("HiRes anchor continuity: second-pass audio is frozen") return output, int(context_length) def apply_motion_context( conditioning: Any, vae: Any, latent: dict[str, Any], context_length: int | str, audio_context_length: int = 24, context_frames: torch.Tensor | None = None, context_latent: dict[str, Any] | None = None, audio_vae: Any | None = None, context_audio: dict[str, Any] | None = None, ) -> tuple[Any, int]: """Attach previous-clip video/audio tails to MiniMax H3 conditioning.""" target_video = _video_from_latent(latent) latent_steps = int(target_video.shape[2]) width = int(target_video.shape[4]) * 16 height = int(target_video.shape[3]) * 16 target_frames = _pixel_frames(latent_steps) if context_latent is not None: source_video = _video_from_latent(context_latent) source_width = int(source_video.shape[4]) * 16 source_height = int(source_video.shape[3]) * 16 if (source_width, source_height) != (width, height): raise ValueError( "easy h3 motion context: context latent resolution " f"{source_width}x{source_height} does not match target " f"{width}x{height}" ) if int(source_video.shape[1]) != int(target_video.shape[1]): raise ValueError( "easy h3 motion context: context and target latent channels differ" ) available_frames = _pixel_frames(int(source_video.shape[2])) source_kind = "latent" else: if context_frames is None: raise ValueError( "easy h3 motion context: connect context_latent or context_frames" ) available_frames = int(context_frames.shape[0]) source_kind = "pixels" requested_frames = int(context_length) pinned_frames = min(requested_frames, available_frames) if pinned_frames < 1: raise ValueError("easy h3 motion context: no frames are available to pin") snapped_frames = next( grid for grid in VIDEO_RUN_GRID if grid <= pinned_frames ) if snapped_frames != pinned_frames: LOGGER.warning( "Context length %d is off the H3 VAE grid; using %d", pinned_frames, snapped_frames, ) pinned_frames = snapped_frames if pinned_frames >= target_frames: raise ValueError( "easy h3 motion context: the pinned context must be shorter than " "the target clip" ) if source_kind == "latent": blocks, offsets, span = _video_tail_from_latent( context_latent, pinned_frames, ) else: assert context_frames is not None tail = _resize_frames( context_frames[available_frames - pinned_frames :], width, height, ) encoded = vae.encode(tail) if getattr(encoded, "ndim", 0) != 5: raise ValueError( "easy h3 motion context: video VAE returned a non-H3 latent" ) steps = int(encoded.shape[2]) offsets = _step_offsets(steps) span = _pixel_frames(steps) if span != pinned_frames: raise RuntimeError( "easy h3 motion context: the video VAE temporal grid changed" ) blocks = [encoded[:, :, index : index + 1] for index in range(steps)] keyframes = [ { "resolved_frame_index": position, "latent": block, } for position, block in zip(offsets, blocks) ] audio_keyframe: dict[str, Any] | None = None audio_frames = 0 audio_steps = 0 if context_latent is not None or context_audio is not None: audio_frames = int(audio_context_length) or span if context_latent is not None: audio_latent, audio_steps, overhang = _audio_tail_from_latent( context_latent, audio_frames, ) else: if audio_vae is None or context_audio is None: raise ValueError( "easy h3 motion context: context_audio requires audio_vae" ) audio_latent, audio_steps = _encode_tail_audio( audio_vae, context_audio, audio_frames / FPS, ) overhang = 0.0 end_frame = float(span) + overhang / FRAME_RESCALE end_frame = round(FRAME_RESCALE * end_frame) / FRAME_RESCALE audio_keyframe = { "resolved_frame_index": end_frame - audio_steps / FRAME_RESCALE, "audio_latent": audio_latent, } output = [] dropped_positions: list[int | float] = [] for embedding, metadata in conditioning: values = metadata.copy() previous_keyframes = values.get("minimax_keyframes") or [] kept_keyframes = [] for keyframe in previous_keyframes: position = keyframe.get("resolved_frame_index", 0) if position >= target_frames: raise ValueError( "easy h3 motion context: conditioning keyframe at frame " f"{position} exceeds the {target_frames}-frame target" ) if position < span: dropped_positions.append(position) continue kept_keyframes.append(dict(keyframe)) native_keyframes = keyframes + ( [audio_keyframe] if audio_keyframe is not None else [] ) values["minimax_keyframes"] = kept_keyframes + native_keyframes output.append([embedding, values]) if dropped_positions: LOGGER.warning( "Dropped existing keyframe anchors inside the pinned head: %s", sorted(set(dropped_positions)), ) LOGGER.info( "Pinned %d H3 frames as %d blocks; trim=%d, audio=%d frames/%d steps", pinned_frames, len(blocks), span, audio_frames, audio_steps, ) return output, span