Image-Text-to-Video
Diffusers
Safetensors
text-to-video
image-to-video
video-to-video
text-to-audio-video
image-to-audio-video
image-text-to-audio-video
video-to-audio-video
audio-to-audio-video
audio-video-generation
multimodal
synchronized-audio-video
reference-to-audio-video
Instructions to use MiniMaxAI/MiniMax-H3 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use MiniMaxAI/MiniMax-H3 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("MiniMaxAI/MiniMax-H3", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
| # SPDX-License-Identifier: Apache-2.0 | |
| # Torch-native normalization for the MiniMax H3 visual VAE. | |
| import math | |
| import os | |
| import torch | |
| import torch.distributed as dist | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from .conv import SpatialParallelConv3d | |
| from .parallel import all_reduce, get_parallel_state | |
| def _validate_activation(activation): | |
| valid_activations = {"identity", "silu", "relu"} | |
| if activation not in valid_activations: | |
| raise ValueError( | |
| f"Unsupported activation: {activation}. Supported: {valid_activations}" | |
| ) | |
| def _apply_activation(x, activation): | |
| _validate_activation(activation) | |
| if activation == "identity": | |
| return x | |
| if activation == "silu": | |
| return F.silu(x) | |
| return F.relu(x) | |
| def _merge_time_to_batch(x): | |
| batch, channels, depth, height, width = x.shape | |
| return ( | |
| x.permute(0, 2, 1, 3, 4) | |
| .contiguous() | |
| .view(batch * depth, channels, 1, height, width) | |
| ) | |
| def _split_time_from_batch(x, batch): | |
| batch_depth, channels, _, height, width = x.shape | |
| depth = batch_depth // batch | |
| return ( | |
| x.view(batch, depth, channels, height, width) | |
| .permute(0, 2, 1, 3, 4) | |
| .contiguous() | |
| ) | |
| def fused_group_norm(x, num_groups, weight, bias, eps=1e-5, activation="silu"): | |
| out = F.group_norm(x, num_groups, weight=weight, bias=bias, eps=eps) | |
| return _apply_activation(out, activation) | |
| def fused_spatial_norm( | |
| f, | |
| num_groups, | |
| norm_weight, | |
| norm_bias, | |
| dynamic_scale, | |
| dynamic_bias, | |
| eps=1e-5, | |
| activation="silu", | |
| ): | |
| norm_f = F.group_norm( | |
| f, | |
| num_groups, | |
| weight=norm_weight, | |
| bias=norm_bias, | |
| eps=eps, | |
| ) | |
| out = norm_f * dynamic_scale + dynamic_bias | |
| return _apply_activation(out, activation) | |
| class DummyAffine(torch.nn.Module): | |
| def __init__(self, num_channels, affine=True): | |
| super().__init__() | |
| if affine: | |
| self.weight = torch.nn.Parameter(torch.ones(num_channels)) | |
| self.bias = torch.nn.Parameter(torch.zeros(num_channels)) | |
| else: | |
| self.register_parameter("weight", None) | |
| self.register_parameter("bias", None) | |
| def forward(self, input): | |
| if self.weight is None: | |
| return input | |
| shape = [1, -1] + [1] * (input.dim() - 2) | |
| return input * self.weight.view(*shape) + self.bias.view(*shape) | |
| class FusedGroupNorm3D(torch.nn.Module): | |
| """Compatibility wrapper implemented with native PyTorch ops.""" | |
| def __init__( | |
| self, | |
| num_groups, | |
| num_channels, | |
| eps=1e-5, | |
| affine=True, | |
| activation="silu", | |
| cond_channels=None, | |
| use_t_isolated_gn=False, | |
| padding_mode="zeros", | |
| padding_mode_t=None, | |
| causal=True, | |
| ): | |
| super().__init__() | |
| _validate_activation(activation) | |
| self.num_groups = num_groups | |
| self.num_channels = num_channels | |
| self.eps = eps | |
| self.affine = affine | |
| self.activation = activation | |
| self.use_t_isolated_gn = use_t_isolated_gn | |
| if cond_channels is not None: | |
| self.use_spatial_affine = True | |
| self.norm_layer = DummyAffine(num_channels, affine=affine) | |
| self.conv_y = SpatialParallelConv3d( | |
| cond_channels, | |
| num_channels, | |
| kernel_size=1, | |
| padding_mode=padding_mode, | |
| padding_mode_t=padding_mode_t, | |
| causal=causal, | |
| ) | |
| self.conv_b = SpatialParallelConv3d( | |
| cond_channels, | |
| num_channels, | |
| kernel_size=1, | |
| padding_mode=padding_mode, | |
| padding_mode_t=padding_mode_t, | |
| causal=causal, | |
| ) | |
| else: | |
| self.use_spatial_affine = False | |
| if self.affine: | |
| self.weight = torch.nn.Parameter(torch.ones(num_channels)) | |
| self.bias = torch.nn.Parameter(torch.zeros(num_channels)) | |
| else: | |
| self.register_parameter("weight", None) | |
| self.register_parameter("bias", None) | |
| def forward(self, f, cond=None): | |
| need_reshape = self.use_t_isolated_gn and f.dim() == 5 | |
| batch = f.shape[0] if need_reshape else None | |
| f_size = f.shape[-3:] | |
| if need_reshape: | |
| f = _merge_time_to_batch(f) | |
| if self.use_spatial_affine: | |
| scale = self.conv_y(cond) | |
| bias = self.conv_b(cond) | |
| if math.prod(scale.shape[-3:]) * math.prod(bias.shape[-3:]) > 1: | |
| scale = F.interpolate(scale, size=f_size, mode="nearest") | |
| bias = F.interpolate(bias, size=f_size, mode="nearest") | |
| if need_reshape: | |
| scale = _merge_time_to_batch(scale) | |
| bias = _merge_time_to_batch(bias) | |
| out = fused_spatial_norm( | |
| f, | |
| self.num_groups, | |
| self.norm_layer.weight, | |
| self.norm_layer.bias, | |
| scale, | |
| bias, | |
| self.eps, | |
| self.activation, | |
| ) | |
| else: | |
| if cond is not None: | |
| raise NotImplementedError("Dynamic affine is not defined") | |
| weight = self.weight if self.affine else None | |
| bias = self.bias if self.affine else None | |
| out = fused_group_norm( | |
| f, self.num_groups, weight, bias, self.eps, self.activation | |
| ) | |
| if need_reshape: | |
| out = _split_time_from_batch(out, batch) | |
| return out | |
| class SpatialParallelGroupNorm(nn.GroupNorm): | |
| def __init__( | |
| self, | |
| *args, | |
| **kwargs, | |
| ): | |
| super().__init__(*args, **kwargs) | |
| self.spatial_parallel = False | |
| def _compute_stats(self, input): | |
| batch, channels = input.shape[0], input.shape[1] | |
| spatial_dims = input.shape[2:] | |
| spatial_size = math.prod(spatial_dims) | |
| groups = self.num_groups | |
| x = input.reshape(batch, groups, channels // groups, -1).to(torch.float32) | |
| local_sum = x.sum(dim=(2, 3)) | |
| local_square_sum = (x * x).sum(dim=(2, 3)) | |
| local_n = (channels // groups) * spatial_size | |
| local_n_tensor = torch.full_like(local_sum, float(local_n)) | |
| stats = torch.stack([local_sum, local_square_sum, local_n_tensor], dim=0) | |
| local_process_group = get_parallel_state()["local_process_group"] | |
| stats = all_reduce(stats, dist.ReduceOp.SUM, local_process_group) | |
| total_sum = stats[0] | |
| total_square_sum = stats[1] | |
| total_n = stats[2] | |
| mean = total_sum / total_n | |
| var = (total_square_sum / total_n) - mean**2 | |
| return mean, var | |
| def forward(self, input): | |
| if not self.spatial_parallel: | |
| return nn.GroupNorm.forward(self, input) | |
| batch, channels = input.shape[0], input.shape[1] | |
| orig_shape = input.shape | |
| mean, var = self._compute_stats(input) | |
| x = input.reshape(batch, self.num_groups, channels // self.num_groups, -1) | |
| mean = mean.unsqueeze(-1).unsqueeze(-1) | |
| var = var.unsqueeze(-1).unsqueeze(-1) | |
| x = (x - mean) / torch.sqrt(var + self.eps) | |
| x = x.reshape(orig_shape) | |
| if self.affine: | |
| shape = [1, -1] + [1] * (len(orig_shape) - 2) | |
| x *= self.weight.view(*shape) | |
| x += self.bias.view(*shape) | |
| return x | |
| class TemporalIsolatedSpatialParallelGroupNorm(SpatialParallelGroupNorm): | |
| def forward(self, input): | |
| if input.dim() == 5: | |
| batch = input.shape[0] | |
| input = _merge_time_to_batch(input) | |
| output = super().forward(input) | |
| return _split_time_from_batch(output, batch) | |
| return super().forward(input) | |
| class SpatialNorm3D(nn.Module): | |
| def __init__( | |
| self, | |
| f_channels, | |
| zq_channels, | |
| padding_mode="zeros", | |
| padding_mode_t=None, | |
| causal=True, | |
| use_t_isolated_gn=False, | |
| ): | |
| super().__init__() | |
| norm_cls = ( | |
| TemporalIsolatedSpatialParallelGroupNorm | |
| if use_t_isolated_gn | |
| else SpatialParallelGroupNorm | |
| ) | |
| self.norm_layer = norm_cls( | |
| num_groups=32, num_channels=f_channels, eps=1e-6, affine=True | |
| ) | |
| self.conv_y = SpatialParallelConv3d( | |
| zq_channels, | |
| f_channels, | |
| kernel_size=1, | |
| padding_mode=padding_mode, | |
| padding_mode_t=padding_mode_t, | |
| causal=causal, | |
| ) | |
| self.conv_b = SpatialParallelConv3d( | |
| zq_channels, | |
| f_channels, | |
| kernel_size=1, | |
| padding_mode=padding_mode, | |
| padding_mode_t=padding_mode_t, | |
| causal=causal, | |
| ) | |
| def forward(self, f, zq): | |
| f_size = f.shape[-3:] | |
| norm_f = self.norm_layer(f) | |
| scale = self.conv_y(zq) | |
| bias = self.conv_b(zq) | |
| if math.prod(scale.shape[-3:]) * math.prod(bias.shape[-3:]) > 1: | |
| scale = F.interpolate(scale, size=f_size, mode="nearest") | |
| bias = F.interpolate(bias, size=f_size, mode="nearest") | |
| return norm_f * scale + bias | |
| def get_spatial_norm_3d( | |
| num_channels, | |
| cond_channels, | |
| *, | |
| padding_mode="zeros", | |
| padding_mode_t=None, | |
| causal=True, | |
| use_t_isolated_gn=False, | |
| ): | |
| if os.environ.get("MINIMAX_H3_USE_FUSED_NORM", "false").lower() == "true": | |
| return FusedGroupNorm3D( | |
| num_groups=32, | |
| num_channels=num_channels, | |
| eps=1e-6, | |
| affine=True, | |
| cond_channels=cond_channels, | |
| use_t_isolated_gn=use_t_isolated_gn, | |
| padding_mode=padding_mode, | |
| padding_mode_t=padding_mode_t, | |
| causal=causal, | |
| ) | |
| return SpatialNorm3D( | |
| num_channels, | |
| cond_channels, | |
| padding_mode=padding_mode, | |
| padding_mode_t=padding_mode_t, | |
| causal=causal, | |
| use_t_isolated_gn=use_t_isolated_gn, | |
| ) | |
| def get_group_norm_3d(num_channels, use_t_isolated_gn=False): | |
| if os.environ.get("MINIMAX_H3_USE_FUSED_NORM", "false").lower() == "true": | |
| return FusedGroupNorm3D( | |
| num_groups=32, | |
| num_channels=num_channels, | |
| eps=1e-6, | |
| affine=True, | |
| use_t_isolated_gn=use_t_isolated_gn, | |
| ) | |
| norm_cls = ( | |
| TemporalIsolatedSpatialParallelGroupNorm | |
| if use_t_isolated_gn | |
| else SpatialParallelGroupNorm | |
| ) | |
| return norm_cls(num_groups=32, num_channels=num_channels, eps=1e-6, affine=True) | |