diff --git a/README.md b/README.md index 5ce450cf..a98a60ff 100644 --- a/README.md +++ b/README.md @@ -57,6 +57,7 @@ AI Toolkit is an easy to use all in one training suite for diffusion models. I t - [Wan-AI/Wan2.2-TI2V-5B-Diffusers](https://huggingface.co/Wan-AI/Wan2.2-TI2V-5B-Diffusers) (Wan 2.2 TI2V 5B) - [Lightricks/LTX-2](https://huggingface.co/Lightricks/LTX-2) (LTX-2) - [Lightricks/LTX-2.3](https://huggingface.co/Lightricks/LTX-2.3) (LTX-2.3) +- [MiniMaxAI/MiniMax-H3](https://huggingface.co/MiniMaxAI/MiniMax-H3) (MiniMaxAI/MiniMax-H3) ### Audio - [ACE-Step/Ace-Step1.5](https://huggingface.co/ACE-Step/Ace-Step1.5) (Ace Step 1.5) diff --git a/extensions_built_in/diffusion_models/__init__.py b/extensions_built_in/diffusion_models/__init__.py index c7e8a147..ea325b8c 100644 --- a/extensions_built_in/diffusion_models/__init__.py +++ b/extensions_built_in/diffusion_models/__init__.py @@ -19,6 +19,7 @@ from .prx_pixel_t2i import PRXPixelT2IModel from .krea2 import Krea2Model from .boogu_image import BooguImageModel, BooguImageEditModel from .mageflow import MageFlowModel, MageFlowEditModel +from .minimax_h3 import MinimaxH3Model AI_TOOLKIT_MODELS = [ # put a list of models here @@ -54,4 +55,5 @@ AI_TOOLKIT_MODELS = [ BooguImageEditModel, MageFlowModel, MageFlowEditModel, + MinimaxH3Model, ] diff --git a/extensions_built_in/diffusion_models/minimax_h3/__init__.py b/extensions_built_in/diffusion_models/minimax_h3/__init__.py new file mode 100644 index 00000000..80eb9f47 --- /dev/null +++ b/extensions_built_in/diffusion_models/minimax_h3/__init__.py @@ -0,0 +1 @@ +from .minimax_h3 import MinimaxH3Model diff --git a/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py b/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py new file mode 100644 index 00000000..74aa388c --- /dev/null +++ b/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py @@ -0,0 +1,947 @@ +"""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. 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`` (default) or ``ref2va``. + +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.memory_management import MemoryManager +from toolkit.metadata import get_meta_for_safetensors +from toolkit.models.base_model import BaseModel +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 toolkit.util.quantize import get_qtype, quantize, quantize_model +from optimum.quanto import freeze + +from .src import packing +from .src.audio_vae import MiniMaxH3AudioVAE, fold_audio_vae_weight_norm +from .src.packing import ( + KEYFRAME_ENCODE_SEED, + KEYFRAME_NOISE_AUG_T, + build_packed_sequence, + pack_audio_latents, + pad_layouts_to_batch, + patchify_video_latents, + remap_sigma, + unpatchify_video_tokens, +) +from .src.pipeline import MiniMaxH3Pipeline +from .src.text_encoder import TEXT_ENCODER_LAYER, encode_minimax_h3_prompt +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_pruned_int8_convrot.safetensors", + "dit_ref2va": "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", +} +# 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 + + +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 + self.latent_space_version = "minimax_h3_v1" + + @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 + + @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, or download it there + when (and only when) it is missing. + + Search order: model_kwargs override, the repo-relative path under + MODELS_PATH (diffusion_models/, text_encoders/, vae/), the bare + filename at the root as a fallback, the same spots under name_or_path + when it is a local folder, then the hub — downloaded to the + repo-relative path under MODELS_PATH. + """ + override = self.model_config.model_kwargs.get(f"{component}_path", None) + if override is not None: + if not os.path.exists(override): + raise FileNotFoundError( + f"model_kwargs.{component}_path does not exist: {override}" + ) + return override + + rel_path = COMFY_FILES[component] + roots = [MODELS_PATH] + name_or_path = self.model_config.name_or_path + if name_or_path and os.path.isdir(name_or_path): + roots.append(name_or_path) + for root in roots: + for rel in (rel_path, os.path.basename(rel_path)): + candidate = os.path.join(root, rel) + if os.path.exists(candidate): + return candidate + + import huggingface_hub + + repo_id = COMFY_REPO + if name_or_path and not os.path.exists(name_or_path) and "/" in name_or_path: + repo_id = name_or_path + self.print_and_status_update( + f"Downloading {rel_path} from {repo_id} into {MODELS_PATH}" + ) + return huggingface_hub.hf_hub_download( + repo_id=repo_id, filename=rel_path, local_dir=MODELS_PATH + ) + + def _dit_component(self) -> str: + partition = str( + self.model_config.model_kwargs.get("partition", "fl2va") + ).lower() + if partition not in ("fl2va", "ref2va"): + raise ValueError( + f"model_kwargs.partition must be fl2va or ref2va, got {partition}" + ) + return f"dit_{partition}" + + def _load_transformer(self) -> MiniMaxH3Transformer: + dtype = self.torch_dtype + dit_path = self._resolve_comfy_file(self._dit_component()) + self.print_and_status_update(f"Loading transformer from {dit_path}") + state_dict = load_file(dit_path) + + params = MiniMaxH3TransformerParams() + table = state_dict.get("adaln_t_table", None) + if table is not None: + # pruned checkpoint: factored timestep table instead of the MLP + params.adaln_t_table_size = table.shape[0] + params.time_embed_dim = table.shape[1] + + with torch.device("meta"): + transformer = MiniMaxH3Transformer(params) + + # attach the pre-quantized (int8 ConvRot) linears onto the toolkit's + # quantization backends; the rest loads at its stored precision (the + # checkpoint's bf16/fp16/fp32 mix is deliberate) + state_dict, num_quantized = import_comfy_quantized_layers( + transformer, state_dict, orig_dtype=dtype + ) + if num_quantized: + self.print_and_status_update( + f" - attached {num_quantized} pre-quantized ConvRot layers" + ) + result = transformer.load_state_dict(state_dict, assign=True, strict=False) + quantized_weight_keys = { + f"{name}.weight" + for name, m in transformer.named_modules() + if isinstance(m, OstrisLinear) + } + bad_missing = [k for k in result.missing_keys if k not in quantized_weight_keys] + if bad_missing or result.unexpected_keys: + raise ValueError( + f"MiniMax-H3 transformer load mismatch: missing {bad_missing[:8]}, " + f"unexpected {result.unexpected_keys[:8]}" + ) + del state_dict + flush() + return transformer + + 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 = Qwen3VLForConditionalGeneration.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 = Qwen3VLForConditionalGeneration(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_sd = load_file(self._resolve_comfy_file("video_vae")) + # normalization stats ride along in the comfy file; the module holds + # them as non-persistent buffers, keep them float32 + video_stats = { + k: video_sd.pop(k).float() + for k in ("latents_mean", "latents_std") + if k in video_sd + } + video_vae = MiniMaxH3VideoVAE() + video_vae.load_state_dict(video_sd, strict=True, assign=True) + for k, v in video_stats.items(): + getattr(video_vae, k).copy_(v) + video_vae.eval().requires_grad_(False) + del video_sd + + self.print_and_status_update("Loading audio VAE") + audio_sd = load_file(self._resolve_comfy_file("audio_vae")) + audio_stats = { + k: audio_sd.pop(k).float() + for k in ("latents_mean", "latents_std") + if k in audio_sd + } + # comfy repack ships the weight norm already folded; fold only if the + # raw parametrization is present (original-repo file) + if any(k.endswith("weight_g") for k in audio_sd.keys()): + audio_sd = fold_audio_vae_weight_norm(audio_sd) + audio_vae = MiniMaxH3AudioVAE() + audio_vae.load_state_dict(audio_sd, strict=True, assign=True) + for k, v in audio_stats.items(): + getattr(audio_vae, k).copy_(v) + audio_vae.to(torch.float32).eval().requires_grad_(False) + del audio_sd + 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() + + if self.model_config.quantize: + self.print_and_status_update("Quantizing transformer") + quantize_model(self, transformer) + flush() + + if ( + self.model_config.layer_offloading + and self.model_config.layer_offloading_transformer_percent > 0 + ): + MemoryManager.attach( + transformer, + self.device_torch, + offload_percent=self.model_config.layer_offloading_transformer_percent, + ) + + if self.model_config.low_vram: + self.print_and_status_update("Keeping transformer on CPU") + transformer.to("cpu") + else: + transformer.to(self.device_torch) + flush() + + tokenizer, processor, text_encoder = self._load_text_encoder() + te_prequantized = any( + isinstance(m, OstrisLinear) for m in text_encoder.modules() + ) + if self.model_config.quantize_te and not te_prequantized: + self.print_and_status_update("Quantizing text encoder") + quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te)) + freeze(text_encoder) + flush() + elif self.model_config.quantize_te: + self.print_and_status_update( + "Text encoder is already nvfp4/int8 quantized; skipping quantize_te" + ) + if ( + self.model_config.layer_offloading + and self.model_config.layer_offloading_text_encoder_percent > 0 + ): + MemoryManager.attach( + text_encoder, + self.device_torch, + offload_percent=self.model_config.layer_offloading_text_encoder_percent, + ) + if self.model_config.low_vram: + text_encoder.to("cpu") + else: + text_encoder.to(self.device_torch) + 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 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( + Image.fromarray(arr.permute(1, 2, 0).cpu().numpy()) + ) + 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, + ) + 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)) + + @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 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 + + # --- i2v first-frame conditioning rows ------------------------- + do_i2v = ( + batch is not None + and batch.dataset_config.do_i2v + and getattr(batch, "num_frames", 1) > 1 + ) + cond_rows = None + if do_i2v: + 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 + ) + cond_rows = patchify_video_latents(first_latents).to(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) + raw_audio = None + if batch is not None and batch.audio_latents is not None: + raw_audio = batch.audio_latents.to(device, torch.float32) + elif batch is not None 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) + 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]) + ) + audio_noise = torch.randn_like(raw_audio) + # model predicts clean - noise; audio_pred is negated below so + # the stored target follows ai-toolkit's noise - clean + batch.audio_target = (audio_noise - raw_audio).detach() + audio_rows = (1.0 - sa) * raw_audio + sa * audio_noise + 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, + ) + + # --- packed layout (per item: text lengths differ) -------------- + layouts = [] + anchors = ("first",) if cond_rows is not None else () + for i in range(batch_size): + layouts.append( + build_packed_sequence( + text_token_tags=text_embeddings.text_token_tags[i].to("cpu"), + num_latent_frames=t_lat, + latent_height=h_lat, + latent_width=w_lat, + num_audio_latents=a_lat, + keyframe_anchors=anchors, + ) + ) + ( + position_ids, + token_tags, + video_indices, + audio_indices, + text_indices, + _, + ) = pad_layouts_to_batch(layouts) + num_cond = layouts[0].num_condition_video_rows + + # per-row timesteps: text/video rows at t_v, audio rows at t_a, + # condition rows pinned at max(t_v, 0.999) + 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) + + # pad text embeds to the batch max length + max_text = int(text_indices.shape[0]) + text_batch = torch.zeros( + batch_size, + max_text, + text_embeddings.text_embeds[0].shape[-1], + device=device, + dtype=dtype, + ) + for i, emb in enumerate(text_embeddings.text_embeds): + 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), + ) + + if batch is not None and batch.audio_target is not None: + # flip to ai-toolkit's noise - clean convention + batch.audio_pred = -audio_pred + + video_pred = video_pred[:, num_cond:] + noise_pred = unpatchify_video_tokens(video_pred, t_lat, h_lat, w_lat) + 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, and the token refiner + # (shipped bf16 in the pre-quantized checkpoints — excluding it makes + # quantize with the checkpoint's own qtype an exact no-op). The + # per-block adaln_proj stays quantizable (half the parameter count). + return [ + "video_patch_proj*", + "audio_patch_proj*", + "time_embedder*", + "final_layer*", + "condition_proj*", + "token_refiner*", + ] + + def convert_lora_weights_before_save(self, state_dict): + # ComfyUI's MiniMax-H3 keys are the original checkpoint keys, so the + # standard diffusion_model prefix maps directly + return { + k.replace("transformer.", "diffusion_model."): v + for k, v in state_dict.items() + } + + def convert_lora_weights_before_load(self, state_dict): + return { + k.replace("diffusion_model.", "transformer."): v + for k, v in state_dict.items() + } diff --git a/extensions_built_in/diffusion_models/minimax_h3/src/audio_vae.py b/extensions_built_in/diffusion_models/minimax_h3/src/audio_vae.py new file mode 100644 index 00000000..504d0b3f --- /dev/null +++ b/extensions_built_in/diffusion_models/minimax_h3/src/audio_vae.py @@ -0,0 +1,504 @@ +"""MiniMax-H3 audio VAE -- frozen DAC/BigVGAN waveform autoencoder. + +32 kHz mono waveform <-> 32-channel latents at 40 latents/second (hop 800, +no mel front-end and no separate vocoder). Stereo is carried by the CALLER +as two batch items. Runs in fp32 (the BigVGAN decoder degrades audibly +under bf16). + +Weight-compatible with the ``MiniMaxAI/MiniMax-H3`` checkpoint +``FL2VA/audio_vae/model.safetensors`` once ``fold_audio_vae_weight_norm`` +collapses its ``weight_g``/``weight_v`` weight-norm pairs into plain conv +weights:: + + vae = MiniMaxH3AudioVAE() + vae.load_state_dict(fold_audio_vae_weight_norm(load_file(path)), strict=True) +""" + +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch import Tensor + +# fmt: off +# Per-channel latent statistics from the released FL2VA/audio_vae/config.json. +_LATENTS_MEAN = [ + -0.020211687488382354, 0.3876466479950502, -0.04398279799186767, + -0.28591514936373, 0.08179686214561671, -0.35782641352446604, + 0.040623809960919084, -0.01552534501956604, -0.223362481667332, + 0.1821006842509091, 0.2941778783780663, -0.07901167601970885, + -0.056815072777201, -0.3699028221860095, -0.31616315591624855, + 0.5905951377425391, -0.052139568068853864, 0.013673160263486295, + -0.03691647864630577, 0.09732660653298163, -0.3394662328788498, + -0.30685677538541667, -0.24504598907458763, -0.034698524462007344, + 0.02868032184767538, -0.21217779266454084, -0.1678263169941987, + 0.3221287889040614, -0.1223055851554907, 0.4356604928128464, + -0.0502599202236253, 0.3979258376211797, +] +_LATENTS_STD = [ + 1.6895524230479284, 2.76263727217653, 1.7945344281264435, + 1.6801681847309828, 1.6390226546605453, 2.7788298348882177, + 1.7659090095747236, 1.6199757612137327, 2.6336525640336896, + 1.8539356672817833, 2.5056497896915633, 1.811019237886178, + 1.9579657790720237, 1.6685498243529284, 1.4922469314453364, + 3.298670198067373, 1.9491804496832168, 1.8720003270431442, + 1.8334080103291832, 1.6488070416529093, 1.6176957696319716, + 1.9131449234774398, 1.5695245398428617, 1.6943659940415912, + 1.8318420762504692, 1.5540637421583379, 1.9344930328968526, + 1.599198216109855, 1.718045989838149, 1.6307219190837705, + 1.8661226051202384, 1.5613768203168363, +] +# fmt: on + + +def fold_audio_vae_weight_norm(state_dict: dict) -> dict: + """Fold weight-norm ``weight_g``/``weight_v`` pairs into plain ``.weight`` + tensors, passing every other key through untouched. + + The checkpoint uses ``nn.utils.weight_norm`` with the default ``dim=0``, + so ``w = g * v / ||v||`` with the norm taken over every dim except 0. + That also holds for the decoder's ConvTranspose1d layers: their + ``weight_v`` is laid out ``[in, out, k]`` and ``weight_g`` is + ``[in, 1, 1]``, i.e. the reduction still runs over dims (1, 2). + """ + folded = {} + for name, tensor in state_dict.items(): + if name.endswith(".weight_g"): + continue + if name.endswith(".weight_v"): + g = state_dict[name[: -len("weight_v")] + "weight_g"] + norm = tensor.norm(p=2, dim=tuple(range(1, tensor.ndim)), keepdim=True) + folded[name[: -len("weight_v")] + "weight"] = g * tensor / norm + else: + folded[name] = tensor + return folded + + +def _kaiser_sinc_filter1d(cutoff: float, half_width: float, kernel_size: int) -> Tensor: + # Kaiser-windowed sinc low-pass, arithmetically identical to + # alias-free-torch. The result is also stored in the checkpoint as a + # buffer, so this init only matters for randomly initialized models. + half_size = kernel_size // 2 + attenuation = 2.285 * (half_size - 1) * math.pi * (4 * half_width) + 7.95 + if attenuation > 50.0: + beta = 0.1102 * (attenuation - 8.7) + elif attenuation >= 21.0: + beta = 0.5842 * (attenuation - 21.0) ** 0.4 + 0.07886 * (attenuation - 21.0) + else: + beta = 0.0 + window = torch.kaiser_window(kernel_size, beta=beta, periodic=False) + if kernel_size % 2 == 0: + time = torch.arange(-half_size, half_size) + 0.5 + else: + time = torch.arange(kernel_size) - half_size + kernel = 2 * cutoff * window * torch.sinc(2 * cutoff * time) + kernel = kernel / kernel.sum() + return kernel.view(1, 1, kernel_size) + + +class Snake1d(nn.Module): + """``x + (alpha + 1e-9)^-1 * sin(alpha * x)^2`` with per-channel alpha + stored as ``[1, C, 1]``. The DAC encoder's activation.""" + + def __init__(self, channels: int): + super().__init__() + self.alpha = nn.Parameter(torch.ones(1, channels, 1)) + + def forward(self, x: Tensor) -> Tensor: + return x + (self.alpha + 1e-9).reciprocal() * torch.sin(self.alpha * x).pow(2) + + +class SnakeBeta(nn.Module): + """BigVGAN activation with separate frequency/magnitude parameters, both + stored log-scale as ``[C]`` vectors: + ``x + (exp(beta) + 1e-9)^-1 * sin(exp(alpha) * x)^2``.""" + + def __init__(self, channels: int): + super().__init__() + self.alpha = nn.Parameter(torch.zeros(channels)) + self.beta = nn.Parameter(torch.zeros(channels)) + + def forward(self, x: Tensor) -> Tensor: + alpha = torch.exp(self.alpha)[None, :, None] + beta = torch.exp(self.beta)[None, :, None] + return x + (beta + 1e-9).reciprocal() * torch.sin(alpha * x).pow(2) + + +class UpSample1d(nn.Module): + """Anti-aliased ratio-x upsampler (depthwise transposed Kaiser-sinc conv).""" + + def __init__(self, ratio: int = 2, kernel_size: int = 12): + super().__init__() + self.ratio = ratio + self.pad = kernel_size // ratio - 1 + self.pad_left = self.pad * ratio + (kernel_size - ratio) // 2 + self.pad_right = self.pad * ratio + (kernel_size - ratio + 1) // 2 + self.register_buffer( + "filter", + _kaiser_sinc_filter1d(0.5 / ratio, 0.6 / ratio, kernel_size), + ) + + def forward(self, x: Tensor) -> Tensor: + channels = x.shape[1] + x = F.pad(x, (self.pad, self.pad), mode="replicate") + x = self.ratio * F.conv_transpose1d( + x, self.filter.expand(channels, -1, -1), stride=self.ratio, groups=channels + ) + return x[..., self.pad_left : -self.pad_right] + + +class LowPassFilter1d(nn.Module): + def __init__(self, cutoff: float, half_width: float, stride: int, kernel_size: int): + super().__init__() + even = kernel_size % 2 == 0 + self.pad_left = kernel_size // 2 - int(even) + self.pad_right = kernel_size // 2 + self.stride = stride + self.register_buffer( + "filter", _kaiser_sinc_filter1d(cutoff, half_width, kernel_size) + ) + + def forward(self, x: Tensor) -> Tensor: + channels = x.shape[1] + x = F.pad(x, (self.pad_left, self.pad_right), mode="replicate") + return F.conv1d( + x, self.filter.expand(channels, -1, -1), stride=self.stride, groups=channels + ) + + +class DownSample1d(nn.Module): + """Anti-aliased ratio-x downsampler.""" + + def __init__(self, ratio: int = 2, kernel_size: int = 12): + super().__init__() + self.lowpass = LowPassFilter1d(0.5 / ratio, 0.6 / ratio, ratio, kernel_size) + + def forward(self, x: Tensor) -> Tensor: + return self.lowpass(x) + + +class AliasFreeActivation1d(nn.Module): + """2x upsample -> activation -> 2x downsample (BigVGAN's alias-free wrapper).""" + + def __init__(self, activation: nn.Module, ratio: int = 2, kernel_size: int = 12): + super().__init__() + self.act = activation + self.upsample = UpSample1d(ratio, kernel_size) + self.downsample = DownSample1d(ratio, kernel_size) + + def forward(self, x: Tensor) -> Tensor: + return self.downsample(self.act(self.upsample(x))) + + +class ResidualUnit(nn.Module): + """DAC residual unit: Snake -> dilated conv k7 -> Snake -> conv k1.""" + + def __init__(self, dim: int, dilation: int): + super().__init__() + self.block = nn.Sequential( + Snake1d(dim), + nn.Conv1d( + dim, dim, 7, dilation=dilation, padding=((7 - 1) * dilation) // 2 + ), + Snake1d(dim), + nn.Conv1d(dim, dim, 1), + ) + + def forward(self, x: Tensor) -> Tensor: + y = self.block(x) + crop = (x.shape[-1] - y.shape[-1]) // 2 + if crop > 0: + x = x[..., crop:-crop] + return x + y + + +class EncoderBlock(nn.Module): + """Three residual units (dilations 1/3/9), Snake, then a strided + channel-doubling conv (kernel 2*stride, padding ceil(stride/2)).""" + + def __init__(self, out_dim: int, stride: int): + super().__init__() + in_dim = out_dim // 2 + self.block = nn.Sequential( + ResidualUnit(in_dim, dilation=1), + ResidualUnit(in_dim, dilation=3), + ResidualUnit(in_dim, dilation=9), + Snake1d(in_dim), + nn.Conv1d( + in_dim, + out_dim, + kernel_size=2 * stride, + stride=stride, + padding=math.ceil(stride / 2), + ), + ) + + def forward(self, x: Tensor) -> Tensor: + return self.block(x) + + +class DacEncoder(nn.Module): + """``[B, 1, samples] -> [B, d_latent, samples / prod(strides)]``.""" + + def __init__(self, d_model: int, strides: tuple, d_latent: int): + super().__init__() + layers = [nn.Conv1d(1, d_model, 7, padding=3)] + for stride in strides: + d_model *= 2 + layers.append(EncoderBlock(d_model, stride)) + layers += [Snake1d(d_model), nn.Conv1d(d_model, d_latent, 3, padding=1)] + self.block = nn.Sequential(*layers) + + def forward(self, x: Tensor) -> Tensor: + return self.block(x) + + +class GeGluMlp(nn.Module): + def __init__(self, in_features: int, hidden_features: int): + super().__init__() + self.norm = nn.LayerNorm(in_features) + self.act = nn.GELU(approximate="tanh") + self.w0 = nn.Linear(in_features, hidden_features) + self.w1 = nn.Linear(in_features, hidden_features) + self.w2 = nn.Linear(hidden_features, in_features) + + def forward(self, x: Tensor) -> Tensor: + x = self.norm(x) + return self.w2(self.act(self.w0(x)) * self.w1(x)) + + +class CausalAttention(nn.Module): + """Causal self-attention that narrows the width from in_dim to out_dim. + + QKV is one bias-less linear; the checkpoint stores separate q/v bias + parameters and a frozen all-zero key bias buffer. The heads are + MEAN-pooled away (not concatenated) and the surviving head dim (256) is + adaptively average-pooled down to out_dim.""" + + def __init__(self, in_dim: int, out_dim: int, num_heads: int): + super().__init__() + self.num_heads = num_heads + self.head_dim = in_dim // num_heads + self.out_dim = out_dim + self.qkv = nn.Linear(in_dim, in_dim * 3, bias=False) + self.q_bias = nn.Parameter(torch.zeros(in_dim)) + self.v_bias = nn.Parameter(torch.zeros(in_dim)) + self.register_buffer("zero_k_bias", torch.zeros(in_dim)) + self.proj = nn.Linear(out_dim, out_dim) + + def forward(self, x: Tensor) -> Tensor: + b, t, _ = x.shape + bias = torch.cat((self.q_bias, self.zero_k_bias, self.v_bias)) + qkv = F.linear(x, self.qkv.weight, bias) + q, k, v = ( + qkv.view(b, t, 3, self.num_heads, self.head_dim) + .permute(2, 0, 3, 1, 4) + .unbind(0) + ) + x = F.scaled_dot_product_attention(q, k, v, is_causal=True) + x = x.mean(dim=1) + x = F.adaptive_avg_pool1d(x, self.out_dim) + return self.proj(x) + + +class AttnProjection(nn.Module): + """``pre_block``: rewires the 2048-wide encoder trunk to the 32-channel + latent width. Operates on ``[B, T, C]``.""" + + def __init__(self, in_dim: int, out_dim: int, num_heads: int, mlp_ratio: int = 2): + super().__init__() + self.norm1 = nn.LayerNorm(in_dim) + self.attn = CausalAttention(in_dim, out_dim, num_heads) + self.proj = nn.Linear(in_dim, out_dim) + self.norm3 = nn.LayerNorm(in_dim) + self.norm2 = nn.LayerNorm(out_dim) + self.mlp = GeGluMlp(out_dim, out_dim * mlp_ratio) + + def forward(self, x: Tensor) -> Tensor: + x = self.proj(self.norm3(x)) + self.attn(self.norm1(x)) + return x + self.mlp(self.norm2(x)) + + +class AMPBlock(nn.Module): + """BigVGAN AMPBlock1: per dilation a (dilated conv, dilation-1 conv) + pair, each conv preceded by its own alias-free SnakeBeta.""" + + def __init__(self, channels: int, kernel_size: int, dilations: tuple): + super().__init__() + self.convs1 = nn.ModuleList( + nn.Conv1d( + channels, + channels, + kernel_size, + dilation=d, + padding=(kernel_size * d - d) // 2, + ) + for d in dilations + ) + self.convs2 = nn.ModuleList( + nn.Conv1d(channels, channels, kernel_size, padding=(kernel_size - 1) // 2) + for _ in dilations + ) + self.activations = nn.ModuleList( + AliasFreeActivation1d(SnakeBeta(channels)) + for _ in range(2 * len(dilations)) + ) + + def forward(self, x: Tensor) -> Tensor: + for i, (conv1, conv2) in enumerate(zip(self.convs1, self.convs2)): + y = conv1(self.activations[2 * i](x)) + y = conv2(self.activations[2 * i + 1](y)) + x = x + y + return x + + +class BigVGANDecoder(nn.Module): + """``[B, in_channels, T] -> [B, 1, T * prod(rates)]``, clamped to [-1, 1].""" + + def __init__( + self, + in_channels: int = 2048, + initial_channels: int = 1024, + rates: tuple = (5, 5, 2, 2, 2, 2, 2), + kernel_sizes: tuple = (9, 9, 4, 4, 4, 4, 4), + resblock_kernel_sizes: tuple = (3, 7, 11), + resblock_dilations: tuple = ((1, 3, 5), (1, 3, 5), (1, 3, 5)), + ): + super().__init__() + self.num_kernels = len(resblock_kernel_sizes) + self.conv_pre = nn.Conv1d(in_channels, initial_channels, 7, padding=3) + + # each upsampler is nested in a one-element ModuleList so the state + # dict keys stay `ups..0.*` like the checkpoint + self.ups = nn.ModuleList() + self.resblocks = nn.ModuleList() + channels = initial_channels + for rate, kernel in zip(rates, kernel_sizes): + self.ups.append( + nn.ModuleList( + [ + nn.ConvTranspose1d( + channels, + channels // 2, + kernel, + stride=rate, + padding=(kernel - rate) // 2, + ) + ] + ) + ) + channels //= 2 + for k, dils in zip(resblock_kernel_sizes, resblock_dilations): + self.resblocks.append(AMPBlock(channels, k, dils)) + + self.activation_post = AliasFreeActivation1d(SnakeBeta(channels)) + self.conv_post = nn.Conv1d(channels, 1, 7, padding=3, bias=False) + + def forward(self, x: Tensor) -> Tensor: + x = self.conv_pre(x) + for i, up in enumerate(self.ups): + x = up[0](x) + acc = None + for j in range(self.num_kernels): + y = self.resblocks[i * self.num_kernels + j](x) + acc = y if acc is None else acc + y + x = acc / self.num_kernels + x = self.activation_post(x) + x = self.conv_post(x) + return torch.clamp(x, min=-1.0, max=1.0) + + +class MiniMaxH3AudioVAE(nn.Module): + """Frozen MiniMax-H3 audio autoencoder with diagonal latent normalization. + + - ``encode``: waveform ``(B, 1, samples)`` at 32 kHz -> latents ``(B, 32, T)`` + - ``decode``: latents ``(B, 32, T)`` -> waveform ``(B, 1, T * 800)`` + """ + + SAMPLE_RATE = 32000 + HOP_LENGTH = 800 + LATENT_CHANNELS = 32 + LATENTS_PER_SECOND = 40.0 + + def __init__(self): + super().__init__() + self.encoder = DacEncoder(d_model=64, strides=(2, 4, 4, 5, 5), d_latent=2048) + self.pre_block = AttnProjection(2048, 32, num_heads=8) + self.mean_proj = nn.Conv1d(32, 32, 1) + # log-STD head (not log-var); unused by encode(), kept for weight parity + self.logs_proj = nn.Conv1d(32, 32, 1) + self.dec_in_proj = nn.Conv1d(32, 2048, 1) + self.decoder = BigVGANDecoder() + + self.register_buffer( + "latents_mean", + torch.tensor(_LATENTS_MEAN, dtype=torch.float32), + persistent=False, + ) + self.register_buffer( + "latents_std", + torch.tensor(_LATENTS_STD, dtype=torch.float32), + persistent=False, + ) + + for p in self.parameters(): + p.requires_grad = False + self.eval() + + @property + def device(self): + return next(self.parameters()).device + + @property + def dtype(self): + return next(self.parameters()).dtype + + @property + def sample_rate(self) -> int: + return self.SAMPLE_RATE + + @property + def downsampling_ratio(self) -> int: + return self.HOP_LENGTH + + def _apply(self, fn, recurse=True): + # This VAE is pinned to fp32 (bf16 decodes are audibly degraded). + # Device moves pass through, but any cast that would land a float + # tensor in a non-fp32 dtype is rewritten to a device-move of the + # ORIGINAL tensor kept at fp32 -- the weights never round-trip + # through a lower precision. + def guarded(t): + out = fn(t) + if ( + torch.is_tensor(out) + and out.is_floating_point() + and out.dtype != torch.float32 + ): + return t.to(device=out.device, dtype=torch.float32) + return out + + return super()._apply(guarded, recurse) + + def encode(self, waveform: Tensor) -> Tensor: + """Mono waveform ``(B, 1, samples)`` at 32 kHz -> normalized latents + ``(B, 32, T)``. Uses the posterior mean (mode), like the released + pipeline; the tail is zero-padded to a multiple of 800 samples.""" + waveform = waveform.float() + remainder = waveform.shape[-1] % self.HOP_LENGTH + if remainder: + waveform = F.pad(waveform, (0, self.HOP_LENGTH - remainder)) + h = self.encoder(waveform) + h = self.pre_block(h.transpose(1, 2)).transpose(1, 2) + z = self.mean_proj(h) + return (z - self.latents_mean[:, None]) / self.latents_std[:, None] + + def decode(self, latents: Tensor) -> Tensor: + """Normalized latents ``(B, 32, T)`` -> waveform ``(B, 1, T * 800)`` + in [-1, 1], fp32.""" + latents = latents.float() + z = latents * self.latents_std[:, None] + self.latents_mean[:, None] + return self.decoder(self.dec_in_proj(z)) + + def forward(self, waveform: Tensor) -> Tensor: + return self.decode(self.encode(waveform)) diff --git a/extensions_built_in/diffusion_models/minimax_h3/src/packing.py b/extensions_built_in/diffusion_models/minimax_h3/src/packing.py new file mode 100644 index 00000000..154170e5 --- /dev/null +++ b/extensions_built_in/diffusion_models/minimax_h3/src/packing.py @@ -0,0 +1,439 @@ +"""Packed-sequence geometry for MiniMax-H3. + +One transformer forward runs over a single packed 1-D sequence: + + [ text (L) | keyframe conditions (C) | target audio (A) | target video (V) ] + +This module owns everything needed to place a row in that sequence and give it +its (t, h, w) rotary coordinate, plus the sigma-shift math that couples the +video (shift 12) and audio (shift 3) flow schedules. + +Rotary coordinates are built in float64 and with numpy's ``linspace`` because +video and audio share one 40-units-per-second rotary clock (video advances +5/3 units per pixel frame at 24 fps, audio one unit per latent at 40/s) and +that shared clock is the released checkpoint's audio/video alignment. +""" + +import math +from dataclasses import dataclass +from typing import List, Optional, Tuple + +import numpy as np +import torch +from PIL import Image + +# Per-row modality tags — these index the transformer's AdaLN table, so the +# values are a checkpoint contract. +VIDEO_TAG = 0 +TEXT_TAG = 1 +AUDIO_TAG = 2 +PAD_TAG = -1 + +FPS = 24 +SHORT_EDGE = 768 +MAX_PIXELS = 768 * 1344 +CANVAS_MULTIPLE = 32 + +# video VAE: 17 pixel frames per chunk -> 5 latent frames, 3 trailing latents +# dropped overall, so 17n+5 pixel frames <-> 5n+2 latent frames +FRAMES_PER_CHUNK = 17 +LATENTS_PER_CHUNK = 5 + +AUDIO_LATENTS_PER_SECOND = 40 +AUDIO_CHANNELS = 2 +AUDIO_SAMPLE_RATE = 32000 + +# released flow shifts (exponential): video 12, audio 3 +VIDEO_SIGMA_SHIFT = 12.0 +AUDIO_SIGMA_SHIFT = 3.0 + +# keyframe conditioning rows are noised to t = 0.999 and pinned there; the +# posterior sample of the keyframe VAE encode uses a fixed seed of 42 +KEYFRAME_NOISE_AUG_T = 0.999 +KEYFRAME_ENCODE_SEED = 42 + +# rotary-time constants: one latent frame spans 5/3 * frames_per_latent units, +# the (1, 4, 4, 4, 4) pattern mirroring the VAE's 17 -> 5 frame grouping +_ROPE_FRAME_RESCALE = 5.0 / 3.0 +_ROPE_FRAMES_PER_LATENT = (1, 4, 4, 4, 4) +_ROPE_SPATIAL_SCALE = 32 + + +# --------------------------------------------------------------------------- +# Frame / canvas arithmetic +# --------------------------------------------------------------------------- + + +def align_num_frames(num_frames: int) -> int: + """Snap a frame count UP to the next 17n+5 the video VAE can encode.""" + if num_frames < 1: + raise ValueError(f"num_frames must be positive, got {num_frames}") + while num_frames % FRAMES_PER_CHUNK != LATENTS_PER_CHUNK: + num_frames += 1 + return num_frames + + +def align_num_frames_down(num_frames: int) -> int: + """Snap a frame count DOWN to the previous 17n+5 (minimum 5).""" + num_frames = max(num_frames, LATENTS_PER_CHUNK) + while num_frames % FRAMES_PER_CHUNK != LATENTS_PER_CHUNK: + num_frames -= 1 + return num_frames + + +def video_latent_num_frames(num_frames: int) -> int: + """17n+5 pixel frames -> 5n+2 latent frames.""" + if num_frames % FRAMES_PER_CHUNK != LATENTS_PER_CHUNK: + raise ValueError(f"num_frames must be of the form 17n+5, got {num_frames}") + return (num_frames - LATENTS_PER_CHUNK) // FRAMES_PER_CHUNK * LATENTS_PER_CHUNK + 2 + + +def audio_latent_num_frames(num_frames: int) -> int: + """Audio latents covering `num_frames` video frames at 24 fps / 40 Hz.""" + return int(round(num_frames / FPS * AUDIO_LATENTS_PER_SECOND)) + + +def resolve_canvas_size(aspect_width: float, aspect_height: float) -> Tuple[int, int]: + """Aspect ratio -> (height, width): short edge 768, area capped at 768*1344, + both axes rounded to the nearest multiple of 32.""" + ratio = aspect_width / aspect_height + if ratio >= 1.0: + width, height = SHORT_EDGE * ratio, float(SHORT_EDGE) + else: + width, height = float(SHORT_EDGE), SHORT_EDGE / ratio + area = width * height + if area > MAX_PIXELS: + scale = (MAX_PIXELS / area) ** 0.5 + width, height = width * scale, height * scale + m = CANVAS_MULTIPLE + return max(m, round(height / m) * m), max(m, round(width / m) * m) + + +def prepare_keyframe_image( + image: Image.Image, height: int, width: int, stretch: bool = True +): + """Put a keyframe onto the target canvas: the geometry anchor is stretched, + a follower keyframe is cover-cropped.""" + if image.size == (width, height): + return image + if stretch: + return image.resize((width, height), Image.Resampling.LANCZOS) + scale = max(width / image.size[0], height / image.size[1]) + resized_size = ( + max(width, round(image.size[0] * scale)), + max(height, round(image.size[1] * scale)), + ) + left = max(0, (resized_size[0] - width) // 2) + top = max(0, (resized_size[1] - height) // 2) + resized = image.resize(resized_size, Image.Resampling.LANCZOS) + return resized.crop((left, top, left + width, top + height)) + + +# --------------------------------------------------------------------------- +# Row packing +# --------------------------------------------------------------------------- + + +def patchify_video_latents(latents: torch.Tensor, patch_size=(1, 2, 2)) -> torch.Tensor: + """(B, C, T, H, W) -> (B, N, C * prod(patch)) rows, frame-major then + row-major, feature order [c, pt, ph, pw].""" + pt, ph, pw = patch_size + b, c, t, h, w = latents.shape + latents = latents.reshape(b, c, t // pt, pt, h // ph, ph, w // pw, pw) + latents = latents.permute(0, 2, 4, 6, 1, 3, 5, 7) + return latents.reshape(b, -1, c * pt * ph * pw).contiguous() + + +def unpatchify_video_tokens( + rows: torch.Tensor, + num_latent_frames: int, + latent_height: int, + latent_width: int, + channels: int = 24, + patch_size=(1, 2, 2), +) -> torch.Tensor: + """(B, N, C * prod(patch)) -> (B, C, T, H, W). Inverse of patchify.""" + pt, ph, pw = patch_size + b = rows.shape[0] + rows = rows.reshape( + b, + num_latent_frames // pt, + latent_height // ph, + latent_width // pw, + channels, + pt, + ph, + pw, + ) + rows = rows.permute(0, 4, 1, 5, 2, 6, 3, 7) + return rows.reshape( + b, channels, num_latent_frames, latent_height, latent_width + ).contiguous() + + +def pack_audio_latents(latents: torch.Tensor) -> torch.Tensor: + """(B, 2, C, T) stereo audio latents -> (B, 2*T, C) channel-major rows + (all T frames of channel 0, then channel 1).""" + return ( + latents.permute(0, 1, 3, 2) + .reshape(latents.shape[0], -1, latents.shape[2]) + .contiguous() + ) + + +def unpack_audio_tokens(rows: torch.Tensor, num_audio_latents: int) -> torch.Tensor: + """(B, 2*T, C) channel-major rows -> (B, 2, C, T).""" + b, _, c = rows.shape + rows = rows.reshape(b, AUDIO_CHANNELS, num_audio_latents, c) + return rows.permute(0, 1, 3, 2).contiguous() + + +# --------------------------------------------------------------------------- +# Rotary grids (float64, numpy linspace — the released grid must reproduce) +# --------------------------------------------------------------------------- + + +def _spatial_position_grid(dim: int, patch: int, sqrt_area: float) -> torch.Tensor: + ratio = dim / sqrt_area + left = (1.0 - ratio) / 2.0 + # numpy linspace(endpoint=False) is start + arange(n) * (stop-start)/n, + # which is not bit-identical to torch.linspace + grid = ( + np.linspace(left, left + ratio, dim // patch, endpoint=False) + * _ROPE_SPATIAL_SCALE + ) + return torch.from_numpy(grid).to(torch.float64) + + +def _temporal_position_grid(num_latent_frames: int, origin: float) -> torch.Tensor: + spans = torch.tensor( + [ + _ROPE_FRAME_RESCALE + * _ROPE_FRAMES_PER_LATENT[i % len(_ROPE_FRAMES_PER_LATENT)] + for i in range(num_latent_frames) + ], + dtype=torch.float64, + ) + return origin + torch.cat( + [torch.zeros(1, dtype=torch.float64), spans[:-1].cumsum(0)] + ) + + +def _temporal_position_span(num_latent_frames: int) -> float: + # numpy pairwise sum on purpose: the reference computes the "last" keyframe + # anchor this way and the summation orders differ in the last ulp + spans = np.ones(num_latent_frames, dtype=np.float64) * _ROPE_FRAME_RESCALE + for i in range(len(_ROPE_FRAMES_PER_LATENT)): + spans[i :: len(_ROPE_FRAMES_PER_LATENT)] *= _ROPE_FRAMES_PER_LATENT[i] + return float(spans.sum()) + + +# --------------------------------------------------------------------------- +# Sequence layout +# --------------------------------------------------------------------------- + + +@dataclass +class PackedLayout: + """Structural description of one packed sequence (one batch item).""" + + sequence_length: int + position_ids: torch.Tensor # (S, 3) float64 + token_tags: torch.Tensor # (S,) long + video_indices: torch.Tensor # condition rows first, then target rows + audio_indices: torch.Tensor + text_indices: torch.Tensor + num_condition_video_rows: int + + +def build_packed_sequence( + text_token_tags: torch.Tensor, # (L,) long: 1 text, 0 for vision-block rows + num_latent_frames: int, + latent_height: int, + latent_width: int, + num_audio_latents: int, + patch_size=(1, 2, 2), + keyframe_anchors: Tuple[str, ...] = (), +) -> PackedLayout: + """Build the [text | keyframe conditions | target audio | target video] + layout used by t2va and fl2va.""" + _, ph, pw = patch_size + rows_per_frame = (latent_height // ph) * (latent_width // pw) + num_text = int(text_token_tags.shape[0]) + num_cond = len(keyframe_anchors) * rows_per_frame + num_audio_rows = num_audio_latents * AUDIO_CHANNELS + num_video_rows = num_latent_frames * rows_per_frame + seq_len = num_text + num_cond + num_audio_rows + num_video_rows + + cond_start = num_text + audio_start = cond_start + num_cond + video_start = audio_start + num_audio_rows + + # text rows sit on the time axis at their row index; the media clock + # continues from there, so prompt length shifts the whole media clock + position_ids = torch.zeros(seq_len, 3, dtype=torch.float64) + position_ids[:num_text, 0] = torch.arange(num_text, dtype=torch.float64) + + sqrt_area = math.sqrt(latent_height * latent_width) + height_grid = _spatial_position_grid(latent_height, ph, sqrt_area) + width_grid = _spatial_position_grid(latent_width, pw, sqrt_area) + frame_grid = torch.stack( + [g.reshape(-1) for g in torch.meshgrid(height_grid, width_grid, indexing="ij")], + dim=-1, + ) + + for i, anchor in enumerate(keyframe_anchors): + if anchor == "first": + anchor_time = float(num_text) + elif anchor == "last": + anchor_time = ( + float(num_text) + + _temporal_position_span(num_latent_frames) + - _ROPE_FRAME_RESCALE + ) + else: + raise ValueError( + f"keyframe anchor must be 'first' or 'last', got {anchor!r}" + ) + rows = slice( + cond_start + i * rows_per_frame, cond_start + (i + 1) * rows_per_frame + ) + position_ids[rows, 0] = anchor_time + position_ids[rows, 1:] = frame_grid + + # audio rows: channel-major, one rotary unit per latent (40/s = 24fps*5/3), + # no height coordinate, width pinned to the grid extremes per channel + audio_time = float(num_text) + torch.arange(num_audio_latents, dtype=torch.float64) + position_ids[audio_start:video_start, 0] = audio_time.repeat(AUDIO_CHANNELS) + position_ids[audio_start:video_start, 2] = torch.cat( + [ + torch.full((num_audio_latents,), float(width_grid[0]), dtype=torch.float64), + torch.full( + (num_audio_latents,), float(width_grid[-1]), dtype=torch.float64 + ), + ] + ) + + video_pos = torch.empty(num_latent_frames, rows_per_frame, 3, dtype=torch.float64) + video_pos[:, :, 0] = _temporal_position_grid(num_latent_frames, float(num_text))[ + :, None + ] + video_pos[:, :, 1:] = frame_grid[None] + position_ids[video_start:] = video_pos.reshape(-1, 3) + + video_indices = torch.cat( + [torch.arange(cond_start, audio_start), torch.arange(video_start, seq_len)] + ) + audio_indices = torch.arange(audio_start, video_start) + text_indices = torch.arange(num_text) + + token_tags = torch.empty(seq_len, dtype=torch.long) + token_tags[text_indices] = text_token_tags.to(torch.long) + token_tags[audio_indices] = AUDIO_TAG + token_tags[video_indices] = VIDEO_TAG + + return PackedLayout( + sequence_length=seq_len, + position_ids=position_ids, + token_tags=token_tags, + video_indices=video_indices, + audio_indices=audio_indices, + text_indices=text_indices, + num_condition_video_rows=num_cond, + ) + + +def build_row_timesteps( + layout: PackedLayout, + video_timestep: float, + audio_timestep: float, + condition_video_timestep: Optional[float] = None, +) -> torch.Tensor: + """Per-row timestep values (S,) float32. Text rows inherit the video + timestep; condition video rows stay pinned at their noise-aug level.""" + if condition_video_timestep is None: + condition_video_timestep = max(video_timestep, KEYFRAME_NOISE_AUG_T) + row_t = torch.full( + (layout.sequence_length,), float(video_timestep), dtype=torch.float32 + ) + row_t[layout.video_indices[: layout.num_condition_video_rows]] = float( + condition_video_timestep + ) + row_t[layout.audio_indices] = float(audio_timestep) + return row_t + + +def pad_layouts_to_batch(layouts: List[PackedLayout]): + """Stack per-item layouts that share the same media geometry but may have + different text lengths into batched transformer inputs. + + Items are right-padded in the TEXT segment to the batch's max text length + (pad rows tagged -1, masked out of attention as keys, positions zero). + Returns (position_ids (B, S, 3) f64, token_tags (B, S) long, + video_indices, audio_indices, text_indices, pad_counts) where the index + tensors describe the shared padded layout: [text_max | cond | audio | video]. + """ + max_text = max(int(l.text_indices.shape[0]) for l in layouts) + ref = layouts[0] + media_len = ref.sequence_length - int(ref.text_indices.shape[0]) + for l in layouts: + if l.sequence_length - int(l.text_indices.shape[0]) != media_len: + raise ValueError( + "all layouts in a batch must share the same media geometry" + ) + seq_len = max_text + media_len + + b = len(layouts) + position_ids = torch.zeros(b, seq_len, 3, dtype=torch.float64) + token_tags = torch.full((b, seq_len), PAD_TAG, dtype=torch.long) + pad_counts = [] + for i, l in enumerate(layouts): + lt = int(l.text_indices.shape[0]) + position_ids[i, :lt] = l.position_ids[:lt] + position_ids[i, max_text:] = l.position_ids[lt:] + token_tags[i, :lt] = l.token_tags[:lt] + token_tags[i, max_text:] = l.token_tags[lt:] + pad_counts.append(max_text - lt) + + offset = max_text - int(ref.text_indices.shape[0]) + video_indices = ref.video_indices + offset + audio_indices = ref.audio_indices + offset + text_indices = torch.arange(max_text) + return ( + position_ids, + token_tags, + video_indices, + audio_indices, + text_indices, + pad_counts, + ) + + +# --------------------------------------------------------------------------- +# Sigma-shift math (video shift 12, audio shift 3, exponential) +# --------------------------------------------------------------------------- + + +def shift_sigma(sigma, shift: float): + """Exponential timeshift: shift * sigma / (1 + (shift - 1) * sigma).""" + return shift * sigma / (1.0 + (shift - 1.0) * sigma) + + +def remap_sigma( + sigma, from_shift: float = VIDEO_SIGMA_SHIFT, to_shift: float = AUDIO_SIGMA_SHIFT +): + """Map a sigma on the `from_shift` schedule onto the `to_shift` schedule + at the same underlying schedule position (the video/audio coupling).""" + base = sigma / (from_shift + sigma * (1.0 - from_shift)) + return shift_sigma(base, to_shift) + + +def build_sigma_schedule( + num_inference_steps: int, shift: float = VIDEO_SIGMA_SHIFT +) -> torch.Tensor: + """The released sampling grid: linspace(1, 0, steps) through the + exponential shift, consecutive duplicates collapsed — the terminal 0 is + part of the count, so `steps` yields `steps - 1` model evaluations.""" + base = torch.linspace(1.0, 0.0, num_inference_steps, dtype=torch.float32) + sigmas = shift_sigma(base, shift) + return torch.unique_consecutive(sigmas) diff --git a/extensions_built_in/diffusion_models/minimax_h3/src/pipeline.py b/extensions_built_in/diffusion_models/minimax_h3/src/pipeline.py new file mode 100644 index 00000000..630a5a69 --- /dev/null +++ b/extensions_built_in/diffusion_models/minimax_h3/src/pipeline.py @@ -0,0 +1,215 @@ +"""MiniMax-H3 sampling pipeline for ai-toolkit training previews. + +Covers t2v (t2va) and first-frame i2v (fl2va), always denoising the joint +audio stream alongside the video (the packed sequence contains audio rows by +construction; decoding the audio track is optional). + +MiniMax-H3 is guidance-distilled: there is no negative prompt, no CFG and +exactly one transformer forward per step. ``unconditional_embeds`` and +``guidance_scale`` are accepted for harness compatibility and ignored. + +Scheduler (the released math, not diffusers'): + - sigma grid: ``linspace(1, 0, steps)`` through the exponential shift + (video 12, audio 3), consecutive duplicates collapsed; the terminal 0 is + part of the count so ``steps`` yields ``steps - 1`` model evaluations + - the model consumes ``t = 1 - sigma`` (t = 1 means clean) and predicts the + data-ward velocity ``clean - noise``: ``denoised = x + sigma * v`` + - Euler update ``x_next = r * x + (1 - r) * denoised`` with + ``r = sigma_next / sigma``, evaluated in float32 +""" + +from typing import Optional + +import numpy as np +import torch +from PIL import Image +from diffusers.utils.torch_utils import randn_tensor + +from . import packing +from .packing import ( + AUDIO_CHANNELS, + AUDIO_SIGMA_SHIFT, + FPS, + KEYFRAME_NOISE_AUG_T, + VIDEO_SIGMA_SHIFT, + build_packed_sequence, + build_row_timesteps, + build_sigma_schedule, + pack_audio_latents, + patchify_video_latents, + remap_sigma, + unpack_audio_tokens, + unpatchify_video_tokens, +) + + +class MiniMaxH3Pipeline: + """Lightweight sampler; receives the MinimaxH3Model (BaseModel subclass) + and reuses its VAEs / transformer / device bookkeeping.""" + + def __init__(self, model): + self.model = model + + @property + def device(self): + return self.model.device_torch + + def to(self, *args, **kwargs): + return self + + def set_progress_bar_config(self, **kwargs): + pass + + @torch.no_grad() + def __call__( + self, + conditional_embeds, # AdvancedPromptEmbeds: text_embeds [(L, 5120)], text_token_tags [(L,)] + unconditional_embeds=None, # ignored: MiniMax-H3 is guidance-distilled + height: int = 768, + width: int = 768, + num_frames: int = 124, + num_inference_steps: int = 28, + guidance_scale: float = 1.0, # ignored + latents: Optional[torch.Tensor] = None, + generator: Optional[torch.Generator] = None, + ctrl_img: Optional[ + Image.Image + ] = None, # first-frame keyframe, already canvas-sized + with_audio: bool = True, + **kwargs, + ): + model = self.model + device = model.device_torch + dtype = model.torch_dtype + transformer = model.transformer + + is_video = num_frames > 1 + if is_video: + num_frames = packing.align_num_frames_down(num_frames) + else: + # a single image still runs through the video path with the + # minimum 5-frame clip; the first frame is returned + num_frames = packing.LATENTS_PER_CHUNK + + t_lat = packing.video_latent_num_frames(num_frames) + h_lat = height // 16 + w_lat = width // 16 + a_lat = packing.audio_latent_num_frames(num_frames) + + text_embeds = conditional_embeds.text_embeds[0].to(device, dtype) + token_tags = conditional_embeds.text_token_tags[0].to("cpu", torch.long) + + # --- packed layout ------------------------------------------------- + anchors = ("first",) if ctrl_img is not None else () + layout = build_packed_sequence( + text_token_tags=token_tags, + num_latent_frames=t_lat, + latent_height=h_lat, + latent_width=w_lat, + num_audio_latents=a_lat, + keyframe_anchors=anchors, + ) + num_cond = layout.num_condition_video_rows + + # --- conditioning rows (draw order: condition noise, video, audio) -- + cond_rows = None + if ctrl_img is not None: + cond_noise = randn_tensor( + (1, 24, 1, h_lat, w_lat), generator=generator, dtype=torch.float32 + ).to(device) + frame = torch.from_numpy(np.array(ctrl_img)).float() + frame = (frame / 255.0) * 2.0 - 1.0 # (H, W, 3) -> [-1, 1] + frame = frame.permute(2, 0, 1)[None, :, None] # (1, 3, 1, H, W) + cond_latents = model.encode_keyframe_latents(frame) # (1, 24, 1, h, w) fp32 + # released noise-aug recipe: x = t * clean + (1 - t) * noise at t = 0.999 + cond_latents = ( + KEYFRAME_NOISE_AUG_T * cond_latents.to(device) + + (1.0 - KEYFRAME_NOISE_AUG_T) * cond_noise + ) + cond_rows = patchify_video_latents(cond_latents) # (1, rows, 96) + + # --- initial noise ------------------------------------------------- + if latents is None: + latents = randn_tensor( + (1, 24, t_lat, h_lat, w_lat), generator=generator, dtype=torch.float32 + ) + video_rows = patchify_video_latents(latents.to(device).float()) # (1, V, 96) + audio_noise = randn_tensor( + (1, AUDIO_CHANNELS, 32, a_lat), generator=generator, dtype=torch.float32 + ).to(device) + audio_rows = pack_audio_latents(audio_noise) # (1, 2*A, 32) + + # --- schedules ----------------------------------------------------- + sigmas_v = build_sigma_schedule(num_inference_steps, VIDEO_SIGMA_SHIFT).to( + device + ) + # the audio schedule follows the video grid through the closed-form + # shift remap so both streams sit at the same underlying position + sigmas_a = remap_sigma(sigmas_v, VIDEO_SIGMA_SHIFT, AUDIO_SIGMA_SHIFT) + + position_ids = layout.position_ids[None].to(device) + tags = layout.token_tags[None].to(device) + video_indices = layout.video_indices.to(device) + audio_indices = layout.audio_indices.to(device) + text_indices = layout.text_indices.to(device) + + # --- denoise loop -------------------------------------------------- + num_steps = sigmas_v.shape[0] - 1 + for i in range(num_steps): + sv, sv_next = sigmas_v[i], sigmas_v[i + 1] + sa, sa_next = sigmas_a[i], sigmas_a[i + 1] + t_v = 1.0 - float(sv) + t_a = 1.0 - float(sa) + + row_t = build_row_timesteps(layout, t_v, t_a)[None].to(device) + + video_in = video_rows + if cond_rows is not None: + video_in = torch.cat([cond_rows, video_rows], dim=1) + + video_pred, audio_pred = transformer( + hidden_states=video_in.to(dtype), + audio_hidden_states=audio_rows.to(dtype), + encoder_hidden_states=text_embeds[None], + row_timesteps=row_t, + token_tags=tags, + position_ids=position_ids, + video_indices=video_indices, + audio_indices=audio_indices, + text_indices=text_indices, + ) + v_video = video_pred[:, num_cond:].float() + v_audio = audio_pred.float() + + denoised_v = video_rows + sv * v_video + ratio_v = sv_next / sv + video_rows = ratio_v * video_rows + (1.0 - ratio_v) * denoised_v + + denoised_a = audio_rows + sa * v_audio + ratio_a = sa_next / sa if float(sa) != 0.0 else 0.0 + audio_rows = ratio_a * audio_rows + (1.0 - ratio_a) * denoised_a + + # --- decode -------------------------------------------------------- + video_latents = unpatchify_video_tokens(video_rows, t_lat, h_lat, w_lat) + video = model.decode_latents(video_latents) # (1, 3, T, H, W) in [-1, 1] + video = ((video.float().clamp(-1, 1) + 1.0) * 127.5).round().to(torch.uint8) + video = video[0].permute(1, 2, 3, 0).cpu() # (T, H, W, C) + + if not is_video: + return [Image.fromarray(video[0].numpy())] + + audio_out = None + if with_audio: + audio_latents = unpack_audio_tokens(audio_rows, a_lat)[0] # (2, 32, A) + waveform = model.decode_audio_latents( + audio_latents.float() + ) # (2, 1, samples) + audio_out = waveform[:, 0].cpu() # (2, samples) stereo + + return { + "video": video, + "fps": FPS, + "audio": audio_out, + "audio_sample_rate": packing.AUDIO_SAMPLE_RATE, + "output_path": None, + } diff --git a/extensions_built_in/diffusion_models/minimax_h3/src/text_encoder.py b/extensions_built_in/diffusion_models/minimax_h3/src/text_encoder.py new file mode 100644 index 00000000..0685042d --- /dev/null +++ b/extensions_built_in/diffusion_models/minimax_h3/src/text_encoder.py @@ -0,0 +1,106 @@ +"""Qwen3-VL conditioning for MiniMax-H3. + +MiniMax-H3 conditions on the **unnormalized** ``hidden_states[50]`` of its +Qwen3-VL-32B conditioner (``hidden_states[0]`` is the embedding output, so +this is the output of decoder layer 49, before the final norm). The LM head +and layers 50..63 are never used, which lets the loader truncate the stack. + +The presentation is raw tokens — no chat template, no special tokens: + + - t2va: the verbatim prompt. + - fl2va: per keyframe, a ``": "`` label plus a vision block + (``<|vision_start|>`` + one ``<|image_pad|>`` per merged vision patch + + ``<|vision_end|>``), then the verbatim prompt. Vision-block rows are + tagged as *video* (0) rather than text (1) — the transformer's AdaLN + modality selection keys off these tags. +""" + +from typing import List, Optional + +import torch + +from .packing import TEXT_TAG, VIDEO_TAG + +TEXT_ENCODER_LAYER = 50 + + +@torch.no_grad() +def encode_minimax_h3_prompt( + text_encoder, # transformers Qwen3VLForConditionalGeneration + tokenizer, # Qwen2TokenizerFast + processor, # Qwen3VLProcessor (needed only when keyframes are present) + prompt: str, + keyframes: Optional[List] = None, # PIL images already on the target canvas + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, +): + """Encode ONE prompt (with optional keyframes) into MiniMax-H3 conditioning. + + Returns (embeds (L, 5120), token_tags (L,) long). The embeds come from + ``hidden_states[50]`` unnormalized. A stack truncated to exactly 50 layers + also works ONLY if the final ``model.norm`` has been replaced with an + Identity (transformers applies the final norm to the last entry of + ``hidden_states``); the loader in minimax_h3.py does exactly that. + """ + num_layers = text_encoder.config.text_config.num_hidden_layers + if num_layers < TEXT_ENCODER_LAYER: + raise ValueError( + f"MiniMax-H3 needs at least {TEXT_ENCODER_LAYER} Qwen3-VL decoder " + f"layers to read hidden_states[{TEXT_ENCODER_LAYER}], got {num_layers}" + ) + if device is None: + device = text_encoder.device + + pixel_values, image_grid_thw = None, None + token_ids: List[int] = [] + token_tags: List[int] = [] + if keyframes: + vision = processor.image_processor(images=keyframes, return_tensors="pt") + pixel_values = vision["pixel_values"] + image_grid_thw = vision["image_grid_thw"] + merge = processor.image_processor.merge_size**2 + vision_start = tokenizer.convert_tokens_to_ids("<|vision_start|>") + vision_end = tokenizer.convert_tokens_to_ids("<|vision_end|>") + image_pad = tokenizer.convert_tokens_to_ids("<|image_pad|>") + for i in range(len(keyframes)): + num_image_tokens = int(image_grid_thw[i].prod()) // merge + label_ids = tokenizer(f": ", add_special_tokens=False)[ + "input_ids" + ] + vision_ids = [vision_start] + [image_pad] * num_image_tokens + [vision_end] + token_ids += label_ids + vision_ids + token_tags += [TEXT_TAG] * len(label_ids) + [VIDEO_TAG] * len(vision_ids) + + prompt_ids = tokenizer(prompt, add_special_tokens=False)["input_ids"] + token_ids += prompt_ids + token_tags += [TEXT_TAG] * len(prompt_ids) + if len(token_ids) == 0: + # empty (unconditional) prompt: a single pad token keeps the sequence + # non-degenerate; the model was not trained with CFG so this is only + # ever a fallback + token_ids = [tokenizer.pad_token_id or 0] + token_tags = [TEXT_TAG] + + input_ids = torch.tensor([token_ids], dtype=torch.long, device=device) + mm_token_type_ids = torch.tensor( + processor.create_mm_token_type_ids([token_ids]), dtype=torch.long, device=device + ) + + # call the inner .model directly: the LM head's vocab projection is dead + # weight here and hidden_states[50] is all that is consumed + outputs = text_encoder.model( + input_ids=input_ids, + attention_mask=torch.ones_like(input_ids), + mm_token_type_ids=mm_token_type_ids, + pixel_values=None + if pixel_values is None + else pixel_values.to(device, text_encoder.dtype), + image_grid_thw=None if image_grid_thw is None else image_grid_thw.to(device), + use_cache=False, + output_hidden_states=True, + ) + layer = min(TEXT_ENCODER_LAYER, len(outputs.hidden_states) - 1) + embeds = outputs.hidden_states[layer][0] + if dtype is not None: + embeds = embeds.to(dtype) + return embeds, torch.tensor(token_tags, dtype=torch.long) diff --git a/extensions_built_in/diffusion_models/minimax_h3/src/transformer.py b/extensions_built_in/diffusion_models/minimax_h3/src/transformer.py new file mode 100644 index 00000000..9fc0bdb4 --- /dev/null +++ b/extensions_built_in/diffusion_models/minimax_h3/src/transformer.py @@ -0,0 +1,480 @@ +"""MiniMax-H3 diffusion transformer (33B), weight-compatible with the original +``MiniMaxAI/MiniMax-H3`` checkpoint keys (``video_patch_proj``, ``blocks.N.*``, +``final_layer.*``, ...). + +MiniMax-H3 runs one stack of blocks over a single packed 1-D sequence holding +``[text | keyframe-condition video rows | audio rows | target video rows]``. +Full self-attention, no cross-attention, no per-modality weights: modality- +specific behavior comes only from the two input patch projections, the per-row +AdaLN modality tag (0 = video, 1 = text, 2 = audio) and the two output heads. + +Deviations from the reference implementations (both are pure generalizations — +with batch size 1 and no padding the math is identical): + + - every structural input (``position_ids``, ``token_tags``, ``row_timesteps``) + carries a batch axis, so items with different prompt lengths can share a + training batch (pad rows are tagged ``-1`` and masked out of attention) + - timesteps are passed per row and deduplicated internally instead of the + caller supplying ``timestep`` + ``timestep_indices`` + +Mixed precision follows the shipped checkpoint: ``video_patch_proj``, +``audio_patch_proj``, ``time_embedder``, ``final_layer.video_out/audio_out`` +and ``rope.inv_freq`` are float32 islands, everything else runs bf16. Inputs +are aligned to each projection's own parameter dtype at the call site. + +Timesteps are consumed unscaled in [0, 1] with t = 1 - sigma (t=1 means clean), +and both heads predict the data-ward velocity ``clean - noise``. +""" + +import math +from dataclasses import dataclass +from typing import Optional, Tuple + +import torch +import torch.nn.functional as F +from torch import nn +from torch.utils.checkpoint import checkpoint + +MODALITY_NUM = 3 # 0 = video, 1 = text, 2 = audio; -1 marks padding rows + + +@dataclass +class MiniMaxH3TransformerParams: + hidden_size: int = 5376 + num_layers: int = 50 + token_refiner_num_layers: int = 2 + num_attention_heads: int = 56 + attention_head_dim: int = 128 # heads * head_dim = 7168 > hidden_size + ffn_hidden_size: int = 14336 + latents_dim: int = 24 + audio_latents_dim: int = 32 + patch_size: Tuple[int, int, int] = (1, 2, 2) + text_dim: int = 5120 + timestep_input_dim: int = 256 + time_embed_hidden_size: int = 5376 + time_embed_dim: int = 2688 + rope_inv_freq_len: int = 16 + rope_theta: float = 10000.0 + norm_eps: float = 1e-5 + qk_norm_eps: float = 1e-5 + final_norm_eps: float = 1e-5 + # "pruned" checkpoints (e.g. Comfy-Org *_pruned_*) replace the timestep + # MLP with a small lookup table: ``adaln_t_table`` of shape + # (adaln_t_table_size, time_embed_dim) sampled by linear interpolation at + # t * (size - 1), consumed by the AdaLN projections WITHOUT the SiLU. + # These checkpoints also shrink time_embed_dim (8 in the released files). + adaln_t_table_size: Optional[int] = None + + @property + def adaln_apply_silu(self) -> bool: + return self.adaln_t_table_size is None + + +class MiniMaxH3Rope(nn.Module): + """3-axis rotary embedding over the packed (t, h, w) coordinates. + + One shared ``inv_freq`` of 16 frequencies per axis; the three 16-angle + blocks concatenate to 48 and duplicate to 96, so the first 96 of the 128 + head channels rotate (rotate-half convention) and the last 32 pass through. + """ + + def __init__(self, inv_freq_len: int = 16, theta: float = 10000.0): + super().__init__() + dim = 2 * inv_freq_len + inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) + # present in the checkpoint, so persistent + self.register_buffer("inv_freq", inv_freq, persistent=True) + + def forward(self, position_ids: torch.Tensor): + """position_ids (B, S, 3) -> cos, sin each (B, S, 96), float32.""" + position_ids = position_ids.to(device=self.inv_freq.device, dtype=torch.float32) + freqs = position_ids.unsqueeze(-1) * self.inv_freq.view(1, 1, 1, -1) + # (B, S, 3, 16) -> (B, S, 48) in (t, h, w) axis order -> duplicate to 96 + freqs = freqs.flatten(2, 3) + freqs = torch.cat([freqs, freqs], dim=-1) + return freqs.cos(), freqs.sin() + + +def apply_rotary_emb(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor): + """x (B, S, H, D); cos/sin (B, S, rot) rotate the leading ``rot`` channels.""" + rot = cos.shape[-1] + x_rot, x_pass = x[..., :rot], x[..., rot:] + cos = cos.to(x.dtype).unsqueeze(2) + sin = sin.to(x.dtype).unsqueeze(2) + x1, x2 = x_rot.chunk(2, dim=-1) + rotated = torch.cat([-x2, x1], dim=-1) + return torch.cat([x_rot * cos + rotated * sin, x_pass], dim=-1) + + +class MiniMaxH3TimeEmbedder(nn.Module): + """Sinusoidal embedding (cos before sin, unscaled t in [0, 1]) -> MLP. + + A float32 island in the checkpoint; the output stays float32 so every + block's AdaLN applies its SiLU at full precision before casting down. + """ + + def __init__(self, freq_dim: int, hidden: int, out_dim: int): + super().__init__() + self.freq_dim = freq_dim + self.proj_in = nn.Linear(freq_dim, hidden, bias=True) + self.proj_out = nn.Linear(hidden, out_dim, bias=True) + + def forward(self, t: torch.Tensor) -> torch.Tensor: + half = self.freq_dim // 2 + freqs = torch.exp( + -math.log(10000.0) + * torch.arange(half, dtype=torch.float32, device=t.device) + / half + ) + args = t.to(torch.float32)[:, None] * freqs[None] + emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) + emb = emb.to(self.proj_in.weight.dtype) + return self.proj_out(F.silu(self.proj_in(emb))) + + +class MiniMaxH3Attention(nn.Module): + """Fused-QKV self-attention with per-head RMSNorm on q/k and partial RoPE.""" + + def __init__(self, hidden: int, heads: int, head_dim: int, qk_norm_eps: float): + super().__init__() + self.heads = heads + self.head_dim = head_dim + inner = heads * head_dim + self.qkv_proj = nn.Linear(hidden, inner * 3, bias=False) + self.q_norm = nn.RMSNorm(head_dim, eps=qk_norm_eps) + self.k_norm = nn.RMSNorm(head_dim, eps=qk_norm_eps) + self.out_proj = nn.Linear(inner, hidden, bias=False) + + def forward( + self, + x: torch.Tensor, # (B, S, hidden) + rotary_emb=None, # (cos, sin) each (B, S, rot) or None + attn_mask: Optional[torch.Tensor] = None, # (B, 1, 1, S) bool, True = attend + ) -> torch.Tensor: + b, s, _ = x.shape + q, k, v = self.qkv_proj(x).chunk(3, dim=-1) + q = q.view(b, s, self.heads, self.head_dim) + k = k.view(b, s, self.heads, self.head_dim) + v = v.view(b, s, self.heads, self.head_dim) + + q = self.q_norm(q) + k = self.k_norm(k) + if rotary_emb is not None: + q = apply_rotary_emb(q, *rotary_emb) + k = apply_rotary_emb(k, *rotary_emb) + + q = q.transpose(1, 2) + k = k.transpose(1, 2) + v = v.transpose(1, 2) + out = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask) + out = out.transpose(1, 2).reshape(b, s, -1) + return self.out_proj(out) + + +class MiniMaxH3Mlp(nn.Module): + """SwiGLU: fc1 packs [gate | up]; silu(gate) * up -> fc2. No biases.""" + + def __init__(self, hidden: int, ffn: int): + super().__init__() + self.fc1 = nn.Linear(hidden, ffn * 2, bias=False) + self.fc2 = nn.Linear(ffn, hidden, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + gate, up = self.fc1(x).chunk(2, dim=-1) + return self.fc2(F.silu(gate) * up) + + +class MiniMaxH3AdalnProj(nn.Module): + """Timestep embedding -> per-(timestep, modality) modulation parameters. + + (M, time_embed_dim) -> ``expand`` tensors of (M * modalities, hidden); row + layout ``[t0_mod0, t0_mod1, t0_mod2, t1_mod0, ...]``, addressed by + ``timestep_index * MODALITY_NUM + tag``. The SiLU runs at temb's own + (float32) precision; only its result is cast to the projection dtype. + """ + + def __init__( + self, + t_dim: int, + hidden: int, + expand: int, + modalities: int, + apply_silu: bool = True, + ): + super().__init__() + self.expand = expand + self.modalities = modalities + self.hidden = hidden + self.apply_silu = apply_silu + self.linear = nn.Linear(t_dim, expand * hidden * modalities, bias=True) + + def forward(self, temb: torch.Tensor): + if self.apply_silu: + temb = F.silu(temb) + x = self.linear(temb.to(self.linear.weight.dtype)) + x = x.view(x.shape[0] * self.modalities, self.expand * self.hidden) + return x.chunk(self.expand, dim=-1) + + +class MiniMaxH3RefinerBlock(nn.Module): + """Plain pre-norm block over the projected text stream. No AdaLN, no RoPE.""" + + def __init__(self, p: MiniMaxH3TransformerParams): + super().__init__() + self.norm1 = nn.RMSNorm(p.hidden_size, eps=p.norm_eps) + self.norm2 = nn.RMSNorm(p.hidden_size, eps=p.norm_eps) + self.attn = MiniMaxH3Attention( + p.hidden_size, p.num_attention_heads, p.attention_head_dim, p.qk_norm_eps + ) + self.mlp = MiniMaxH3Mlp(p.hidden_size, p.ffn_hidden_size) + + def forward(self, x: torch.Tensor, attn_mask=None) -> torch.Tensor: + x = x + self.attn(self.norm1(x), attn_mask=attn_mask) + x = x + self.mlp(self.norm2(x)) + return x + + +class MiniMaxH3TokenRefiner(nn.Module): + def __init__(self, p: MiniMaxH3TransformerParams): + super().__init__() + self.blocks = nn.ModuleList( + [MiniMaxH3RefinerBlock(p) for _ in range(p.token_refiner_num_layers)] + ) + self.final_norm = nn.RMSNorm(p.hidden_size, eps=p.final_norm_eps) + self.gradient_checkpointing = False + + def forward(self, x: torch.Tensor, attn_mask=None) -> torch.Tensor: + for block in self.blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + x = checkpoint(block, x, attn_mask, use_reentrant=False) + else: + x = block(x, attn_mask) + return self.final_norm(x) + + +class MiniMaxH3Block(nn.Module): + """Pre-norm attention + SwiGLU MLP, each modulated by AdaLN parameters + gathered per row from the (timestep, modality) table.""" + + def __init__(self, p: MiniMaxH3TransformerParams): + super().__init__() + self.norm1 = nn.RMSNorm(p.hidden_size, eps=p.norm_eps) + self.norm2 = nn.RMSNorm(p.hidden_size, eps=p.norm_eps) + self.attn = MiniMaxH3Attention( + p.hidden_size, p.num_attention_heads, p.attention_head_dim, p.qk_norm_eps + ) + self.mlp = MiniMaxH3Mlp(p.hidden_size, p.ffn_hidden_size) + self.adaln_proj = MiniMaxH3AdalnProj( + p.time_embed_dim, + p.hidden_size, + expand=6, + modalities=MODALITY_NUM, + apply_silu=p.adaln_apply_silu, + ) + + def forward( + self, + x: torch.Tensor, # (B, S, hidden) + temb: torch.Tensor, # (M, time_embed_dim) float32 + adaln_indices: torch.Tensor, # (B, S) long into the (M * 3) table + rotary_emb, # (cos, sin) + attn_mask: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( + self.adaln_proj(temb) + ) + dt = x.dtype # pruned checkpoints store the adaln projections fp16 + + h = self.norm1(x) * (1.0 + scale_msa[adaln_indices].to(dt)) + shift_msa[ + adaln_indices + ].to(dt) + x = x + gate_msa[adaln_indices].to(dt) * self.attn(h, rotary_emb, attn_mask) + + h = self.norm2(x) * (1.0 + scale_mlp[adaln_indices].to(dt)) + shift_mlp[ + adaln_indices + ].to(dt) + x = x + gate_mlp[adaln_indices].to(dt) * self.mlp(h) + return x + + +class MiniMaxH3FinalLayer(nn.Module): + """Shared shift/scale-modulated RMSNorm + the two per-modality heads. + + Both heads run over every row (matching the reference); the caller selects + each modality's rows from the results. Heads are float32 islands. + """ + + def __init__(self, p: MiniMaxH3TransformerParams): + super().__init__() + video_patch_dim = ( + p.latents_dim * p.patch_size[0] * p.patch_size[1] * p.patch_size[2] + ) + self.norm = nn.RMSNorm(p.hidden_size, eps=p.final_norm_eps) + self.adaln_proj = MiniMaxH3AdalnProj( + p.time_embed_dim, + p.hidden_size, + expand=2, + modalities=1, + apply_silu=p.adaln_apply_silu, + ) + self.video_out = nn.Linear(p.hidden_size, video_patch_dim, bias=True) + self.audio_out = nn.Linear(p.hidden_size, p.audio_latents_dim, bias=True) + + def forward( + self, + x: torch.Tensor, # (B, S, hidden) + temb: torch.Tensor, # (M, time_embed_dim) + timestep_indices: torch.Tensor, # (B, S) long into the M rows + ): + shift, scale = self.adaln_proj(temb) + dt = x.dtype + h = self.norm(x) * (1.0 + scale[timestep_indices].to(dt)) + shift[ + timestep_indices + ].to(dt) + h = h.to(self.video_out.weight.dtype) + return self.video_out(h), self.audio_out(h) + + +class MiniMaxH3Transformer(nn.Module): + def __init__(self, params: Optional[MiniMaxH3TransformerParams] = None): + super().__init__() + if params is None: + params = MiniMaxH3TransformerParams() + self.params = params + p = params + + video_patch_dim = ( + p.latents_dim * p.patch_size[0] * p.patch_size[1] * p.patch_size[2] + ) + self.video_patch_proj = nn.Linear(video_patch_dim, p.hidden_size, bias=True) + self.audio_patch_proj = nn.Linear(p.audio_latents_dim, p.hidden_size, bias=True) + self.condition_proj = nn.Linear(p.text_dim, p.hidden_size, bias=True) + + if p.adaln_t_table_size is not None: + # pruned checkpoints: the timestep MLP is replaced by a lookup + # table sampled with linear interpolation at t * (size - 1) + self.time_embedder = None + self.register_buffer( + "adaln_t_table", + torch.zeros(p.adaln_t_table_size, p.time_embed_dim), + persistent=True, + ) + else: + self.time_embedder = MiniMaxH3TimeEmbedder( + p.timestep_input_dim, p.time_embed_hidden_size, p.time_embed_dim + ) + self.rope = MiniMaxH3Rope(p.rope_inv_freq_len, p.rope_theta) + self.token_refiner = MiniMaxH3TokenRefiner(p) + self.blocks = nn.ModuleList([MiniMaxH3Block(p) for _ in range(p.num_layers)]) + self.final_layer = MiniMaxH3FinalLayer(p) + + self.gradient_checkpointing = False + + # float32 islands of the shipped checkpoint; used by the loader to keep + # these keys at full precision when the rest is cast to bf16 + FP32_KEY_PREFIXES = ( + "video_patch_proj", + "audio_patch_proj", + "time_embedder", + "final_layer.video_out", + "final_layer.audio_out", + "rope", + ) + + def enable_gradient_checkpointing(self, enable: bool = True): + self.gradient_checkpointing = enable + self.token_refiner.gradient_checkpointing = enable + + def disable_gradient_checkpointing(self): + self.enable_gradient_checkpointing(False) + + @property + def device(self): + return self.condition_proj.weight.device + + @property + def dtype(self): + # working dtype of the block stack (norm weights are never quantized) + return self.token_refiner.final_norm.weight.dtype + + def _time_embedding(self, t: torch.Tensor) -> torch.Tensor: + """(M,) float32 in [0, 1] -> (M, time_embed_dim) float32.""" + if self.time_embedder is not None: + return self.time_embedder(t) + table = self.adaln_t_table.float() + pos = t.clamp(0.0, 1.0) * (table.shape[0] - 1) + lo = pos.floor().long() + hi = (lo + 1).clamp(max=table.shape[0] - 1) + frac = (pos - lo.float()).unsqueeze(1) + return table[lo] * (1.0 - frac) + table[hi] * frac + + def forward( + self, + hidden_states: torch.Tensor, # (B, Nv, 96) patchified video rows (cond + target) + audio_hidden_states: torch.Tensor, # (B, Na, 32) audio rows + encoder_hidden_states: torch.Tensor, # (B, L, 5120) text conditioning + row_timesteps: torch.Tensor, # (B, S) float in [0, 1], t = 1 - sigma + token_tags: torch.Tensor, # (B, S) long: 0 video, 1 text, 2 audio, -1 pad + position_ids: torch.Tensor, # (B, S, 3) float (t, h, w) rotary coords + video_indices: torch.Tensor, # (Nv,) long positions of video rows in the pack + audio_indices: torch.Tensor, # (Na,) long + text_indices: torch.Tensor, # (L,) long + ): + """Returns (video_out (B, Nv, 96), audio_out (B, Na, 32)) — the + data-ward velocity ``clean - noise`` for every row, in input order. + Conditioning rows come back unmasked; discarding them is the caller's + job.""" + batch_size, seq_len = token_tags.shape + + rotary_emb = self.rope(position_ids) + + video_embeds = self.video_patch_proj( + hidden_states.to(self.video_patch_proj.weight.dtype) + ) + audio_embeds = self.audio_patch_proj( + audio_hidden_states.to(self.audio_patch_proj.weight.dtype) + ) + text_embeds = self.condition_proj(encoder_hidden_states.to(self.dtype)) + + # pad rows never act as attention keys; as queries they see everything + # (their outputs are discarded), which keeps SDPA rows finite + attn_mask = None + text_attn_mask = None + is_pad = token_tags < 0 + if bool(is_pad.any()): + live = ~is_pad + attn_mask = live[:, None, None, :] + text_attn_mask = live[:, text_indices][:, None, None, :] + + text_embeds = self.token_refiner(text_embeds, text_attn_mask) + + x = text_embeds.new_zeros((batch_size, seq_len, text_embeds.shape[-1])) + x = x.index_copy(1, text_indices, text_embeds) + x = x.index_copy(1, video_indices, video_embeds.to(x.dtype)) + x = x.index_copy(1, audio_indices, audio_embeds.to(x.dtype)) + + # one timestep embedding per distinct noise level across the batch + unique_t, inverse = torch.unique( + row_timesteps.to(torch.float32), sorted=True, return_inverse=True + ) + temb = self._time_embedding(unique_t) + adaln_indices = inverse * MODALITY_NUM + token_tags.clamp(min=0) + + for block in self.blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + x = checkpoint( + block, + x, + temb, + adaln_indices, + rotary_emb, + attn_mask, + use_reentrant=False, + ) + else: + x = block(x, temb, adaln_indices, rotary_emb, attn_mask) + + video_all, audio_all = self.final_layer(x, temb, inverse) + video_out = video_all.index_select(1, video_indices) + audio_out = audio_all.index_select(1, audio_indices) + return video_out, audio_out diff --git a/extensions_built_in/diffusion_models/minimax_h3/src/vae.py b/extensions_built_in/diffusion_models/minimax_h3/src/vae.py new file mode 100644 index 00000000..9880f054 --- /dev/null +++ b/extensions_built_in/diffusion_models/minimax_h3/src/vae.py @@ -0,0 +1,786 @@ +"""MiniMax-H3 video VAE -- causal 3D-CNN encoder + non-causal 36-layer ViT decoder. + +16x spatial / 4x temporal compression, 24 latent channels, per-channel +latents_mean/latents_std normalization. The temporal geometry is fixed by +``clip_length`` (17 pixel frames per encoder chunk) and ``token_drop`` (3 +trailing latent frames dropped after encoding): ``17n + 5`` pixel frames +map to ``5n + 2`` latent frames. ``decode`` mirrors the chunking -- 5-token +chunks with a 2-token overlap, cross-faded over 5 pixel frames. A single +frame (image / i2v keyframe) skips the temporal chunking entirely and maps +to a single latent frame. + +The reference pixel convention is ImageNet-normalized RGB over [0, 1]; +``encode``/``decode`` here speak ai-toolkit's [-1, 1] and convert inside. +Spatial tiling (256 px tiles, >= 64 px overlap, linearly blended) is ON by +default for both encode and decode -- the model was released that way and +the released frames are the blended-tile ones. The checkpoint is float32 +and the released decode recipe is fp16 autocast over the fp32 weights, with +the token embedder, output projection, and every norm kept in fp32. +""" + +import math +from contextlib import nullcontext + +import torch +import torch.nn as nn +import torch.nn.functional as F + +IMAGENET_MEAN = (0.485, 0.456, 0.406) +IMAGENET_STD = (0.229, 0.224, 0.225) + +LATENTS_MEAN = [ + 0.858090341091156, + -0.9606591463088989, + 1.0661640167236328, + -0.5090325474739075, + -0.2727581858634949, + -1.3675414323806763, + -0.2553254961967468, + -0.26907554268836975, + -0.5376840829849243, + -0.0464097298681736, + 0.6657370328903198, + 0.19690127670764923, + -0.5460608005523682, + -0.4035342037677765, + -0.23683024942874908, + 0.25928452610969543, + -0.30133944749832153, + 0.211341992020607, + -1.1206848621368408, + 0.3581933379173279, + -0.04225143790245056, + 0.2604829967021942, + 0.22864092886447906, + 0.7056031823158264, +] + +LATENTS_STD = [ + 1.2223774194717407, + 1.2767263650894165, + 1.68317747116088865, + 1.7549455165863037, + 1.5636216402053833, + 2.194143533706665, + 0.96531379222869875, + 1.05698859691619875, + 0.841948926448822, + 0.7729952931404114, + 1.8955937623977661, + 0.946841835975647, + 0.7996809482574463, + 0.44988900423049925, + 0.7197399735450745, + 0.69362932443618775, + 2.961095094680786, + 2.7694199085235595, + 3.0496184825897215, + 2.1088054180145265, + 3.276226282119751, + 3.1627357006073, + 2.28168129920959475, + 2.6127843856811525, +] + + +# --------------------------------------------------------------------------- +# causal 3D-CNN encoder +# --------------------------------------------------------------------------- + + +class CausalConv3d(nn.Conv3d): + """Conv3d with symmetric reflect spatial padding and causal (front-only, + zeros) temporal padding of ``kernel_t - 1`` frames.""" + + def __init__(self, in_channels, out_channels, kernel_size, stride=1, spatial_pad=0): + super().__init__( + in_channels, out_channels, kernel_size=kernel_size, stride=stride + ) + self.spatial_pad = spatial_pad + + def forward(self, x): + if self.spatial_pad > 0: + p = self.spatial_pad + x = F.pad(x, (p, p, p, p, 0, 0), mode="reflect") + t_pad = self.kernel_size[0] - 1 + if t_pad > 0: + x = F.pad(x, (0, 0, 0, 0, t_pad, 0)) + return F.conv3d(x, self.weight, self.bias, stride=self.stride) + + +class FrameGroupNorm(nn.GroupNorm): + """GroupNorm with statistics computed per frame (``use_t_isolated_gn``): + the temporal axis is folded into the batch axis.""" + + def forward(self, x): + b, c, t, h, w = x.shape + x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) + x = super().forward(x) + return x.view(b, t, c, h, w).permute(0, 2, 1, 3, 4) + + +class ResnetBlock3d(nn.Module): + def __init__(self, in_channels, out_channels, norm_num_groups=32): + super().__init__() + self.norm1 = FrameGroupNorm(norm_num_groups, in_channels, eps=1e-6) + self.conv1 = CausalConv3d(in_channels, out_channels, 3, spatial_pad=1) + self.norm2 = FrameGroupNorm(norm_num_groups, out_channels, eps=1e-6) + self.conv2 = CausalConv3d(out_channels, out_channels, 3, spatial_pad=1) + if in_channels != out_channels: + self.nin_shortcut = CausalConv3d(in_channels, out_channels, 1) + else: + self.nin_shortcut = None + + def forward(self, x): + h = self.conv1(F.silu(self.norm1(x))) + h = self.conv2(F.silu(self.norm2(h))) + if self.nin_shortcut is not None: + x = self.nin_shortcut(x) + return x + h + + +class Downsample3d(nn.Module): + """Strided 3x3x3 conv; a spatial stride of 2 is preceded by an asymmetric + bottom/right reflect pad of 1 so the output is exactly ceil(size / 2).""" + + def __init__(self, channels, time_stride, space_stride): + super().__init__() + self.space_stride = space_stride + self.conv = CausalConv3d( + channels, channels, 3, stride=(time_stride, space_stride, space_stride) + ) + + def forward(self, x): + if self.space_stride == 2: + x = F.pad(x, (0, 1, 0, 1, 0, 0), mode="reflect") + return self.conv(x) + + +class Encoder3d(nn.Module): + def __init__( + self, + in_channels, + out_channels, + block_out_channels, + layers_per_block, + space_down, + time_down, + norm_num_groups, + ): + super().__init__() + block_in = (block_out_channels[0],) + tuple(block_out_channels[:-1]) + + self.conv_in = CausalConv3d(in_channels, block_in[0], 3, spatial_pad=1) + + self.down = nn.ModuleList() + for i in range(len(block_out_channels)): + level = nn.Module() + level.block = nn.ModuleList( + ResnetBlock3d( + block_in[i] if j == 0 else block_out_channels[i], + block_out_channels[i], + norm_num_groups=norm_num_groups, + ) + for j in range(layers_per_block) + ) + if space_down[i] * time_down[i] > 1: + level.downsample = Downsample3d( + block_out_channels[i], time_down[i], space_down[i] + ) + self.down.append(level) + + self.norm_out = FrameGroupNorm( + norm_num_groups, block_out_channels[-1], eps=1e-6 + ) + self.conv_out = CausalConv3d( + block_out_channels[-1], out_channels, 3, spatial_pad=1 + ) + + def forward(self, x): + h = self.conv_in(x) + for level in self.down: + for block in level.block: + h = block(h) + if hasattr(level, "downsample"): + h = level.downsample(h) + return self.conv_out(F.silu(self.norm_out(h))) + + +# --------------------------------------------------------------------------- +# non-causal ViT decoder +# --------------------------------------------------------------------------- + + +class RMSNormFP32(nn.RMSNorm): + """The reference normalizes in float32 regardless of the compute dtype + (or the storage dtype — the Comfy repack ships this VAE in fp16).""" + + def forward(self, x): + weight = None if self.weight is None else self.weight.float() + out = F.rms_norm(x.float(), self.normalized_shape, weight, self.eps) + return out.to(x.dtype) + + +class RotaryEmbedding3d(nn.Module): + """3-axis rotary table. Coordinates are length-normalized to [-1, 1) per + axis and scaled by 2*pi; the per-axis angles are concatenated and then + duplicated, so the first ``rope_dim_ratio * head_dim`` channels of every + head are rotated (rotate-half convention).""" + + def __init__(self, dim, theta=100.0, num_axes=3): + super().__init__() + inv_freq = 1.0 / theta ** torch.arange( + 0, 1, 2 * num_axes / dim, dtype=torch.float32 + ) + self.register_buffer("inv_freq", inv_freq, persistent=False) + + def forward(self, position_ids): + angles = ( + 2.0 + * math.pi + * position_ids[:, :, :, None].float() + * self.inv_freq[None, None, None, :] + ) + angles = angles.flatten(2, 3).tile(2).unsqueeze(2) # (B, S, 1, rot_dim) + return angles.cos(), angles.sin() + + +def _rotate_half(x): + x1, x2 = x.chunk(2, dim=-1) + return torch.cat([-x2, x1], dim=-1) + + +def _apply_rope(x, cos, sin): + rot = cos.shape[-1] + x_rot, x_pass = x[..., :rot], x[..., rot:] + x_rot = x_rot * cos + _rotate_half(x_rot) * sin + return torch.cat([x_rot, x_pass], dim=-1) + + +class Attention(nn.Module): + def __init__(self, heads, head_dim, eps=1e-5): + super().__init__() + self.heads = heads + self.head_dim = head_dim + inner = heads * head_dim + self.norm_q = RMSNormFP32(head_dim, eps=eps, elementwise_affine=False) + self.norm_k = RMSNormFP32(head_dim, eps=eps, elementwise_affine=False) + self.to_qkv = nn.Linear(inner, inner * 3, bias=True) + self.to_out = nn.Linear(inner, inner, bias=True) + + def forward(self, x, rotary_emb): + b, s, _ = x.shape + # per-head-interleaved qkv layout: each head's 3 * head_dim slab is + # split into q | k | v, NOT [Q_all | K_all | V_all] + qkv = self.to_qkv(x).view(b, s, self.heads, 3 * self.head_dim) + q, k, v = qkv.chunk(3, dim=-1) + q = self.norm_q(q) + k = self.norm_k(k) + cos, sin = rotary_emb + cos = cos.to(q.dtype) + sin = sin.to(q.dtype) + q = _apply_rope(q, cos, sin) + k = _apply_rope(k, cos, sin) + out = F.scaled_dot_product_attention( + q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) + ) + return self.to_out(out.transpose(1, 2).reshape(b, s, -1)) + + +class FeedForward(nn.Module): + """Gated SiLU: w1 output chunks into gate | value, silu(gate) * value.""" + + def __init__(self, dim, mult=4): + super().__init__() + inner = dim * mult + self.w1 = nn.Linear(dim, inner * 2, bias=True) + self.w2 = nn.Linear(inner, dim, bias=True) + + def forward(self, x): + gate, x = self.w1(x).chunk(2, dim=-1) + return self.w2(F.silu(gate) * x) + + +class TransformerBlock(nn.Module): + def __init__(self, heads, head_dim, ffn_mult=4, eps=1e-5): + super().__init__() + dim = heads * head_dim + self.norm1 = RMSNormFP32(dim, eps=eps, elementwise_affine=True) + self.attn = Attention(heads, head_dim, eps=eps) + self.scale1 = nn.Parameter(torch.zeros(dim)) + self.norm2 = RMSNormFP32(dim, eps=eps, elementwise_affine=True) + self.ff = FeedForward(dim, mult=ffn_mult) + self.scale2 = nn.Parameter(torch.zeros(dim)) + + def forward(self, x, rotary_emb): + x = x + self.attn(self.norm1(x), rotary_emb) * self.scale1 + return x + self.ff(self.norm2(x)) * self.scale2 + + +class ViTDecoder3d(nn.Module): + """Every latent voxel is one token; 4 learned register tokens plus one + all-zero token are appended (all at rope coordinate 0), attended over with + full self-attention, and dropped before the patch projection expands each + token into a patch_size_t x patch_size x patch_size pixel block.""" + + def __init__( + self, + in_channels, + out_channels, + patch_size, + patch_size_t, + num_layers, + heads, + head_dim, + num_register_tokens=4, + ffn_mult=4, + rope_theta=100.0, + rope_dim_ratio=0.75, + eps=1e-5, + ): + super().__init__() + dim = heads * head_dim + self.patch_size = patch_size + self.patch_size_t = patch_size_t + self.out_channels = out_channels + self.num_register_tokens = num_register_tokens + + self.rope = RotaryEmbedding3d(int(head_dim * rope_dim_ratio), theta=rope_theta) + self.x_embedder = nn.Linear(in_channels, dim) + self.register_tokens = nn.Parameter(torch.zeros(1, num_register_tokens, dim)) + # unused at inference; kept so the checkpoint loads strict + self.register_buffer("mask_token", torch.zeros(1, 1, dim)) + self.transformer_blocks = nn.ModuleList( + TransformerBlock(heads, head_dim, ffn_mult=ffn_mult, eps=eps) + for _ in range(num_layers) + ) + self.norm_out = nn.LayerNorm(dim, eps=eps) + self.proj_out = nn.Linear( + dim, out_channels * patch_size_t * patch_size * patch_size + ) + + def forward(self, z): + b, c, t, h, w = z.shape + tokens = z.permute(0, 2, 3, 4, 1).reshape(b, t * h * w, c) + + # released recipe: token embedder / output projection / norms compute + # fp32 even when the matmuls run fp16 (autocast or fp16 storage); the + # block stack then runs at the storage dtype + with torch.autocast(device_type=z.device.type, enabled=False): + x = F.linear( + tokens.float(), + self.x_embedder.weight.float(), + self.x_embedder.bias.float(), + ) + x = x.to(self.x_embedder.weight.dtype) + num_patches = x.shape[1] + + x = torch.cat( + [ + x, + self.register_tokens.expand(b, -1, -1).to(x.dtype), + torch.zeros_like(x[:, :1]), + ], + dim=1, + ) + + grids = [ + 2.0 * (torch.arange(0.5, size, dtype=torch.float32, device=z.device) / size) + - 1.0 + for size in (t, h, w) + ] + position_ids = torch.stack( + torch.meshgrid(*grids, indexing="ij"), dim=-1 + ).flatten(0, 2) + position_ids = position_ids.unsqueeze(0).expand(b, -1, -1) + suffix_ids = position_ids.new_zeros((b, self.num_register_tokens + 1, 3)) + with torch.autocast(device_type=z.device.type, enabled=False): + rotary_emb = self.rope(torch.cat([position_ids, suffix_ids], dim=1)) + + for block in self.transformer_blocks: + x = block(x, rotary_emb) + + with torch.autocast(device_type=z.device.type, enabled=False): + x = F.layer_norm( + x.float(), + self.norm_out.normalized_shape, + self.norm_out.weight.float(), + self.norm_out.bias.float(), + self.norm_out.eps, + ) + x = F.linear(x, self.proj_out.weight.float(), self.proj_out.bias.float()) + x = x[:, :num_patches] + + pt, ps = self.patch_size_t, self.patch_size + x = x.view(b, t, h, w, self.out_channels, pt, ps, ps) + x = x.permute(0, 4, 1, 5, 2, 6, 3, 7) + return x.reshape(b, self.out_channels, t * pt, h * ps, w * ps) + + +# --------------------------------------------------------------------------- +# full VAE +# --------------------------------------------------------------------------- + + +class MiniMaxH3VideoVAE(nn.Module): + def __init__( + self, + in_channels=3, + out_channels=3, + latent_channels=24, + block_out_channels=(128, 256, 256, 512, 512, 1024), + layers_per_block=2, + space_down=(2, 2, 2, 2, 1, 1), + time_down=(1, 2, 2, 1, 1, 1), + norm_num_groups=32, + decoder_num_layers=36, + decoder_heads=32, + decoder_head_dim=64, + decoder_num_register_tokens=4, + decoder_ffn_mult=4, + rope_theta=100.0, + rope_dim_ratio=0.75, + clip_length=17, + token_drop=3, + latents_mean=None, + latents_std=None, + tile_size=256, + tile_overlap_min=64, + tiling=True, + ): + super().__init__() + self.latent_channels = latent_channels + self.spatial_compression = int(math.prod(space_down)) + self.temporal_compression = int(math.prod(time_down)) + self.clip_length = clip_length + self.token_drop = token_drop + + # clip_length is not a multiple of the temporal ratio, so each decoded + # chunk carries an implicit leading pad; token_drop leaves an overlap + # between consecutive chunks that gets cross-faded in pixel space + self.frame_pre_padding = (-clip_length) % self.temporal_compression + self.tokens_chunk_size = math.ceil(clip_length / self.temporal_compression) + self.token_overlap = (-token_drop) % self.tokens_chunk_size + self.frame_overlap = max( + self.token_overlap * self.temporal_compression - self.frame_pre_padding, 0 + ) + + self.use_tiling = tiling + self.tile_size = tile_size + self.tile_overlap_min = tile_overlap_min + + self.encoder = Encoder3d( + in_channels=in_channels, + out_channels=2 * latent_channels, + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + space_down=space_down, + time_down=time_down, + norm_num_groups=norm_num_groups, + ) + self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, 1) + self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, 1) + self.decoder = ViTDecoder3d( + in_channels=latent_channels, + out_channels=out_channels, + patch_size=self.spatial_compression, + patch_size_t=self.temporal_compression, + num_layers=decoder_num_layers, + heads=decoder_heads, + head_dim=decoder_head_dim, + num_register_tokens=decoder_num_register_tokens, + ffn_mult=decoder_ffn_mult, + rope_theta=rope_theta, + rope_dim_ratio=rope_dim_ratio, + ) + + if latents_mean is None: + latents_mean = ( + LATENTS_MEAN + if latent_channels == len(LATENTS_MEAN) + else [0.0] * latent_channels + ) + if latents_std is None: + latents_std = ( + LATENTS_STD + if latent_channels == len(LATENTS_STD) + else [1.0] * latent_channels + ) + self.register_buffer( + "latents_mean", + torch.tensor(latents_mean, dtype=torch.float32), + persistent=False, + ) + self.register_buffer( + "latents_std", + torch.tensor(latents_std, dtype=torch.float32), + persistent=False, + ) + self.register_buffer( + "pixel_mean", + torch.tensor(IMAGENET_MEAN, dtype=torch.float32).view(1, 3, 1, 1, 1), + persistent=False, + ) + self.register_buffer( + "pixel_std", + torch.tensor(IMAGENET_STD, dtype=torch.float32).view(1, 3, 1, 1, 1), + persistent=False, + ) + + @property + def device(self): + return next(self.parameters()).device + + @property + def dtype(self): + return next(self.parameters()).dtype + + @staticmethod + def latent_frames(num_pixel_frames: int) -> int: + """17n + 5 pixel frames -> 5n + 2 latent frames; 1 -> 1 (keyframe).""" + if num_pixel_frames == 1: + return 1 + return 5 * math.ceil(num_pixel_frames / 17) - 3 + + @staticmethod + def pixel_frames(num_latent_frames: int) -> int: + """5n + 2 latent frames -> 17n + 5 pixel frames; 1 -> 1 (keyframe).""" + if num_latent_frames == 1: + return 1 + n, rem = divmod(num_latent_frames - 2, 5) + if rem != 0 or n < 0: + raise ValueError( + f"latent frame count must be 5n + 2 or 1, got {num_latent_frames}" + ) + return 17 * n + 5 + + # -- spatial tiling ------------------------------------------------------ + + def _split_tiles(self, length): + """Lay tile_size-wide tiles over length pixels; the slack beyond the + minimum overlap is spread round-robin in whole spatial_compression + steps so every tile boundary stays latent-aligned.""" + tile_size = self.tile_size + if tile_size >= length: + return [0], [length], [] + + num_tiles = math.ceil(length / tile_size) + while tile_size * num_tiles - self.tile_overlap_min * (num_tiles - 1) < length: + num_tiles += 1 + + overlaps = [self.tile_overlap_min] * (num_tiles - 1) + remaining = tile_size * num_tiles - sum(overlaps) - length + for i in range(remaining // self.spatial_compression): + overlaps[i % (num_tiles - 1)] += self.spatial_compression + + starts = [0] + for i in range(num_tiles - 1): + starts.append(starts[-1] + tile_size - overlaps[i]) + return starts, [tile_size] * num_tiles, overlaps + + @staticmethod + def _blend(a, b, blend_extent, dim): + blend_extent = min(a.shape[dim], b.shape[dim], blend_extent) + positions = torch.arange(blend_extent, device=b.device, dtype=b.dtype) + shape = [1] * b.ndim + shape[dim] = blend_extent + weight_b = (positions / blend_extent).view(shape) + + slice_a = [slice(None)] * a.ndim + slice_a[dim] = slice(-blend_extent, None) + slice_b = [slice(None)] * b.ndim + slice_b[dim] = slice(0, blend_extent) + blended = a[tuple(slice_a)] * (1 - weight_b) + b[tuple(slice_b)] * weight_b + + if blend_extent == b.shape[dim]: + return blended + slice_rest = [slice(None)] * b.ndim + slice_rest[dim] = slice(blend_extent, None) + return torch.cat([blended, b[tuple(slice_rest)]], dim=dim) + + def _stitch_tiles(self, rows, y_overlaps, x_overlaps): + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + if i > 0: + tile = self._blend(rows[i - 1][j], tile, y_overlaps[i - 1], dim=-2) + if j > 0: + tile = self._blend(row[j - 1], tile, x_overlaps[j - 1], dim=-1) + if i < len(rows) - 1: + tile = tile[..., : -y_overlaps[i], :] + if j < len(row) - 1: + tile = tile[..., :, : -x_overlaps[j]] + result_row.append(tile) + result_rows.append(torch.cat(result_row, dim=-1)) + return torch.cat(result_rows, dim=-2) + + def _encode_clip(self, x): + """Encode one temporal clip to moments, spatially tiled when enabled.""" + if not self.use_tiling: + return self.quant_conv(self.encoder(x)) + + y_starts, y_lens, y_overlaps = self._split_tiles(x.shape[-2]) + x_starts, x_lens, x_overlaps = self._split_tiles(x.shape[-1]) + rows = [] + for i_pos, i_len in zip(y_starts, y_lens): + row = [] + for j_pos, j_len in zip(x_starts, x_lens): + tile = x[..., i_pos : i_pos + i_len, j_pos : j_pos + j_len] + row.append(self.quant_conv(self.encoder(tile))) + rows.append(row) + ratio = self.spatial_compression + return self._stitch_tiles( + rows, [o // ratio for o in y_overlaps], [o // ratio for o in x_overlaps] + ) + + def _decode_clip(self, z): + """Decode one temporal clip; tiles are laid out in pixel space and + mapped back onto the latent grid, blended in pixel space.""" + if not self.use_tiling: + return self.decoder(self.post_quant_conv(z)) + + ratio = self.spatial_compression + y_starts, y_lens, y_overlaps = self._split_tiles(z.shape[-2] * ratio) + x_starts, x_lens, x_overlaps = self._split_tiles(z.shape[-1] * ratio) + rows = [] + for i_pos, i_len in zip(y_starts, y_lens): + row = [] + for j_pos, j_len in zip(x_starts, x_lens): + tile = z[ + ..., + i_pos // ratio : (i_pos + i_len) // ratio, + j_pos // ratio : (j_pos + j_len) // ratio, + ] + row.append(self.decoder(self.post_quant_conv(tile))) + rows.append(row) + return self._stitch_tiles(rows, y_overlaps, x_overlaps) + + # -- temporal chunking --------------------------------------------------- + + def _encode_video(self, x): + if x.shape[2] % self.clip_length != 0: + pad = x[:, :, -1:].repeat(1, 1, (-x.shape[2]) % self.clip_length, 1, 1) + x = torch.cat([x, pad], dim=2) + moments = torch.cat( + [ + self._encode_clip( + x[:, :, i * self.clip_length : (i + 1) * self.clip_length] + ) + for i in range(x.shape[2] // self.clip_length) + ], + dim=2, + ) + if self.token_drop > 0: + moments = moments[:, :, : -self.token_drop] + return moments + + def _decode_video(self, z): + tcs = self.tokens_chunk_size + ratio_t = self.temporal_compression + chunk_frames = tcs * ratio_t + split_count = 2 if self.token_drop > 0 else 1 + + num_tokens = z.shape[2] + self.token_drop + pad_tokens = (-num_tokens) % tcs + num_chunks = (num_tokens + pad_tokens) // tcs - (split_count - 1) + if num_chunks < 1: + # too few tokens for one chunk (e.g. the 2-latent minimum clip) + pad_tokens += tcs + num_chunks += 1 + if pad_tokens > 0: + z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2) + + decoded = [] + overlap = None + for i in range(num_chunks): + start = i * tcs + clip = self._decode_clip(z[:, :, start : start + tcs + self.token_overlap]) + for j in range(split_count): + part = clip[:, :, j * chunk_frames : (j + 1) * chunk_frames] + part = part[:, :, self.frame_pre_padding :] + if j == 0: + if overlap is not None: + part = self._blend(overlap, part, self.frame_overlap, dim=-3) + decoded.append(part) + else: + overlap = part + if overlap is not None: + decoded.append(overlap) + dec = torch.cat(decoded, dim=2) + + if pad_tokens > 0: + # repeated latent frames produced trailing pixel frames that were + # never requested; a chunk's last token only covers + # clip_length % ratio_t frames, the others cover ratio_t + intra_tail = self.clip_length % ratio_t + before_pad = z.shape[2] - pad_tokens + pad_frames = sum( + intra_tail if intra_tail and (before_pad + k) % tcs == 0 else ratio_t + for k in range(pad_tokens) + ) + dec = dec[:, :, :-pad_frames] + return dec + + # -- public interface ---------------------------------------------------- + + def encode(self, pixels, sample=True, generator=None, fp16_round=False): + """pixels (B, 3, T, H, W) in [-1, 1], T == 17n + 5, or T == 1 (single + keyframe: spatial encode only, no temporal chunking). Returns + normalized latents (B, latent_channels, t, h, w): posterior sampled + (or the mean when sample=False), optionally rounded through fp16 + BEFORE the (z - latents_mean) / latents_std normalization + (fp16_round=True is the released first-frame-conditioning recipe).""" + x = pixels + if x.ndim == 4: + x = x.unsqueeze(2) + x = (x.float() + 1.0) * 0.5 + x = (x - self.pixel_mean) / self.pixel_std + x = x.to(self.dtype) + + if x.shape[2] == 1: + moments = self._encode_clip(x)[:, :, -1:] + else: + moments = self._encode_video(x) + + mean, logvar = moments.float().chunk(2, dim=1) + if sample: + std = torch.exp(0.5 * logvar.clamp(-30.0, 20.0)) + noise = torch.randn( + mean.shape, + generator=generator, + device=generator.device if generator is not None else mean.device, + dtype=mean.dtype, + ).to(mean.device) + z = mean + std * noise + else: + z = mean + + if fp16_round: + z = z.to(torch.float16).float() + shape = (1, -1, 1, 1, 1) + return (z - self.latents_mean.view(shape)) / self.latents_std.view(shape) + + def decode(self, latents, autocast_fp16=True): + """Normalized latents (B, latent_channels, t, h, w) -> pixels + (B, 3, T, H, W) in [-1, 1]. Runs under fp16 autocast over the fp32 + weights on cuda (the released recipe) unless autocast_fp16=False.""" + shape = (1, -1, 1, 1, 1) + z = latents.float() * self.latents_std.view(shape) + self.latents_mean.view( + shape + ) + + ctx = ( + torch.autocast(device_type="cuda", dtype=torch.float16) + if autocast_fp16 and z.is_cuda + else nullcontext() + ) + with ctx: + if z.shape[2] == 1: + dec = self._decode_clip(z)[:, :, -1:] + else: + dec = self._decode_video(z) + + # out-of-place so gradients can flow through decode (pixel-space losses) + dec = dec.float() * self.pixel_std + self.pixel_mean + return dec.clamp(0.0, 1.0) * 2.0 - 1.0 diff --git a/jobs/process/BaseSDTrainProcess.py b/jobs/process/BaseSDTrainProcess.py index ff0a50d2..b6777eec 100644 --- a/jobs/process/BaseSDTrainProcess.py +++ b/jobs/process/BaseSDTrainProcess.py @@ -1988,7 +1988,12 @@ class BaseSDTrainProcess(BaseTrainProcess): # we cannot merge in if quantized or offloading. note: torchao quantized weights can # still be force merged at save time for the merge-and-reset method (see save logic), # but we keep can_merge_in False here so sampling never merges in/out. - if self.model_config.quantize or self.model_config.layer_offloading: + # models loaded from pre-quantized checkpoints (e.g. comfy convrot/nvfp4 + # imports) never set model_config.quantize, so detect their layers too + model_is_prequantized = any( + getattr(m, 'is_ostris_quantized', False) for m in unet.modules() + ) if unet is not None else False + if self.model_config.quantize or self.model_config.layer_offloading or model_is_prequantized: # todo find a way around this self.network.can_merge_in = False diff --git a/toolkit/data_transfer_object/data_loader.py b/toolkit/data_transfer_object/data_loader.py index 2aef0e57..e93d0747 100644 --- a/toolkit/data_transfer_object/data_loader.py +++ b/toolkit/data_transfer_object/data_loader.py @@ -60,6 +60,14 @@ class FileItemDTO( self.sample_rate = kwargs.get("sample_rate", 48000) self.num_frames = self.dataset_config.num_frames self.temporal_compression = kwargs.get("temporal_compression", 8) + # module-level function (picklable) for models whose valid frame + # counts are not temporal_compression * n + 1; None = default math + _sd = kwargs.get("sd", None) + self.frame_count_snapper = ( + _sd.get_frame_count_snapper() + if _sd is not None and hasattr(_sd, "get_frame_count_snapper") + else None + ) size_database = kwargs.get("size_database", {}) dataset_root = kwargs.get("dataset_root", None) self.encode_control_in_text_embeddings = kwargs.get( diff --git a/toolkit/dataloader_mixins.py b/toolkit/dataloader_mixins.py index 6adee1bf..42b14400 100644 --- a/toolkit/dataloader_mixins.py +++ b/toolkit/dataloader_mixins.py @@ -508,15 +508,19 @@ class ImageProcessingDTOMixin: if self.dataset_config.auto_frame_count: # allow for any length video here but make sure it is temporally compressable. vid_length_seconds = total_frames / video_fps - + desired_num_frames = int(vid_length_seconds * self.dataset_config.fps) - - # make sure it is divisible by temporal_compression - desired_num_frames = desired_num_frames // self.temporal_compression * self.temporal_compression - - # TODO, all models currently add a key frame, but future models may not, update here if this changes. - desired_num_frames += 1 # add one for the key frame that is always added - + + if getattr(self, 'frame_count_snapper', None) is not None: + # model-specific valid-frame-count grid (e.g. minimax_h3's 17n+5) + desired_num_frames = self.frame_count_snapper(desired_num_frames) + else: + # make sure it is divisible by temporal_compression + desired_num_frames = desired_num_frames // self.temporal_compression * self.temporal_compression + + # TODO, all models currently add a key frame, but future models may not, update here if this changes. + desired_num_frames += 1 # add one for the key frame that is always added + self.num_frames = desired_num_frames @@ -705,31 +709,39 @@ class ImageProcessingDTOMixin: else: target_duration = source_duration - waveform, sample_rate = torchaudio.load(self.path) # [channels, samples] - - waveform = waveform_to_stereo(waveform) # Convert to stereo if not already - - if self.dataset_config.audio_normalize: - peak = waveform.abs().amax() # global peak across channels - eps = 1e-9 - target_peak = 0.999 # ~ -0.01 dBFS - gain = target_peak / (peak + eps) - waveform = waveform * gain + # torchcodec's AudioDecoder raises when a video has no audio + # track, so probe for a stream before decoding. + import av + with av.open(self.path) as container: + has_audio_stream = len(container.streams.audio) > 0 - # Slice to the selected clip region (when we have a meaningful time range) - if source_duration > 0.0: - start_sample = int(round(clip_start_time * sample_rate)) - end_sample = int(round(clip_end_time * sample_rate)) - start_sample = max(0, min(start_sample, waveform.shape[-1])) - end_sample = max(0, min(end_sample, waveform.shape[-1])) - if end_sample > start_sample: - waveform = waveform[..., start_sample:end_sample] + waveform = None + if has_audio_stream: + waveform, sample_rate = torchaudio.load(self.path) # [channels, samples] + + waveform = waveform_to_stereo(waveform) # Convert to stereo if not already + + if self.dataset_config.audio_normalize: + peak = waveform.abs().amax() # global peak across channels + eps = 1e-9 + target_peak = 0.999 # ~ -0.01 dBFS + gain = target_peak / (peak + eps) + waveform = waveform * gain + + # Slice to the selected clip region (when we have a meaningful time range) + if source_duration > 0.0: + start_sample = int(round(clip_start_time * sample_rate)) + end_sample = int(round(clip_end_time * sample_rate)) + start_sample = max(0, min(start_sample, waveform.shape[-1])) + end_sample = max(0, min(end_sample, waveform.shape[-1])) + if end_sample > start_sample: + waveform = waveform[..., start_sample:end_sample] + else: + # No valid audio segment + waveform = None else: - # No valid audio segment + # If we can't compute a meaningful time range, treat as no-audio waveform = None - else: - # If we can't compute a meaningful time range, treat as no-audio - waveform = None if waveform is not None and waveform.numel() > 0: target_samples = int(round(target_duration * sample_rate)) diff --git a/toolkit/models/base_model.py b/toolkit/models/base_model.py index 419060fc..eb75ddb2 100644 --- a/toolkit/models/base_model.py +++ b/toolkit/models/base_model.py @@ -274,12 +274,24 @@ class BaseModel: except: # if we have a custom vae, it might not have this divisibility = 8 - + # flux packs this again, if self.is_flux: divisibility = divisibility * 2 return divisibility + def get_frame_count_snapper(self): + """Optional hook for video models whose VAE accepts frame counts on a + grid other than the default ``temporal_compression * n + 1``. + + Return a MODULE-LEVEL function ``(num_frames: int) -> int`` that snaps + an arbitrary frame count DOWN to the nearest count the video VAE can + encode (it must be picklable by reference — file items travel into + dataloader workers, so no lambdas or bound methods). Returning None + keeps the default auto_frame_count behavior. + """ + return None + # these must be implemented in child classes def load_model(self): # override this in child classes diff --git a/toolkit/paths.py b/toolkit/paths.py index 9f6a8cfb..28b31e32 100644 --- a/toolkit/paths.py +++ b/toolkit/paths.py @@ -8,7 +8,7 @@ DIFFUSERS_CONFIGS_ROOT = os.path.join(TOOLKIT_ROOT, "toolkit", "diffusers_config COMFY_MODELS_PATH = None # check if ENV variable is set -if 'MODELS_PATH' in os.environ: +if 'MODELS_PATH' in os.environ and os.environ['MODELS_PATH'].strip() != "": MODELS_PATH = os.environ['MODELS_PATH'] else: MODELS_PATH = os.path.join(TOOLKIT_ROOT, "models") diff --git a/toolkit/util/comfy_quant_import.py b/toolkit/util/comfy_quant_import.py new file mode 100644 index 00000000..85f31994 --- /dev/null +++ b/toolkit/util/comfy_quant_import.py @@ -0,0 +1,175 @@ +"""Import ComfyUI pre-quantized checkpoints onto toolkit modules. + +ComfyUI quantized checkpoints mark each quantized submodule with a +``.comfy_quant`` uint8 tensor holding a JSON config, alongside the +quantized ``weight`` and its scale tensors. This module walks those markers +and converts the matching submodules in place: + + - ``{"format": "int8_tensorwise", "convrot": true, "convrot_groupsize": G}`` + per-output-row symmetric int8 on regular-Hadamard-rotated weights — the + exact storage of the toolkit's convrot8 backend + (toolkit/util/convrot_quant.py:ConvRotInt8Quantizer), so the tensors are + attached to its buffers directly (no requantization). Without the + ``convrot`` flag the rotation block is 1, i.e. plain per-row int8, which + the same backend also decodes (rotate is the identity at rot_size 1). + - ``{"format": "nvfp4"}`` block-16 fp4 with e4m3 block scales, an fp32 + per-tensor scale and an optional AWQ ``pre_quant_scale`` — attached to + the nvfp4 backend (toolkit/util/nvfp4_quant.py). + - an int8 marker on an ``nn.Embedding`` swaps in :class:`Int8Embedding` + (per-row scales, dequantized per lookup). + +Linears become OstrisLinear (class swap in place, like +convert_linear_to_ostris), so LoRA attachment, memory management and the +quantized save paths all work unchanged. +""" + +import json +from typing import Dict, Tuple + +import torch + +from toolkit.util.nvfp4_quant import ( + Nvfp4Quantizer, + swap_nvfp4_nibbles, + unswizzle_nvfp4_scales, +) +from toolkit.util.ostris_quant import OstrisLinear, get_ostris_quantizer + + +def parse_comfy_quant_blob(blob: torch.Tensor) -> dict: + return json.loads(bytes(blob.cpu().tolist()).decode("utf-8")) + + +class Int8Embedding(torch.nn.Module): + """An embedding table stored as per-row symmetric int8. Rows are + dequantized per lookup, so the full-precision table never materializes.""" + + def __init__(self, qweight: torch.Tensor, scales: torch.Tensor, dtype: torch.dtype): + super().__init__() + self.num_embeddings, self.embedding_dim = qweight.shape + self.output_dtype = dtype + self.register_buffer("qweight", qweight.contiguous(), persistent=False) + self.register_buffer( + "scales", + scales.detach().float().reshape(-1).contiguous().view(torch.uint8), + persistent=False, + ) + + @property + def weight(self): + # full dequantized table, for code that inspects it + scales = self.scales.view(torch.float32) + return (self.qweight.float() * scales.unsqueeze(1)).to(self.output_dtype) + + def forward(self, input_ids: torch.Tensor) -> torch.Tensor: + flat = input_ids.reshape(-1) + rows = self.qweight.index_select(0, flat).float() + scales = self.scales.view(torch.float32).index_select(0, flat) + out = rows * scales.unsqueeze(1) + return out.to(self.output_dtype).reshape(*input_ids.shape, self.embedding_dim) + + +def _to_ostris(module: torch.nn.Linear, quantizer, orig_dtype: torch.dtype) -> OstrisLinear: + if "weight" in module._parameters: + del module._parameters["weight"] + module.ostris_quantizer = quantizer + module.ostris_orig_dtype = orig_dtype + if module.bias is not None: + module.bias.requires_grad_(False) + module.__class__ = OstrisLinear + return module + + +@torch.no_grad() +def import_comfy_quantized_layers( + root: torch.nn.Module, + state_dict: Dict[str, torch.Tensor], + orig_dtype: torch.dtype = torch.bfloat16, + key_map=None, +) -> Tuple[Dict[str, torch.Tensor], int]: + """Convert every module a ``comfy_quant`` marker points at and attach its + quantized tensors. Consumes the quantized entries from ``state_dict`` and + returns ``(remaining_state_dict, num_converted)`` — load the remainder + with the regular load_state_dict. + + ``key_map`` optionally maps a checkpoint prefix to the module path in + ``root`` (e.g. comfy text encoder keys onto transformers module paths). + """ + state_dict = dict(state_dict) + converted = 0 + + marker_keys = [k for k in state_dict.keys() if k.endswith(".comfy_quant")] + for marker_key in marker_keys: + prefix = marker_key[: -len(".comfy_quant")] + conf = parse_comfy_quant_blob(state_dict.pop(marker_key)) + fmt = conf.get("format") + module_path = key_map(prefix) if key_map is not None else prefix + module = root.get_submodule(module_path) + + weight = state_dict.pop(f"{prefix}.weight") + weight_scale = state_dict.pop(f"{prefix}.weight_scale", None) + + if isinstance(module, torch.nn.Embedding): + if fmt != "int8_tensorwise": + raise ValueError( + f"Unsupported comfy quant format {fmt!r} on embedding {prefix}" + ) + parent_path, _, attr = module_path.rpartition(".") + parent = root.get_submodule(parent_path) if parent_path else root + setattr(parent, attr, Int8Embedding(weight, weight_scale, orig_dtype)) + converted += 1 + continue + + if not isinstance(module, torch.nn.Linear): + raise ValueError( + f"comfy_quant marker {prefix} points at {type(module).__name__}, " + "expected nn.Linear or nn.Embedding" + ) + + if fmt == "int8_tensorwise": + rot = int(conf.get("convrot_groupsize", 256)) if conf.get("convrot") else 1 + quantizer = get_ostris_quantizer("convrot8") + module.register_buffer("cr8_qdata", weight.contiguous(), persistent=False) + module.register_buffer( + "cr8_scales", + weight_scale.detach().float().reshape(-1).contiguous().view(torch.uint8), + persistent=False, + ) + module.cr8_rot_size = rot + elif fmt == "nvfp4": + quantizer = get_ostris_quantizer("nvfp4") + # normalize comfy_kitchen's storage to the toolkit's conventions: + # fp4 pairs are packed high-nibble-first and the e4m3 block scales + # are stored in the swizzled cuBLAS 128x4 tile layout + scales = unswizzle_nvfp4_scales( + weight_scale.view(torch.float8_e4m3fn), + module.out_features, + module.in_features // 16, + ) + Nvfp4Quantizer.attach_( + module, + packed=swap_nvfp4_nibbles(weight), + scales=scales, + pts=state_dict.pop(f"{prefix}.weight_scale_2"), + pre_scale=state_dict.pop(f"{prefix}.pre_quant_scale", None), + ) + else: + raise ValueError( + f"Unsupported comfy quant format {fmt!r} on {prefix} " + "(supported: int8_tensorwise, nvfp4)" + ) + + # drop unused calibration extras if present + state_dict.pop(f"{prefix}.input_scale", None) + + _to_ostris(module, quantizer, orig_dtype) + bias = state_dict.pop(f"{prefix}.bias", None) + if bias is not None and module.bias is not None: + # bias may still be a meta parameter when the model was built under + # a meta device context + module._parameters["bias"] = torch.nn.Parameter( + bias.detach().clone(), requires_grad=False + ) + converted += 1 + + return state_dict, converted diff --git a/toolkit/util/nvfp4_quant.py b/toolkit/util/nvfp4_quant.py new file mode 100644 index 00000000..72d71556 --- /dev/null +++ b/toolkit/util/nvfp4_quant.py @@ -0,0 +1,155 @@ +"""NVFP4 (ModelOpt/ComfyUI-style) OstrisQuantizer backend — qtype "nvfp4". + +Plain block-16 nvfp4 weight storage without ConvRot's rotation: fp4 e2m1 +codes packed two per byte, one fp8 e4m3 scale per 16 elements, one fp32 +per-tensor scale, plus an optional AWQ ``pre_quant_scale`` applied +elementwise to the input activation before the matmul (ModelOpt convention, +matching ComfyUI's quantized ops). + +This is the layout ComfyUI checkpoints tagged ``{"format": "nvfp4"}`` carry +(e.g. the Comfy-Org MiniMax-H3 text encoder). Those exports set +``full_precision_matrix_mult`` — the activations are NOT fp4-quantized — so +the forward here is always the dequantized matmul in the activation's dtype: +weights stay at ~4.25 bits in memory and the math runs on any GPU (or CPU), +no Blackwell fp4 tensor cores required. The triton dequant kernel in +convrot_quant is used when available; pure torch otherwise. + +Quantized state attached to each module (uint8 byte views, like the convrot +backends, so nn.Module._apply dtype casts can't corrupt them): + nv4_qdata packed e2m1 codes (uint8, out x in/2; low nibble = even column) + nv4_scales e4m3 block scales (out x in/16) + nv4_pts fp32 per-tensor scale (1 element) + nv4_pre_scale optional fp32 AWQ input scale (in,) +""" + +from typing import Optional + +import torch +import torch.nn.functional as F + +from toolkit.util.convrot_quant import BLOCK, dequantize_nvfp4, quantize_nvfp4 +from toolkit.util.ostris_quant import OstrisQuantizer +from toolkit.print import print_acc + +NVFP4_QTYPES = ("nvfp4",) + +_skip_warned = set() + + +def unswizzle_nvfp4_scales(scales: torch.Tensor, rows: int, cols: int) -> torch.Tensor: + """Undo the cuBLAS 128x4-tile block-scale layout (comfy_kitchen's + ``to_blocked``) back to a row-major (rows, cols) matrix. ComfyUI nvfp4 + checkpoints store ``weight_scale`` swizzled; when the dims are already + tile-aligned the shape is unchanged and only the element order differs.""" + n_row_blocks = (rows + 127) // 128 + n_col_blocks = (cols + 3) // 4 + padded_rows = n_row_blocks * 128 + padded_cols = n_col_blocks * 4 + x = scales.reshape(-1, 32, 16) + x = x.reshape(-1, 32, 4, 4).transpose(1, 2) + x = x.reshape(n_row_blocks, n_col_blocks, 4, 32, 4) + x = x.reshape(n_row_blocks, n_col_blocks, 128, 4) + x = x.permute(0, 2, 1, 3).reshape(padded_rows, padded_cols) + return x[:rows, :cols].contiguous() + + +def swap_nvfp4_nibbles(packed: torch.Tensor) -> torch.Tensor: + """ComfyUI packs fp4 pairs high-nibble-first; the toolkit's decode is + low-nibble-first. Swapping nibbles converts between the two.""" + return ((packed << 4) | (packed >> 4)).contiguous() + + +class Nvfp4Quantizer(OstrisQuantizer): + """Block-16 nvfp4 weights, full-precision activations. One instance is + shareable across modules.""" + + def can_quantize(self, module: torch.nn.Linear) -> bool: + if module.in_features % BLOCK != 0: + if module.in_features not in _skip_warned: + _skip_warned.add(module.in_features) + print_acc( + f"nvfp4: skipping linears with in_features={module.in_features} " + f"(needs in divisible by {BLOCK})" + ) + return False + return True + + def quantize_(self, module: torch.nn.Linear, weight_fp32: torch.Tensor) -> None: + packed, scales, pts = quantize_nvfp4(weight_fp32, optimize_scales=True) + self.attach_(module, packed, scales, pts, pre_scale=None) + + @staticmethod + def attach_( + module: torch.nn.Module, + packed: torch.Tensor, # uint8 (out, in/2) + scales: torch.Tensor, # float8_e4m3fn (out, in/16) + pts: torch.Tensor, # fp32 scalar per-tensor scale + pre_scale: Optional[torch.Tensor] = None, # (in,) AWQ input scale + ) -> None: + """Register the quantized representation on the module. Used both by + quantize_ and by importers of pre-quantized checkpoints.""" + module.register_buffer("nv4_qdata", packed.contiguous(), persistent=False) + module.register_buffer( + "nv4_scales", scales.contiguous().view(torch.uint8), persistent=False + ) + module.register_buffer( + "nv4_pts", + pts.detach().float().clone().reshape(1).view(torch.uint8), + persistent=False, + ) + if pre_scale is not None: + module.register_buffer( + "nv4_pre_scale", + pre_scale.detach().float().clone().contiguous().view(torch.uint8), + persistent=False, + ) + + @staticmethod + def _pts(module) -> torch.Tensor: + return module.nv4_pts.view(torch.float32).reshape(()) + + @staticmethod + def _pre_scale(module) -> Optional[torch.Tensor]: + buf = getattr(module, "nv4_pre_scale", None) + return None if buf is None else buf.view(torch.float32) + + def _dequantize_weight(self, module, dtype: torch.dtype) -> torch.Tensor: + return dequantize_nvfp4( + module.nv4_qdata, + module.nv4_scales.view(torch.float8_e4m3fn), + self._pts(module), + module.out_features, + module.in_features, + dtype, + ) + + def dequantize(self, module) -> torch.Tensor: + """The weight as stored. NOTE: with an AWQ pre_quant_scale present the + stored weight expects pre-scaled activations; folding the scale back + (w * pre_scale per column) would reconstruct the original-basis weight + but is deliberately not done here — forward() owns that contract.""" + return self._dequantize_weight(module, torch.float32) + + def dequantize_folded(self, module) -> torch.Tensor: + """Weight for raw (un-pre-scaled) activations: the AWQ pre_quant_scale + multiplies the input elementwise, which folds into the weight columns.""" + w = self._dequantize_weight(module, torch.float32) + pre_scale = self._pre_scale(module) + if pre_scale is not None: + w = w * pre_scale.unsqueeze(0) + return w + + def requantize_(self, module, fp_weight: torch.Tensor) -> None: + w = fp_weight.to(device=module.nv4_qdata.device, dtype=torch.float32) + packed, scales, pts = quantize_nvfp4(w, optimize_scales=True) + module.nv4_qdata = packed + module.nv4_scales = scales.view(torch.uint8) + module.nv4_pts = pts.detach().clone().reshape(1).view(torch.uint8) + + def forward(self, module, x: torch.Tensor) -> torch.Tensor: + pre_scale = self._pre_scale(module) + if pre_scale is not None: + x = x * pre_scale.to(dtype=x.dtype) + with torch.no_grad(): + w = self._dequantize_weight(module, x.dtype) + return F.linear(x, w, module.bias) diff --git a/toolkit/util/ostris_quant.py b/toolkit/util/ostris_quant.py index a647295a..6384674d 100644 --- a/toolkit/util/ostris_quant.py +++ b/toolkit/util/ostris_quant.py @@ -51,6 +51,13 @@ class OstrisQuantizer: """Reconstruct the full weight in the original basis, in float32.""" raise NotImplementedError + def dequantize_folded(self, module: "OstrisLinear") -> torch.Tensor: + """The full weight with any ACTIVATION-side transform folded in, i.e. a + weight that computes the same output on raw activations. Used when + re-quantizing into a different backend (which won't know about this + backend's activation transforms). Default: same as dequantize.""" + return self.dequantize(module) + def requantize_(self, module: "OstrisLinear", fp_weight: torch.Tensor) -> None: """Re-quantize in place from a full precision weight in the original basis (used by the continuous merge/reset method).""" @@ -179,6 +186,7 @@ def get_ostris_quantizer(qtype: str) -> Optional[OstrisQuantizer]: """Resolve a qtype string to a quantizer backend instance, or None if the qtype does not belong to a custom backend. Add new backends here.""" from toolkit.util.convrot_quant import CONVROT_QTYPES, get_convrot_quantizer + from toolkit.util.nvfp4_quant import NVFP4_QTYPES, Nvfp4Quantizer from toolkit.util.orbit_quant import ORBIT_QTYPES, OrbitQuantizer from toolkit.util.orbit_vq_quant import ORBIT_VQ_QTYPES, OrbitVQQuantizer from toolkit.util.uintx_quant import UINTX_QTYPES, UIntXQuantizer @@ -190,6 +198,8 @@ def get_ostris_quantizer(qtype: str) -> Optional[OstrisQuantizer]: quantizer = OrbitVQQuantizer(**ORBIT_VQ_QTYPES[qtype]) elif qtype in CONVROT_QTYPES: quantizer = get_convrot_quantizer(qtype) + elif qtype in NVFP4_QTYPES: + quantizer = Nvfp4Quantizer() elif qtype in UINTX_QTYPES: quantizer = UIntXQuantizer(UINTX_QTYPES[qtype]) if quantizer is not None: @@ -324,8 +334,33 @@ def convert_linear_to_ostris( module: torch.nn.Linear, quantizer: OstrisQuantizer ) -> bool: """Quantize an nn.Linear in place (class swap). Returns True if the module was - converted (or already was), False if it is not a candidate.""" + converted (or already was), False if it is not a candidate. + + A module that is ALREADY quantized (e.g. loaded from a pre-quantized + checkpoint) is re-quantized into the requested backend when the qtypes + differ: the weight is dequantized and re-quantized layer by layer, so the + full-precision transient never exceeds one layer's weight. Same qtype is + a no-op (the shipped quantization is kept).""" if isinstance(module, OstrisLinear): + current_qtype = getattr(module.ostris_quantizer, "qtype", None) + if quantizer.qtype is None or current_qtype == quantizer.qtype: + return True + if not quantizer.can_quantize(module): + return True # keep the existing quantization rather than dropping it + # fold any activation-side transform (e.g. an AWQ pre_quant_scale) into + # the weight so the new backend computes the same function on raw inputs + weight = module.ostris_quantizer.dequantize_folded(module).to( + module.ostris_orig_dtype + ) + # backend state lives exclusively in buffers; leftover scalar attrs from + # the old backend are inert + module._buffers.clear() + if quantizer.wants_fp32_weight: + quantizer.quantize_(module, weight.to(torch.float32)) + else: + quantizer.quantize_(module, weight) + del weight + module.ostris_quantizer = quantizer return True weight = getattr(module, "weight", None) if not isinstance(weight, torch.nn.Parameter) or not weight.dtype.is_floating_point: diff --git a/toolkit/util/quantize.py b/toolkit/util/quantize.py index 9307f3d4..cf983c36 100644 --- a/toolkit/util/quantize.py +++ b/toolkit/util/quantize.py @@ -176,7 +176,16 @@ def quantize( try: # check if m is QLinear or QConv2d if m.__class__.__name__ in Q_MODULES: - continue + # OstrisLinear may still be RE-quantized into a different + # ostris qtype (same qtype is a per-layer no-op); every other + # already-quantized module type is always left alone, which + # also keeps quanto/torchao from double-quantizing + # pre-quantized checkpoints + if not ( + isinstance(weights, ostristype) + and m.__class__.__name__ == "OstrisLinear" + ): + continue if ( isinstance(weights, aotype) and not isinstance(m, torch.nn.Linear) @@ -193,7 +202,10 @@ def quantize( continue orig_device = None if quantize_device is not None and next(m.children(), None) is None: + # OstrisLinear layers being re-quantized hold buffers, not params param = next(m.parameters(recurse=False), None) + if param is None: + param = next(m.buffers(recurse=False), None) if param is not None: orig_device = param.device m.to(quantize_device) diff --git a/ui/cron/actions/startJob.ts b/ui/cron/actions/startJob.ts index 0312b038..f0f3534f 100644 --- a/ui/cron/actions/startJob.ts +++ b/ui/cron/actions/startJob.ts @@ -3,7 +3,7 @@ import { Job } from '@prisma/client'; import { spawn, ChildProcess } from 'child_process'; import path from 'path'; import fs from 'fs'; -import { TOOLKIT_ROOT, getTrainingFolder, getHFToken } from '../paths'; +import { TOOLKIT_ROOT, getTrainingFolder, getHFToken, getModelsPath } from '../paths'; import { resolveDetachedPythonPath } from '../pythonPath'; const isWindows = process.platform === 'win32'; @@ -250,6 +250,15 @@ const startAndWatchJob = (job: Job) => { additionalEnv.HF_TOKEN = hfToken; } + // MODELS_PATH - one set in the env always takes precedence (it passes + // through via process.env); only fall back to the setting if it is not set + if (!process.env.MODELS_PATH || process.env.MODELS_PATH.trim() === '') { + const modelsPath = await getModelsPath(); + if (modelsPath && modelsPath.trim() !== '') { + additionalEnv.MODELS_PATH = modelsPath; + } + } + const args = [runFilePath, configPath]; // Where the Windows relay reports the job's real pid back to us. diff --git a/ui/cron/paths.ts b/ui/cron/paths.ts index a2a4eab6..54d1084d 100644 --- a/ui/cron/paths.ts +++ b/ui/cron/paths.ts @@ -5,6 +5,7 @@ export const TOOLKIT_ROOT = path.resolve('@', '..', '..'); export const defaultTrainFolder = path.join(TOOLKIT_ROOT, 'output'); export const defaultDatasetsFolder = path.join(TOOLKIT_ROOT, 'datasets'); export const defaultDataRoot = path.join(TOOLKIT_ROOT, 'data'); +export const defaultModelsFolder = path.join(TOOLKIT_ROOT, 'models'); // Forked file-server workers set AI_TOOLKIT_QUIET_PATHS so this line prints // once per launched process group, not once per worker. @@ -39,3 +40,17 @@ export const getHFToken = async () => { } return token; }; + +export const getModelsPath = async () => { + const key = 'MODELS_PATH'; + let row = await prisma.settings.findFirst({ + where: { + key: key, + }, + }); + let modelsPath = ''; + if (row?.value && row.value !== '' && row.value !== defaultModelsFolder) { + modelsPath = row.value; + } + return modelsPath; +}; diff --git a/ui/src/app/api/settings/route.ts b/ui/src/app/api/settings/route.ts index 24722a2a..46cf47ee 100644 --- a/ui/src/app/api/settings/route.ts +++ b/ui/src/app/api/settings/route.ts @@ -1,6 +1,6 @@ import { NextResponse } from 'next/server'; import prisma from '@/server/prisma'; -import { defaultTrainFolder, defaultDatasetsFolder } from '@/paths'; +import { defaultTrainFolder, defaultDatasetsFolder, defaultModelsFolder } from '@/paths'; import { flushCache } from '@/server/settings'; export async function GET() { @@ -18,6 +18,13 @@ export async function GET() { if (!settingsObject.DATASETS_FOLDER || settingsObject.DATASETS_FOLDER === '') { settingsObject.DATASETS_FOLDER = defaultDatasetsFolder; } + // MODELS_PATH from the env file always takes precedence over the setting + if (process.env.MODELS_PATH && process.env.MODELS_PATH.trim() !== '') { + settingsObject.MODELS_PATH = process.env.MODELS_PATH; + } else if (!settingsObject.MODELS_PATH || settingsObject.MODELS_PATH === '') { + // if MODELS_PATH is not set, use default + settingsObject.MODELS_PATH = defaultModelsFolder; + } return NextResponse.json(settingsObject); } catch (error) { return NextResponse.json({ error: 'Failed to fetch settings' }, { status: 500 }); @@ -27,7 +34,7 @@ export async function GET() { export async function POST(request: Request) { try { const body = await request.json(); - const { HF_TOKEN, TRAINING_FOLDER, DATASETS_FOLDER } = body; + const { HF_TOKEN, TRAINING_FOLDER, DATASETS_FOLDER, MODELS_PATH } = body; // Upsert both settings await Promise.all([ @@ -46,6 +53,11 @@ export async function POST(request: Request) { update: { value: DATASETS_FOLDER }, create: { key: 'DATASETS_FOLDER', value: DATASETS_FOLDER }, }), + prisma.settings.upsert({ + where: { key: 'MODELS_PATH' }, + update: { value: MODELS_PATH }, + create: { key: 'MODELS_PATH', value: MODELS_PATH }, + }), ]); flushCache(); diff --git a/ui/src/app/jobs/new/SimpleJob.tsx b/ui/src/app/jobs/new/SimpleJob.tsx index 36ec7fd1..efb8491c 100644 --- a/ui/src/app/jobs/new/SimpleJob.tsx +++ b/ui/src/app/jobs/new/SimpleJob.tsx @@ -324,6 +324,30 @@ export default function SimpleJob({ placeholder="" /> )} + {modelArch?.modelNotes && ( +
+ +
+ )} {modelArch?.gateUrl && (
+ +
+ + +
diff --git a/ui/src/hooks/useSettings.tsx b/ui/src/hooks/useSettings.tsx index 35fcc538..e5cf58a9 100644 --- a/ui/src/hooks/useSettings.tsx +++ b/ui/src/hooks/useSettings.tsx @@ -7,6 +7,7 @@ export interface Settings { HF_TOKEN: string; TRAINING_FOLDER: string; DATASETS_FOLDER: string; + MODELS_PATH: string; } export default function useSettings() { @@ -14,6 +15,7 @@ export default function useSettings() { HF_TOKEN: '', TRAINING_FOLDER: '', DATASETS_FOLDER: '', + MODELS_PATH: '', }); const [isSettingsLoaded, setIsLoaded] = useState(false); useEffect(() => { @@ -26,6 +28,7 @@ export default function useSettings() { HF_TOKEN: data.HF_TOKEN || '', TRAINING_FOLDER: data.TRAINING_FOLDER || '', DATASETS_FOLDER: data.DATASETS_FOLDER || '', + MODELS_PATH: data.MODELS_PATH || '', }); setIsLoaded(true); }) diff --git a/ui/src/paths.ts b/ui/src/paths.ts index 92311f43..54fb1c65 100644 --- a/ui/src/paths.ts +++ b/ui/src/paths.ts @@ -3,3 +3,4 @@ export const TOOLKIT_ROOT = path.resolve('@', '..', '..'); export const defaultTrainFolder = path.join(TOOLKIT_ROOT, 'output'); export const defaultDatasetsFolder = path.join(TOOLKIT_ROOT, 'datasets'); export const defaultDataRoot = path.join(TOOLKIT_ROOT, 'data'); +export const defaultModelsFolder = path.join(TOOLKIT_ROOT, 'models'); diff --git a/version.py b/version.py index 15ab0ed6..f792f6fe 100644 --- a/version.py +++ b/version.py @@ -1 +1 @@ -VERSION = "0.12.1" +VERSION = "0.12.2"