"""MiniMax-H3 (33B joint video+audio DiT) for ai-toolkit. Supports t2v (t2va) and first-frame i2v (fl2va) training and sampling, with joint audio when the dataset provides it. Image datasets train as single latent frames (keyframe-row geometry) and sampling with num_frames 1 renders a single image. The architecture lives in ./src/: - transformer.py: packed-sequence DiT, weight-compatible with the original ``MiniMaxAI/MiniMax-H3`` checkpoint keys - vae.py / audio_vae.py: the video VAE (causal CNN encoder + ViT decoder, 16x/17n+5->5n+2) and the waveform audio VAE (DAC/BigVGAN, 32 kHz, 40 Hz) - packing.py: packed-sequence geometry, rotary grids, sigma-shift math - text_encoder.py: Qwen3-VL-32B conditioning (unnormalized hidden_states[50], ": " + vision block presentation for keyframes) - pipeline.py: the released sampler (no CFG — the model is guidance-distilled) Weights load from the Comfy-Org repack (``Comfy-Org/MiniMax-H3``) by default: the pruned int8-ConvRot transformer, the nvfp4 AWQ Qwen3-VL text encoder (kept quantized through the toolkit's Ostris quantization backends — convrot8 and nvfp4 — with dequantized-matmul fallbacks for GPUs without the fast kernels), and the fp16/fp32 single-file VAEs. Files are resolved under ``MODELS_PATH`` (checked first, both at the repo-relative location and flat at the root) and downloaded from the hub into ``MODELS_PATH`` when missing. Individual files can be overridden via ``model_kwargs``: ``dit__path``, ``text_encoder_path``, ``video_vae_path``, ``audio_vae_path``; ``model_kwargs.partition`` picks ``fl2va``, ``fl2va_pruned`` (default), ``ref2va``, or ``ref2va_pruned``. Conventions bridged to ai-toolkit: - the model consumes t = 1 - sigma in [0, 1] (t=1 clean) and predicts the data-ward velocity ``clean - noise``; ai-toolkit targets ``noise - clean`` on a 0..1000 timestep scale, so timesteps are flipped and the prediction negated in get_noise_prediction - the audio stream runs on its own flow shift (3 vs video's 12): its sigma is derived per step from the video sigma via the closed-form remap, in training and sampling alike """ import os from functools import partial from typing import TYPE_CHECKING, List, Optional import torch import yaml from PIL import Image from safetensors.torch import load_file, save_file from toolkit.accelerator import unwrap_model from toolkit.advanced_prompt_embeds import AdvancedPromptEmbeds from toolkit.basic import flush from toolkit.config_modules import GenerateImageConfig, ModelConfig from toolkit.dto import DTO from toolkit.metadata import get_meta_for_safetensors from toolkit.models.base_model import BaseModel from toolkit.models.v2.text_encoders.qwen3_vl import Qwen3VLTextEncoder from toolkit.models.v2.resolver import ( find_file_recursive, repo_id_from_name_or_path, resolve_comfy_file, ) from toolkit.paths import MODELS_PATH from toolkit.util.comfy_quant_import import import_comfy_quantized_layers from toolkit.util.ostris_quant import OstrisLinear from toolkit.samplers.custom_flowmatch_sampler import ( CustomFlowMatchEulerDiscreteScheduler, ) from .src import packing packing_video_exts = [".mp4", ".avi", ".mov", ".webm", ".mkv", ".wmv", ".m4v", ".flv"] from .src.audio_vae import MiniMaxH3AudioVAE from .src.packing import ( KEYFRAME_ENCODE_SEED, KEYFRAME_NOISE_AUG_T, build_packed_sequence, pack_audio_latents, pad_layouts_to_batch, unpack_audio_tokens, patchify_video_latents, remap_sigma, unpatchify_video_tokens, ) from .src.pipeline import MiniMaxH3Pipeline from .src.ref_video_cache import ( load_ref_video_latent, load_video_ref_for_te, ref_frame_indices, static_image_video_ref, ) from .src.text_encoder import ( TEXT_ENCODER_LAYER, VideoRef, encode_minimax_h3_prompt, load_video_ref, trim_caption_tokens, ) from .src.transformer import MiniMaxH3Transformer, MiniMaxH3TransformerParams from .src.vae import MiniMaxH3VideoVAE if TYPE_CHECKING: from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO scheduler_config = { "num_train_timesteps": 1000, "shift": packing.VIDEO_SIGMA_SHIFT, "use_dynamic_shifting": False, } # Comfy-Org repack of the released weights, at ComfyUI's repo-relative paths # under MODELS_PATH (diffusion_models/, text_encoders/, vae/). Files are used # in place when present and downloaded to exactly these locations only when # missing. COMFY_REPO = "Comfy-Org/MiniMax-H3" COMFY_FILES = { "dit_fl2va": "diffusion_models/minimax_h3_fl2va_int8_convrot.safetensors", "dit_fl2va_pruned": "diffusion_models/minimax_h3_fl2va_pruned_int8_convrot.safetensors", "dit_ref2va": "diffusion_models/minimax_h3_ref2va_int8_convrot.safetensors", "dit_ref2va_pruned": "diffusion_models/minimax_h3_ref2va_pruned_int8_convrot.safetensors", "text_encoder": "text_encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors", "video_vae": "vae/minimax_h3_video_vae_fp16.safetensors", "audio_vae": "vae/minimax_h3_audio_vae_fp32.safetensors", "dit_fasth3": "diffusion_models/minimax_h3_fasth3_preview_v0.2_int8_convrot.safetensors", } # FastH3 (FastVideo 4-step VSA distill) int8-convrot repack, produced by # scripts/convert_minimax_h2_fastvideo.py from the FastVideo diffusers repo. # Hub fallback repo for the file (flat at the repo root, downloaded into # MODELS_PATH/diffusion_models/); not published there yet — until it is, the # file must exist locally or be built with the converter. FASTH3_REPO = "Kijai/MiniMax-H3-experimental" FASTH3_SOURCE_REPO = "FastVideo/FastVideo-Minimax-FastH3-Preview-v0.2" # tokenizer/processor/text-encoder config come from the original repo (tiny files) ORIGINAL_REPO = "MiniMaxAI/MiniMax-H3" def new_save_image_function( self: GenerateImageConfig, image, count=0, max_count=0, **kwargs ): # video (+ audio) previews save as mp4 try: from diffusers.utils import encode_video except ImportError: from diffusers.pipelines.ltx2.export_utils import encode_video image["output_path"] = self.get_image_path(count, max_count) os.makedirs(os.path.dirname(image["output_path"]), exist_ok=True) if image.get("audio", None) is None: image.pop("audio", None) image.pop("audio_sample_rate", None) encode_video(**image) flush() def blank_log_image_function(self, *args, **kwargs): # todo handle wandb logging of videos with audio return class MiniMaxH3VaeBundle(torch.nn.Module): """Holds both frozen autoencoders behind the single ``self.vae`` handle.""" def __init__(self, video_vae: MiniMaxH3VideoVAE, audio_vae: MiniMaxH3AudioVAE): super().__init__() self.video_vae = video_vae self.audio_vae = audio_vae @property def device(self): return self.video_vae.device @property def dtype(self): return self.video_vae.dtype def enable_gradient_checkpointing(self, enable: bool = True): self.video_vae.enable_gradient_checkpointing(enable) self.audio_vae.enable_gradient_checkpointing(enable) def disable_gradient_checkpointing(self): self.enable_gradient_checkpointing(False) class MinimaxH3Model(BaseModel): arch = "minimax_h3" use_old_lokr_format = False def __init__( self, device, model_config: ModelConfig, dtype="bf16", custom_pipeline=None, noise_scheduler=None, **kwargs, ): super().__init__( device, model_config, dtype, custom_pipeline, noise_scheduler, **kwargs ) self.is_flow_matching = True self.is_transformer = True self.target_lora_modules = ["MiniMaxH3Transformer"] self.supports_model_paths = True # keyframes ride into the Qwen3-VL conditioning as vision blocks, so # sampling (and control_path datasets) pass control images to # get_prompt_embeds self.encode_control_in_text_embeddings = True self.processor = None # Qwen3VLProcessor self._warned_frame_trim = False # video-ref presentation context: dataset config while caching training # embeds; the sample's frame cap while encoding sample prompts self._ref_video_dataset_config = None self._sample_ref_max_frames = None self.latent_space_version = "minimax_h3_v1" # caption token cap (vision blocks are never truncated); the released # stack has no limit — set 0 to disable self.max_text_length = int( self.model_config.model_kwargs.get("max_text_length", 512) ) @staticmethod def get_train_scheduler(): return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config) def get_bucket_divisibility(self): # 16x VAE spatial compression * 2x2 transformer patch return 32 def get_frame_count_snapper(self): # auto_frame_count: snap dataset clips down to the VAE's 17n+5 grid return packing.align_num_frames_down def prepare_sample_prompt_context(self, gen_config): # sample prompts: video refs are treated at the sample's length, not # the dataset's (which only applies while caching training embeds) self._ref_video_dataset_config = None self._sample_ref_max_frames = max(int(gen_config.num_frames), 5) @property def video_vae(self) -> MiniMaxH3VideoVAE: return self.vae.video_vae @property def audio_vae(self) -> MiniMaxH3AudioVAE: return self.vae.audio_vae # ------------------------------------------------------------------ # Loading # ------------------------------------------------------------------ def _resolve_comfy_file(self, component: str) -> str: """Find a weight file at its local location (model_kwargs override, comfy layout under MODELS_PATH or a local name_or_path dir), or download it there when (and only when) it is missing — see toolkit/models/v2/resolver.py for the search order.""" name_or_path = self.model_config.name_or_path extra_roots = ( [name_or_path] if name_or_path and os.path.isdir(name_or_path) else [] ) return resolve_comfy_file( COMFY_FILES[component], repo_id=repo_id_from_name_or_path(name_or_path, COMFY_REPO), override_path=self.model_config.model_kwargs.get(f"{component}_path", None), extra_roots=extra_roots, status_fn=self.print_and_status_update, ) def _dit_component(self) -> str: partition = str( self.model_config.model_kwargs.get("partition", "fl2va_pruned") ).lower() if partition not in ("fl2va", "fl2va_pruned", "ref2va", "ref2va_pruned"): raise ValueError( "model_kwargs.partition must be fl2va, fl2va_pruned, ref2va, " f"or ref2va_pruned, got {partition}" ) return f"dit_{partition}" def load_training_adapter(self, transformer: MiniMaxH3Transformer): """Load an assistant LoRA (e.g. a de-distillation adapter) as a LIVE module: active during training, deactivated by the sampler. It is deliberately NOT merged into the base weights — the transformer is pre-quantized, and a merge would resample every int8 scale. Path resolution: a local path is used as-is; otherwise the loras folder under MODELS_PATH is searched recursively for the filename; otherwise a ``user/repo/file.safetensors`` hub path downloads into MODELS_PATH/loras/training_adapters/. """ from toolkit.config_modules import NetworkConfig from toolkit.lora_special import LoRASpecialNetwork self.print_and_status_update("Loading assistant LoRA") lora_path = self.model_config.assistant_lora_path if not os.path.exists(lora_path): filename = os.path.basename(lora_path) found = find_file_recursive(os.path.join(MODELS_PATH, "loras"), filename) if found is not None: lora_path = found else: lora_splits = lora_path.split("/") if len(lora_splits) != 3: raise ValueError( f"Assistant LoRA path {lora_path} is not a local path, a " f"file under {os.path.join(MODELS_PATH, 'loras')}, or a " "'user/repo/file.safetensors' hub path." ) import huggingface_hub target_dir = os.path.join(MODELS_PATH, "loras", "training_adapters") os.makedirs(target_dir, exist_ok=True) try: lora_path = huggingface_hub.hf_hub_download( repo_id="/".join(lora_splits[:2]), filename=lora_splits[2], local_dir=target_dir, ) except Exception as e: raise ValueError( f"Failed to download assistant LoRA from {lora_path}: {e}" ) self.model_config.assistant_lora_path = lora_path # load the adapter; it stays a live module (never merged) and the # sampler toggles it off for previews lora_state_dict = load_file(lora_path) dim_key = next( k for k in lora_state_dict if k.endswith("lora_A.weight") or k.endswith("lora_down.weight") ) dim = int(lora_state_dict[dim_key].shape[0]) lora_state_dict = self.convert_lora_weights_before_load(lora_state_dict) network_config = NetworkConfig( **{ "type": "lora", "linear": dim, "linear_alpha": dim, "transformer_only": True, } ) LoRASpecialNetwork.LORA_PREFIX_UNET = "lora_transformer" network = LoRASpecialNetwork( text_encoder=None, unet=transformer, lora_dim=network_config.linear, multiplier=1.0, alpha=network_config.linear_alpha, train_unet=True, train_text_encoder=False, network_config=network_config, network_type=network_config.type, transformer_only=network_config.transformer_only, is_transformer=True, target_lin_modules=self.target_lora_modules, is_assistant_adapter=True, is_ara=True, ) network.apply_to(None, transformer, apply_text_encoder=False, apply_unet=True) network.force_to(self.device_torch, dtype=self.torch_dtype) network._update_torch_multiplier() network.load_weights(lora_state_dict) # frozen: the adapter shapes the training distribution but is never # itself trained, so its params must not collect gradients network.is_merged_in = False for param in network.parameters(): param.requires_grad_(False) network.eval() self.assistant_lora: LoRASpecialNetwork = network # live during training; the sampler's non-inverted assistant path # (BaseModel.generate_images) deactivates it for previews and turns # it back on afterwards self.assistant_lora.multiplier = 1.0 self.assistant_lora.is_active = True self.invert_assistant_lora = False def _load_transformer(self) -> MiniMaxH3Transformer: dit_path = self._resolve_comfy_file(self._dit_component()) self.print_and_status_update(f"Loading transformer from {dit_path}") # the mixin single-file path: config sniffed from the checkpoint # (adaln_t_table), pre-quantized ConvRot linears attached, everything # else at its stored precision (the bf16/fp16/fp32 mix is deliberate) return MiniMaxH3Transformer.load_model(dit_path, dtype=self.torch_dtype) def _load_text_encoder(self): from accelerate import init_empty_weights from transformers import ( AutoConfig, AutoProcessor, AutoTokenizer, Qwen3VLForConditionalGeneration, ) tokenizer = AutoTokenizer.from_pretrained( ORIGINAL_REPO, subfolder="FL2VA/tokenizer" ) processor = AutoProcessor.from_pretrained( ORIGINAL_REPO, subfolder="FL2VA/processor" ) te_path = self.model_config.te_name_or_path if te_path is not None and os.path.isdir(te_path): # transformers-format folder (e.g. the original repo's text_encoder) self.print_and_status_update( f"Loading Qwen3-VL text encoder from {te_path}" ) config = AutoConfig.from_pretrained(te_path) config.text_config.num_hidden_layers = TEXT_ENCODER_LAYER text_encoder = Qwen3VLTextEncoder.from_pretrained( te_path, config=config, torch_dtype=self.te_torch_dtype ) else: if te_path is not None: te_file = te_path else: te_file = self._resolve_comfy_file("text_encoder") self.print_and_status_update( f"Loading Qwen3-VL text encoder from {te_file}" ) # single-file ComfyUI checkpoint: 50 decoder layers, no final norm, # no lm_head; LM linears nvfp4 (AWQ), embeddings int8, vision bf16 config = AutoConfig.from_pretrained( ORIGINAL_REPO, subfolder="FL2VA/text_encoder" ) # only hidden_states[50] is consumed: truncate the decoder stack to # 50 layers; the final norm is neutralized below so # hidden_states[-1] stays the unnormalized layer-49 output config.text_config.num_hidden_layers = TEXT_ENCODER_LAYER config.tie_word_embeddings = False with init_empty_weights(): text_encoder = Qwen3VLTextEncoder(config) text_encoder.lm_head = None state_dict = load_file(te_file) def key_map(prefix: str) -> str: if prefix.startswith("model."): return "model.language_model." + prefix[len("model.") :] if prefix.startswith("visual."): return "model." + prefix return prefix state_dict, num_quantized = import_comfy_quantized_layers( text_encoder, state_dict, orig_dtype=self.te_torch_dtype, key_map=key_map, ) self.print_and_status_update( f" - attached {num_quantized} pre-quantized nvfp4/int8 layers" ) state_dict = { key_map(k[: k.rfind(".")]) + k[k.rfind(".") :]: v for k, v in state_dict.items() } result = text_encoder.load_state_dict(state_dict, assign=True, strict=False) quantized_keys = set() for name, m in text_encoder.named_modules(): if isinstance(m, OstrisLinear): quantized_keys.add(f"{name}.weight") allowed_missing_prefixes = ( "lm_head", "model.language_model.norm", "model.language_model.embed_tokens", ) bad_missing = [ k for k in result.missing_keys if k not in quantized_keys and not k.startswith(allowed_missing_prefixes) ] if bad_missing or result.unexpected_keys: raise ValueError( f"MiniMax-H3 text encoder load mismatch: missing {bad_missing[:8]}, " f"unexpected {result.unexpected_keys[:8]}" ) del state_dict text_encoder.model.language_model.norm = torch.nn.Identity() text_encoder.eval() text_encoder.requires_grad_(False) flush() return tokenizer, processor, text_encoder def _load_vaes(self) -> MiniMaxH3VaeBundle: self.print_and_status_update("Loading video VAE") video_vae = MiniMaxH3VideoVAE.load_model(self._resolve_comfy_file("video_vae")) self.print_and_status_update("Loading audio VAE") audio_vae = MiniMaxH3AudioVAE.load_model(self._resolve_comfy_file("audio_vae")) flush() return MiniMaxH3VaeBundle(video_vae, audio_vae) def load_model(self): dtype = self.torch_dtype self.print_and_status_update("Loading MiniMax-H3 model") transformer = self._load_transformer() # load assistant lora if specified (merged into the quantized weights) if self.model_config.assistant_lora_path is not None: self.load_training_adapter(transformer) # quantize + offload + placement, all driven by model_config transformer.aitk_post_load(**self.component_load_kwargs("transformer")) flush() tokenizer, processor, text_encoder = self._load_text_encoder() if any(isinstance(m, OstrisLinear) for m in text_encoder.modules()): # already nvfp4/int8 quantized; aitk_post_load skips quantize_te text_encoder.aitk_is_quantized = True # quantize + offload + placement, all driven by model_config text_encoder.aitk_post_load(**self.component_load_kwargs("te")) flush() vae_bundle = self._load_vaes() vae_bundle.to(self.vae_device_torch) self.noise_scheduler = MinimaxH3Model.get_train_scheduler() self.vae = vae_bundle self.text_encoder = text_encoder self.tokenizer = tokenizer self.processor = processor self.model = transformer self.pipeline = MiniMaxH3Pipeline(self) self.print_and_status_update("Model Loaded") # ------------------------------------------------------------------ # Text conditioning # ------------------------------------------------------------------ def _present_image_control(self, image: Image.Image): """Hook: how a control IMAGE enters the Qwen3-VL presentation. The default is a plain ```` image; ref2va can turn it into a static video reference (``image_refs_as_video``).""" return image def get_prompt_embeds(self, prompt, control_images=None) -> AdvancedPromptEmbeds: if isinstance(prompt, str): prompt = [prompt] if self.text_encoder.device == torch.device("cpu"): self.text_encoder.to(self.device_torch) # control tensors arrive in [0, 1]; the Qwen3-VL processor wants PIL keyframes_per_prompt = [None] * len(prompt) if control_images is not None: if isinstance(control_images, torch.Tensor): images = [control_images[i] for i in range(control_images.shape[0])] elif isinstance(control_images, list): images = [ c[0] if isinstance(c, torch.Tensor) and c.ndim == 4 else c for c in control_images ] else: images = [control_images] pil_images = [] for img in images: if isinstance(img, torch.Tensor): if img.ndim == 4: img = img[0] arr = (img.float().clamp(0, 1) * 255).round().to(torch.uint8) pil_images.append( self._present_image_control( Image.fromarray(arr.permute(1, 2, 0).cpu().numpy()) ) ) elif isinstance(img, str): # a control VIDEO path: 2 fps timestamped presentation over # the SAME frames the latent rows use (dataset treatment # when caching training embeds, sample-length at sampling) ds_cfg = getattr(self, "_ref_video_dataset_config", None) pil_images.append( load_video_ref_for_te( self, img, ds_cfg, max_frames=self._sample_ref_max_frames ) ) else: pil_images.append(img) if len(pil_images) == 1: keyframes_per_prompt = [pil_images] * len(prompt) elif len(pil_images) == len(prompt): keyframes_per_prompt = [[img] for img in pil_images] else: keyframes_per_prompt = [pil_images] * len(prompt) embeds_list, tags_list = [], [] for p, keyframes in zip(prompt, keyframes_per_prompt): embeds, tags = encode_minimax_h3_prompt( self.text_encoder, self.tokenizer, self.processor, p.strip(), keyframes=keyframes, device=self.device_torch, dtype=self.torch_dtype, max_length=self.max_text_length, ) embeds_list.append(embeds) tags_list.append(tags) pe = AdvancedPromptEmbeds(text_embeds=embeds_list, text_token_tags=tags_list) pe.frozen_dtype_keys = ["text_token_tags"] return pe # ------------------------------------------------------------------ # VAE encode / decode # ------------------------------------------------------------------ @torch.no_grad() def encode_images(self, image_list, device=None, dtype=None): """Images (C, H, W) or videos (T, C, H, W) in [-1, 1] -> normalized video latents (B, 24, t, h, w). Video frame counts are trimmed down to the VAE's 17n+5 grid when needed.""" if device is None: device = self.vae_device_torch if dtype is None: dtype = self.vae_torch_dtype if self.vae.device == torch.device("cpu"): self.vae.to(self.vae_device_torch) items = [] for image in image_list: if image.ndim == 3: items.append(image.unsqueeze(1)) # (C, 1, H, W) elif image.ndim == 4: items.append(image.permute(1, 0, 2, 3)) # (C, T, H, W) else: raise ValueError(f"Invalid image shape: {image.shape}") num_frames = items[0].shape[1] if num_frames > 1: aligned = packing.align_num_frames_down(num_frames) if aligned != num_frames and not self._warned_frame_trim: print( f"MiniMax-H3: trimming {num_frames}-frame clips to {aligned} " f"frames (the video VAE needs 17n+5: 5, 22, 39, 56, ...). Set " f"the dataset num_frames accordingly to avoid wasted decode." ) self._warned_frame_trim = True items = [it[:, :aligned] for it in items] batch = torch.stack(items).to(self.vae_device_torch, self.video_vae.dtype) latents = self.video_vae.encode(batch, sample=True) return latents.to(device, dtype=dtype) @torch.no_grad() def encode_keyframe_latents(self, frames: torch.Tensor) -> torch.Tensor: """(B, 3, 1, H, W) in [-1, 1] -> normalized latents (B, 24, 1, h, w), with the released conditioning recipe: seeded posterior sample (seed 42, independent of the request seed) rounded to fp16 before normalization.""" if self.vae.device == torch.device("cpu"): self.vae.to(self.vae_device_torch) generator = torch.Generator(device="cpu").manual_seed(KEYFRAME_ENCODE_SEED) latents = self.video_vae.encode( frames.to(self.vae_device_torch, self.video_vae.dtype), sample=True, generator=generator, fp16_round=True, ) return latents.float() def decode_latents(self, latents: torch.Tensor, device=None, dtype=None): # differentiable: pixel-space losses backprop through the video VAE if self.vae.device == torch.device("cpu"): self.vae.to(self.vae_device_torch) video = self.video_vae.decode(latents.to(self.vae.device, self.video_vae.dtype)) if device is not None: video = video.to(device, dtype=dtype) return video def decode_audio_latents(self, latents: torch.Tensor): # differentiable, like decode_latents """(B, 32, T) normalized -> waveform (B, 1, T*800) at 32 kHz.""" if self.vae.device == torch.device("cpu"): self.vae.to(self.vae_device_torch) return self.audio_vae.decode(latents.to(self.audio_vae.device, torch.float32)) @property def audio_sample_rate(self) -> int: return packing.AUDIO_SAMPLE_RATE def decode_packed_audio_rows(self, rows: torch.Tensor) -> torch.Tensor: # differentiable: audio perceptual losses backprop through the audio VAE """Packed channel-major audio rows (B, 2*T, 32) -> stereo waveform (B, 2, T*800) at 32 kHz. Each stereo channel decodes as its own batch item through the mono audio VAE.""" a_lat = rows.shape[1] // packing.AUDIO_CHANNELS latents = unpack_audio_tokens(rows, a_lat) # (B, 2, 32, T) b = latents.shape[0] waveform = self.decode_audio_latents( latents.reshape(b * packing.AUDIO_CHANNELS, latents.shape[2], a_lat) ) # (B*2, 1, samples) return waveform.reshape(b, packing.AUDIO_CHANNELS, -1) @torch.no_grad() def encode_audio(self, audio_data_list): """[{"waveform": (C, L), "sample_rate": int}, ...] -> packed audio rows (B, 2*T, 32), normalized, channel-major stereo.""" import torchaudio if self.vae.device == torch.device("cpu"): self.vae.to(self.device_torch) packed = [] for audio_data in audio_data_list: waveform = audio_data["waveform"].to(self.audio_vae.device, torch.float32) sample_rate = int(audio_data["sample_rate"]) if waveform.dim() == 1: waveform = waveform.unsqueeze(0) if waveform.shape[0] == 1: waveform = waveform.repeat(2, 1) # mono -> stereo elif waveform.shape[0] > 2: waveform = waveform[:2] if sample_rate != packing.AUDIO_SAMPLE_RATE: waveform = torchaudio.functional.resample( waveform, sample_rate, packing.AUDIO_SAMPLE_RATE ) # the mono VAE sees each stereo channel as its own batch item z = self.audio_vae.encode(waveform.unsqueeze(1)) # (2, 32, T) packed.append(pack_audio_latents(z.unsqueeze(0))) # (1, 2*T, 32) max_len = max(p.shape[1] for p in packed) packed = [ torch.nn.functional.pad(p, (0, 0, 0, max_len - p.shape[1])) for p in packed ] return torch.cat(packed, dim=0).to(self.device_torch, self.torch_dtype) # ------------------------------------------------------------------ # Training forward # ------------------------------------------------------------------ def _build_condition( self, batch: "DataLoaderBatchDTO", latent_shape, device, dtype ): """Build the packed sequence's condition segment for one train step. ``latent_shape`` is the target's ``(t_lat, h_lat, w_lat)``. Returns ``(cond_rows, cond_audio_rows, keyframe_anchors, ref_blocks)`` where ``cond_rows`` is ``(B, num_condition_rows, 96)`` or None. The base model implements fl2va: the clip's first frame as a keyframe when the dataset asks for i2v. MinimaxH3Ref2VAModel overrides this with image references from the control images.""" do_i2v = ( batch is not None and batch.dataset_config.do_i2v and getattr(batch, "num_frames", 1) > 1 ) if not do_i2v: return None, None, (), () if batch.first_frame_latents is not None: first_latents = batch.first_frame_latents.to(device, torch.float32) else: frames = batch.tensor if frames is None: raise ValueError( "do_i2v needs the first frame; no cached " "first_frame_latents or raw tensors in batch" ) first_frames = frames[:, 0] if frames.ndim == 5 else frames first_latents = self.encode_keyframe_latents( first_frames.unsqueeze(2).to(device) ) if first_latents.ndim == 4: first_latents = first_latents.unsqueeze(2) cond_noise = torch.randn_like(first_latents) first_latents = ( KEYFRAME_NOISE_AUG_T * first_latents + (1.0 - KEYFRAME_NOISE_AUG_T) * cond_noise ) return patchify_video_latents(first_latents).to(dtype), None, ("first",), () def get_noise_prediction( self, latent_model_input: torch.Tensor, # (B, 24, t, h, w) noisy latents timestep: torch.Tensor, # (B,) on the 0..1000 scale, 1000 = pure noise text_embeddings: AdvancedPromptEmbeds, batch: "DataLoaderBatchDTO" = None, **kwargs, ): device = self.device_torch dtype = self.torch_dtype if self.model.device == torch.device("cpu"): self.model.to(device) batch_size, _, t_lat, h_lat, w_lat = latent_model_input.shape with torch.no_grad(): sigma_v = (timestep.to(device, torch.float32) / 1000.0).clamp(1e-6, 1.0) if sigma_v.dim() == 0: sigma_v = sigma_v.unsqueeze(0) if sigma_v.shape[0] != batch_size: sigma_v = sigma_v.expand(batch_size) sigma_a = remap_sigma(sigma_v) t_v = 1.0 - sigma_v t_a = 1.0 - sigma_a # --- conditioning rows (fl2va keyframe / ref2va references) ---- ( cond_rows, cond_audio_rows, keyframe_anchors, ref_blocks, ) = self._build_condition(batch, (t_lat, h_lat, w_lat), device, dtype) # --- audio rows ------------------------------------------------- if batch is not None and getattr(batch, "num_frames", None): num_frames = batch.num_frames else: # invert 17n+5 -> 5n+2 from the latent frame count num_frames = (t_lat - 2) // 5 * 17 + 5 if t_lat > 1 else 1 a_lat = packing.audio_latent_num_frames(num_frames) # audio only trains for video batches from datasets that asked for # it. Cached latents can carry audio after do_audio was turned off, # and image (single frame) batches must never pick up a soundtrack # — either way it rides along as silence with no audio loss. do_audio = ( batch is not None and batch.dataset_config is not None and batch.dataset_config.do_audio and num_frames > 1 ) raw_audio = None if do_audio and batch.audio_latents is not None: raw_audio = batch.audio_latents.to(device, torch.float32) elif do_audio and getattr(batch, "audio_data", None) is not None: raw_audio = self.encode_audio(batch.audio_data).to( device, torch.float32 ) sa = sigma_a.view(-1, 1, 1) audio_target = None noisy_audio_rows = None if raw_audio is not None: expected_rows = a_lat * packing.AUDIO_CHANNELS if raw_audio.shape[1] > expected_rows: raw_audio = raw_audio[:, :expected_rows] elif raw_audio.shape[1] < expected_rows: raw_audio = torch.nn.functional.pad( raw_audio, (0, 0, 0, expected_rows - raw_audio.shape[1]) ) # the audio noise is drawn once per step and shared by every # pass (prior, primary, cfg/guidance, preservation) so they all # see the same soundtrack and every pass's target matches. It # rides on the latents DTO along with the trimmed audio so # on-the-fly encodes aren't repeated per pass. audio_noise = ( batch.latents.get("audio_noise") if isinstance(batch.latents, DTO) else None ) if audio_noise is not None and audio_noise.shape == raw_audio.shape: audio_noise = audio_noise.to(device, torch.float32) else: audio_noise = torch.randn_like(raw_audio) if batch.latents is not None: batch.latents = DTO( batch.latents, audio=raw_audio, audio_noise=audio_noise ) audio_rows = (1.0 - sa) * raw_audio + sa * audio_noise # model predicts clean - noise; audio_pred is negated below so # the target follows ai-toolkit's noise - clean convention audio_target = (audio_noise - raw_audio).detach() # what audio perceptual losses need to rebuild the clean # estimate (x0 = noisy - sigma_a * pred); rides the pred DTO noisy_audio_rows = audio_rows else: # no soundtrack: silence (zeros) noised at the audio sigma # rides along without contributing to the loss audio_rows = sa * torch.randn( batch_size, a_lat * packing.AUDIO_CHANNELS, 32, device=device, dtype=torch.float32, ) # embeds cached with a longer max_text_length: cap the caption # tail (vision blocks are never touched) trimmed = [ trim_caption_tokens(e, t, self.max_text_length) for e, t in zip( text_embeddings.text_embeds, text_embeddings.text_token_tags ) ] text_embed_list = [e for e, _ in trimmed] text_tag_list = [t for _, t in trimmed] # --- packed layout (per item: text lengths differ) -------------- layouts = [] for i in range(batch_size): layouts.append( build_packed_sequence( text_token_tags=text_tag_list[i].to("cpu"), num_latent_frames=t_lat, latent_height=h_lat, latent_width=w_lat, num_audio_latents=a_lat, keyframe_anchors=keyframe_anchors, ref_blocks=ref_blocks, ) ) ( position_ids, token_tags, video_indices, audio_indices, text_indices, _, ) = pad_layouts_to_batch(layouts) num_cond = layouts[0].num_condition_video_rows num_cond_audio = layouts[0].num_condition_audio_rows # per-row timesteps: text/video rows at t_v, audio rows at t_a, # condition rows pinned at max(t_v, 0.999); ref soundtracks clean row_t = t_v.view(-1, 1).expand(-1, token_tags.shape[1]).clone() row_t[:, audio_indices] = t_a.view(-1, 1) if num_cond > 0: cond_t = torch.maximum(t_v, torch.full_like(t_v, KEYFRAME_NOISE_AUG_T)) row_t[:, video_indices[:num_cond]] = cond_t.view(-1, 1) if num_cond_audio > 0: row_t[:, audio_indices[:num_cond_audio]] = 1.0 audio_rows = torch.cat( [cond_audio_rows.to(audio_rows.dtype), audio_rows], dim=1 ) # pad text embeds to the batch max length max_text = int(text_indices.shape[0]) text_batch = torch.zeros( batch_size, max_text, text_embed_list[0].shape[-1], device=device, dtype=dtype, ) for i, emb in enumerate(text_embed_list): text_batch[i, : emb.shape[0]] = emb.to(device, dtype) video_rows = patchify_video_latents( latent_model_input.to(device, torch.float32) ).to(dtype) if cond_rows is not None: video_rows = torch.cat([cond_rows, video_rows], dim=1) video_pred, audio_pred = self.model( hidden_states=video_rows, audio_hidden_states=audio_rows.to(dtype), encoder_hidden_states=text_batch, row_timesteps=row_t.to(device), token_tags=token_tags.to(device), position_ids=position_ids.to(device), video_indices=video_indices.to(device), audio_indices=audio_indices.to(device), text_indices=text_indices.to(device), # target-video token grid (patch 1x2x2); consumed only by VSA models vsa_video_grid=(t_lat, h_lat // 2, w_lat // 2), ) if num_cond_audio > 0: # reference soundtrack rows are conditioning, not targets audio_pred = audio_pred[:, num_cond_audio:] video_pred = video_pred[:, num_cond:] noise_pred = unpatchify_video_tokens(video_pred, t_lat, h_lat, w_lat) if audio_target is not None: # every pass's DTO carries its own audio stream; preds flipped to # ai-toolkit's noise - clean convention return DTO( -noise_pred, audio=-audio_pred, audio_target=audio_target, audio_noisy=noisy_audio_rows, audio_sigma=sigma_a, ) return -noise_pred def get_loss_target(self, *args, **kwargs): noise = kwargs.get("noise") batch = kwargs.get("batch") return (noise - batch.latents).detach() # ------------------------------------------------------------------ # Sampling (training previews) # ------------------------------------------------------------------ def get_generation_pipeline(self): return MiniMaxH3Pipeline(self) def generate_single_image( self, pipeline: MiniMaxH3Pipeline, gen_config: GenerateImageConfig, conditional_embeds: AdvancedPromptEmbeds, unconditional_embeds: AdvancedPromptEmbeds, generator: torch.Generator, extra: dict, ): if self.model.device == torch.device("cpu"): self.model.to(self.device_torch) sc = self.get_bucket_divisibility() gen_config.width = max(sc, int(gen_config.width // sc * sc)) gen_config.height = max(sc, int(gen_config.height // sc * sc)) is_video = gen_config.num_frames > 1 if is_video: gen_config.num_frames = packing.align_num_frames_down(gen_config.num_frames) gen_config.fps = packing.FPS gen_config.save_image = partial(new_save_image_function, gen_config) gen_config.log_image = partial(blank_log_image_function, gen_config) gen_config.output_ext = "mp4" ctrl_img = None if gen_config.ctrl_img is not None: ctrl_img = Image.open(gen_config.ctrl_img).convert("RGB") ctrl_img = packing.prepare_keyframe_image( ctrl_img, gen_config.height, gen_config.width, stretch=True ) with_audio = bool(self.model_config.model_kwargs.get("sample_audio", True)) result = pipeline( conditional_embeds=conditional_embeds, unconditional_embeds=unconditional_embeds, height=gen_config.height, width=gen_config.width, num_frames=gen_config.num_frames, num_inference_steps=gen_config.num_inference_steps, guidance_scale=gen_config.guidance_scale, latents=gen_config.latents, generator=generator, ctrl_img=ctrl_img, with_audio=with_audio and is_video, ) if is_video: return result # dict consumed by new_save_image_function return result[0] # ------------------------------------------------------------------ # Saving / bookkeeping # ------------------------------------------------------------------ def get_model_has_grad(self): return False def get_te_has_grad(self): return False def save_model(self, output_path, meta, save_dtype): from toolkit.util.quantize import dequantize_if_quantized transformer: MiniMaxH3Transformer = unwrap_model(self.model) os.makedirs(os.path.join(output_path, "transformer"), exist_ok=True) state_dict = transformer.state_dict() save_dict = {} for k, v in state_dict.items(): v = dequantize_if_quantized(v) if v.is_floating_point() and not k.startswith( MiniMaxH3Transformer.FP32_KEY_PREFIXES ): v = v.to(save_dtype) save_dict[k] = v.clone().to("cpu") meta_st = get_meta_for_safetensors(meta, name="minimax_h3") save_file( save_dict, os.path.join(output_path, "transformer", "model.safetensors"), metadata=meta_st, ) with open(os.path.join(output_path, "aitk_meta.yaml"), "w") as f: yaml.dump(meta, f) def get_base_model_version(self): return "minimax_h3" def get_transformer_block_names(self) -> Optional[List[str]]: return ["blocks"] def get_quantization_exclude_modules(self) -> Optional[List[str]]: # float32 islands, the conditioning projection, the token refiner and # the AdaLN projections — all shipped unquantized in the pre-quantized # checkpoints (pruned files carry tiny fp16 adaln linears fed by the # 8-dim time table), so excluding them makes quantize with the # checkpoint's own qtype an exact no-op and keeps the sensitive # modulation path at full precision under any other qtype. return [ "video_patch_proj*", "audio_patch_proj*", "time_embedder*", "final_layer*", "condition_proj*", "token_refiner*", "*adaln_proj*", ] # ComfyUI's MiniMax-H3 keys are the original checkpoint keys, so the # standard diffusion_model prefix maps directly lora_keys_use_comfy_prefix = True class MinimaxH3Ref2VAModel(MinimaxH3Model): """Reference-to-video (ref2va): the control images ride along as reference blocks in the packed sequence (plus ``: `` vision blocks in the Qwen3-VL conditioning) instead of anchoring the first frame. Image references only for now. Every reference keeps its OWN aspect ratio and is resized to the TARGET's pixel area (axes snapped to /32), then sits on its own aspect-normalized rotary grid — like the released ref2va's per-reference grids, but area-matched to the target instead of the 2048px reference short edge. Each reference block advances the rotary media clock by 1.0. At sampling, ctrl images are ALWAYS references, never first frames. ``model_kwargs.image_refs_as_video`` (default off) routes still-image references through the VIDEO reference path instead: the image is held for ``image_ref_video_frames`` frames (17n+5, default 5) as a silent static clip — video sizing (true area match), multi-frame latent block, temporal-span rotary advance, and a ``