MiniMaxAI/MiniMax-H3
6k3.6m
1# SPDX-License-Identifier: Apache-2.02# Spatial-parallel 3D convolution for the MiniMax H3 visual VAE.3import torch4import torch.nn as nn5import torch.nn.functional as F6 7from .parallel import get_parallel_state, exchange_borders8 9 10 11 12class BaseConv3d(nn.Conv3d):13 def __init__(14 self,15 in_channels,16 out_channels,17 kernel_size,18 stride=1,19 padding=0,20 bias=True,21 padding_mode="zeros",22 padding_mode_t=None,23 causal=True,24 ):25 super().__init__(26 in_channels,27 out_channels,28 kernel_size=kernel_size,29 stride=stride,30 padding=padding,31 bias=bias,32 padding_mode=padding_mode,33 )34 padding_mode = "constant" if padding_mode == "zeros" else padding_mode35 padding_mode_t = "constant" if padding_mode_t == "zeros" else padding_mode_t36 self.pad_mode = padding_mode37 self.pad_mode_t = padding_mode_t or ("constant" if causal else "replicate")38 self.causal = causal39 40 def _apply_temporal_padding(self, x):41 B, C, D, H, W = x.shape42 if D > 1:43 pad_size = (44 0,45 0,46 0,47 0,48 self.padding[0] * 2 if self.causal else self.padding[0],49 0 if self.causal else self.padding[0],50 )51 return F.pad(x, pad_size, mode=self.pad_mode_t)52 else:53 if self.pad_mode_t == "constant":54 assert self.causal, "Zeros padding is only supported for causal mode"55 zeros = torch.zeros_like(x[:, :, :1, :, :]).expand(56 -1, -1, self.kernel_size[0] - 1, -1, -157 )58 return torch.cat([zeros, x], dim=2)59 else:60 return x.expand(-1, -1, self.kernel_size[0], -1, -1)61 62 def _apply_padding(self, x):63 if sum(self.padding) == 0:64 return x65 66 x = F.pad(67 x,68 (self.padding[2], self.padding[2], self.padding[1], self.padding[1], 0, 0),69 mode=self.pad_mode,70 )71 72 x = self._apply_temporal_padding(x)73 return x74 75 def forward(self, x):76 if sum(self.padding) == 0:77 return super().forward(x)78 79 x = self._apply_padding(x)80 return F.conv3d(81 x,82 self.weight,83 self.bias,84 stride=self.stride,85 padding=0,86 dilation=self.dilation,87 )88 89 90class SpatialParallelConv3d(BaseConv3d):91 def __init__(92 self,93 in_channels,94 out_channels,95 kernel_size,96 stride=1,97 padding=0,98 bias=True,99 padding_mode="zeros",100 padding_mode_t=None,101 causal=True,102 ):103 super().__init__(104 in_channels,105 out_channels,106 kernel_size=kernel_size,107 stride=stride,108 padding=padding,109 bias=bias,110 padding_mode=padding_mode,111 padding_mode_t=padding_mode_t,112 causal=causal,113 )114 self.spatial_parallel = False115 self.chunk_dim = -1116 117 def _exchange_borders(self, x, sp_rank, sp_size):118 if self.chunk_dim == -1:119 pad = self.padding[2]120 elif self.chunk_dim == -2:121 pad = self.padding[1]122 else:123 raise ValueError(f"Invalid chunk dimension: {self.chunk_dim}")124 125 if pad == 0:126 return x127 128 local_process_group = get_parallel_state()["sp_process_group"]129 return exchange_borders(130 x,131 pad,132 self.pad_mode,133 sp_rank,134 sp_size,135 local_process_group,136 dim=self.chunk_dim,137 )138 139 def _apply_padding(self, x):140 if not self.spatial_parallel:141 return super()._apply_padding(x)142 143 state = get_parallel_state()144 145 x = self._exchange_borders(x, state["sp_rank"], state["sp_size"])146 147 if self.chunk_dim == -1:148 x = F.pad(149 x, (0, 0, self.padding[1], self.padding[1], 0, 0), mode=self.pad_mode150 )151 elif self.chunk_dim == -2:152 x = F.pad(153 x, (self.padding[2], self.padding[2], 0, 0, 0, 0), mode=self.pad_mode154 )155 else:156 raise ValueError(f"Invalid chunk dimension: {self.chunk_dim}")157 158 x = self._apply_temporal_padding(x)159 return x160 