Fix VAE encoder to load canonical checkpoint key names
This commit is contained in:
parent
807bd64a82
commit
9bec53ac36
1 changed files with 123 additions and 150 deletions
|
|
@ -6,15 +6,18 @@ upcast to FP32 before mean/std normalization (the reference contract).
|
||||||
|
|
||||||
Causal-conv semantics: spatial padding is reflect; temporal padding is causal
|
Causal-conv semantics: spatial padding is reflect; temporal padding is causal
|
||||||
(front-only zeros) with a stride grid that starts at the first input frame
|
(front-only zeros) with a stride grid that starts at the first input frame
|
||||||
(Comfy autopad "same" / ``causal``). For a single input frame the temporal
|
(Comfy autopad "same" / ``causal``). For a single input frame the temporal taps
|
||||||
taps of the kernel are truncated (Comfy ``autopad="causal_zero"``) so the frame
|
of the kernel are truncated (Comfy ``autopad="causal_zero"``) so a keyframe is
|
||||||
is not convolved against zero frames.
|
never convolved against zero frames.
|
||||||
|
|
||||||
|
Weights are read straight from the checkpoint's ``encoder.*`` / ``quant_conv.*``
|
||||||
|
keys (a direct name-for-name copy into plain tensors) and applied by the
|
||||||
|
stateless kernels below.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import math
|
import math
|
||||||
import os
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
@ -28,182 +31,152 @@ 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_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)
|
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)
|
||||||
|
|
||||||
|
CH = 128
|
||||||
|
CH_MULT = (1, 2, 2, 4, 4, 8)
|
||||||
|
SPACE_DOWN = (2, 2, 2, 2, 1, 1)
|
||||||
|
TIME_DOWN = (1, 2, 2, 1, 1, 1)
|
||||||
|
NUM_RES_BLOCKS = 2
|
||||||
|
VAE_RATIO = 16
|
||||||
|
NIN_LEVELS = frozenset({1, 3, 5})
|
||||||
|
DOWNSAMPLE_LEVELS = frozenset({0, 1, 2, 3})
|
||||||
|
|
||||||
|
|
||||||
def _causal_front_padding(t_in: int, kernel: int, stride: int) -> int:
|
def _causal_front_padding(t_in: int, kernel: int, stride: int) -> int:
|
||||||
"""Front-only zero padding matching the reference causal 3D conv.
|
"""Front-only zero padding: ``kernel-1-(ceil(t_in/stride)-1)*stride``."""
|
||||||
|
|
||||||
``padding = kernel - 1 - (t_out - 1) * stride`` with ``t_out = ceil(t_in /
|
|
||||||
stride)`` (the reference computes the output length from the unpadded
|
|
||||||
input). The result is non-negative: with a stride-s grid and a (2k+1)
|
|
||||||
kernel, ``ceil(t/stride) >= 1 + (t-1) // stride`` for all t, so
|
|
||||||
``(t_in - 1) % stride * stride >= (kernel - 1) % (2 * stride)``.
|
|
||||||
"""
|
|
||||||
t_out = math.ceil(t_in / stride)
|
t_out = math.ceil(t_in / stride)
|
||||||
return max(0, kernel - 1 - (t_out - 1) * stride)
|
return max(0, kernel - 1 - (t_out - 1) * stride)
|
||||||
|
|
||||||
|
|
||||||
class _CausalConv3d(nn.Module):
|
def _causal_conv3d(x, weight, bias, *, kernel_size, stride, spatial_padding):
|
||||||
"""3D conv: reflect spatial padding, causal (front-zero) temporal padding."""
|
"""Reflect-spatial / causal-temporal 3D conv (single-frame -> truncated taps)."""
|
||||||
|
t = x.shape[2]
|
||||||
def __init__(self, in_channels: int, out_channels: int, kernel_size: int, stride: int | tuple[int, int, int] = 1, spatial_padding: int = 0):
|
if t == 1:
|
||||||
super().__init__()
|
if spatial_padding > 0:
|
||||||
self.kernel_size = kernel_size
|
half = (kernel_size - 1) // 2
|
||||||
self.stride = stride if isinstance(stride, tuple) else (stride, stride, stride)
|
kernel = weight[:, :, half : half + 2 * spatial_padding + 1]
|
||||||
self.spatial_padding = spatial_padding
|
return F.conv3d(x, kernel, bias, (1, stride[1], stride[2]), (0, spatial_padding, spatial_padding))
|
||||||
self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride=self.stride, padding=(0, spatial_padding, spatial_padding))
|
return F.conv3d(x, weight, bias, stride, (0, spatial_padding, spatial_padding))
|
||||||
|
front = _causal_front_padding(t, kernel_size, stride[0])
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
if front > 0:
|
||||||
t = x.shape[2]
|
x = F.pad(x, (0, 0, 0, 0, front, 0))
|
||||||
if t == 1:
|
return F.conv3d(x, weight, bias, stride, (0, spatial_padding, spatial_padding))
|
||||||
# A single input frame never convolves against zero frames: the
|
|
||||||
# temporal taps are truncated (Comfy autopad="causal_zero").
|
|
||||||
if self.spatial_padding > 0:
|
|
||||||
half = (self.kernel_size - 1) // 2
|
|
||||||
kernel = self.conv.weight[:, :, half:half + 2 * self.spatial_padding + 1]
|
|
||||||
return F.conv3d(x, kernel, self.conv.bias, (1, self.stride[1], self.stride[2]), (0, self.spatial_padding, self.spatial_padding))
|
|
||||||
return self.conv(x)
|
|
||||||
front = _causal_front_padding(t, self.kernel_size, self.stride[0])
|
|
||||||
if front > 0:
|
|
||||||
x = F.pad(x, (0, 0, 0, 0, front, 0))
|
|
||||||
return self.conv(x)
|
|
||||||
|
|
||||||
|
|
||||||
class TemporalIsolatedGroupNorm(nn.GroupNorm):
|
def _group_norm_3d(x, weight, bias):
|
||||||
"""GroupNorm with statistics computed per frame (time merged into batch)."""
|
"""GroupNorm (32 groups, eps 1e-6) with per-frame statistics."""
|
||||||
|
b, c, t, h, w = x.shape
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
y = F.group_norm(x.permute(0, 2, 1, 3, 4).contiguous().view(b * t, c, 1, h, w), 32, weight, bias, 1e-6)
|
||||||
if x.dim() != 5:
|
return y.view(b, t, c, h, w).permute(0, 2, 1, 3, 4).contiguous()
|
||||||
return super().forward(x)
|
|
||||||
b, c, t, h, w = x.shape
|
|
||||||
x = x.permute(0, 2, 1, 3, 4).contiguous().view(b * t, c, 1, h, w)
|
|
||||||
x = super().forward(x)
|
|
||||||
return x.view(b, t, c, h, w).permute(0, 2, 1, 3, 4).contiguous()
|
|
||||||
|
|
||||||
|
|
||||||
def group_norm_3d(num_channels: int) -> TemporalIsolatedGroupNorm:
|
def _resnet(x, p):
|
||||||
return TemporalIsolatedGroupNorm(num_groups=32, num_channels=num_channels, eps=1e-6, affine=True)
|
residual = x if p["nin"] is None else _causal_conv3d(x, p["nin"][0], p["nin"][1], kernel_size=1, stride=(1, 1, 1), spatial_padding=0)
|
||||||
|
h = _causal_conv3d(F.silu(_group_norm_3d(x, p["norm1_w"], p["norm1_b"])), p["conv1_w"], p["conv1_b"], kernel_size=3, stride=(1, 1, 1), spatial_padding=1)
|
||||||
|
h = _causal_conv3d(F.silu(_group_norm_3d(h, p["norm2_w"], p["norm2_b"])), p["conv2_w"], p["conv2_b"], kernel_size=3, stride=(1, 1, 1), spatial_padding=1)
|
||||||
|
return h.add_(residual)
|
||||||
|
|
||||||
|
|
||||||
class Downsample3D(nn.Module):
|
def _downsample(x, p):
|
||||||
def __init__(self, in_channels: int, out_channels: int, time_stride: int = 1, space_stride: int = 2):
|
if p["space"] == 2:
|
||||||
super().__init__()
|
x = F.pad(x, (0, 1, 0, 1, 0, 0), mode="reflect")
|
||||||
self.space_stride = space_stride
|
return _causal_conv3d(x, p["w"], p["b"], kernel_size=3, stride=(p["time"], p["space"], p["space"]), spatial_padding=0)
|
||||||
self.conv = _CausalConv3d(in_channels, out_channels, kernel_size=3, stride=(time_stride, space_stride, space_stride), spatial_padding=0)
|
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
||||||
if self.space_stride == 2:
|
|
||||||
x = F.pad(x, (0, 1, 0, 1, 0, 0), mode="reflect")
|
|
||||||
return self.conv(x)
|
|
||||||
|
|
||||||
|
|
||||||
class ResnetBlock3D(nn.Module):
|
def _encoder_run(x, E):
|
||||||
def __init__(self, in_channels: int, out_channels: int | None = None):
|
h = _causal_conv3d(x, E["conv_in"][0], E["conv_in"][1], kernel_size=3, stride=(1, 1, 1), spatial_padding=1)
|
||||||
super().__init__()
|
for level in E["down"]:
|
||||||
self.in_channels = in_channels
|
for blk in level["blocks"]:
|
||||||
self.out_channels = in_channels if out_channels is None else out_channels
|
h = _resnet(h, blk)
|
||||||
self.norm1 = group_norm_3d(in_channels)
|
if level["down"] is not None:
|
||||||
self.norm2 = group_norm_3d(self.out_channels)
|
h = _downsample(h, level["down"])
|
||||||
self.conv1 = _CausalConv3d(in_channels, self.out_channels, kernel_size=3, spatial_padding=1)
|
h = F.silu(_group_norm_3d(h, E["norm_out_w"], E["norm_out_b"]))
|
||||||
self.conv2 = _CausalConv3d(self.out_channels, self.out_channels, kernel_size=3, spatial_padding=1)
|
return _causal_conv3d(h, E["conv_out"][0], E["conv_out"][1], kernel_size=3, stride=(1, 1, 1), spatial_padding=1)
|
||||||
if self.in_channels != self.out_channels:
|
|
||||||
self.nin_shortcut = _CausalConv3d(self.in_channels, self.out_channels, kernel_size=1, spatial_padding=0)
|
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
||||||
h = self.conv1(F.silu(self.norm1(x), inplace=True))
|
|
||||||
h = self.conv2(F.silu(self.norm2(h), inplace=True))
|
|
||||||
if self.in_channels != self.out_channels:
|
|
||||||
x = self.nin_shortcut(x)
|
|
||||||
return h.add_(x)
|
|
||||||
|
|
||||||
|
|
||||||
class EncoderFCN3D(nn.Module):
|
|
||||||
def __init__(self, ch: int, ch_mult: tuple[int, ...], space_down: tuple[int, ...], time_down: tuple[int, ...], num_res_blocks: int, in_channels: int, z_channels: int, double_z: bool = True):
|
|
||||||
super().__init__()
|
|
||||||
self.num_levels = len(ch_mult)
|
|
||||||
self.num_res_blocks = [num_res_blocks] * self.num_levels
|
|
||||||
block_mid = [ch * ch_mult[i] for i in range(self.num_levels)]
|
|
||||||
block_in = [block_mid[0]] + block_mid[:-1]
|
|
||||||
self.conv_in = _CausalConv3d(in_channels, block_in[0], kernel_size=3, spatial_padding=1)
|
|
||||||
self.down = nn.ModuleList()
|
|
||||||
for i_level in range(self.num_levels):
|
|
||||||
down = nn.Module()
|
|
||||||
down.block = nn.ModuleList(
|
|
||||||
ResnetBlock3D(block_in[i_level] if i == 0 else block_mid[i_level], block_mid[i_level])
|
|
||||||
for i in range(self.num_res_blocks[i_level])
|
|
||||||
)
|
|
||||||
if space_down[i_level] * time_down[i_level] > 1:
|
|
||||||
down.downsample = Downsample3D(block_mid[i_level], block_mid[i_level], time_stride=time_down[i_level], space_stride=space_down[i_level])
|
|
||||||
self.down.append(down)
|
|
||||||
self.norm_out = group_norm_3d(block_mid[-1])
|
|
||||||
self.conv_out = _CausalConv3d(block_mid[-1], 2 * z_channels if double_z else z_channels, kernel_size=3, spatial_padding=1)
|
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
||||||
h = self.conv_in(x)
|
|
||||||
for i_level in range(self.num_levels):
|
|
||||||
for i_block in range(self.num_res_blocks[i_level]):
|
|
||||||
h = self.down[i_level].block[i_block](h)
|
|
||||||
if hasattr(self.down[i_level], "downsample"):
|
|
||||||
h = self.down[i_level].downsample(h)
|
|
||||||
h = F.silu(self.norm_out(h), inplace=True)
|
|
||||||
return self.conv_out(h)
|
|
||||||
|
|
||||||
|
|
||||||
class MiniMaxH3VideoVAEEncoder(nn.Module):
|
class MiniMaxH3VideoVAEEncoder(nn.Module):
|
||||||
"""Encoder-only H3 VAE. ``encode`` matches the public contract of upstream_vae.py."""
|
"""Encoder-only H3 VAE. ``encode`` matches the public contract of upstream_vae.py.
|
||||||
|
|
||||||
def __init__(self, *, device: torch.device | str | None = None, tiling: bool = True):
|
Weights are plain tensors loaded from the checkpoint by canonical name into
|
||||||
|
``self.W`` (a dict), so no ``nn.Module`` sub-hierarchy is needed.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, *, tiling: bool = True):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.vae_ratio, self.vae_ratio_t = 16, 4
|
self.vae_ratio, self.vae_ratio_t = VAE_RATIO, 4
|
||||||
self.clip_length, self.token_drop = 17, 3
|
self.clip_length, self.token_drop = 17, 3
|
||||||
self.tiling, self.tile_size, self.tile_overlap_min = tiling, 256, 64
|
self.tiling, self.tile_size, self.tile_overlap_min = tiling, 256, 64
|
||||||
self.encoder = EncoderFCN3D(ch=128, ch_mult=(1, 2, 2, 4, 4, 8), space_down=(2, 2, 2, 2, 1, 1), time_down=(1, 2, 2, 1, 1, 1), num_res_blocks=2, in_channels=3, z_channels=24, double_z=True)
|
self.quant_conv = nn.Conv3d(48, 48, 1)
|
||||||
self.quant_conv = nn.Conv3d(48, 48, 1, device=device)
|
|
||||||
self.register_buffer("latents_mean", torch.tensor(LATENTS_MEAN), persistent=False)
|
self.register_buffer("latents_mean", torch.tensor(LATENTS_MEAN), persistent=False)
|
||||||
self.register_buffer("latents_std", torch.tensor(LATENTS_STD), persistent=False)
|
self.register_buffer("latents_std", torch.tensor(LATENTS_STD), persistent=False)
|
||||||
self.register_buffer("pixel_mean", torch.tensor(IMAGENET_MEAN).view(1, 3, 1, 1, 1), persistent=False)
|
self.register_buffer("pixel_mean", torch.tensor(IMAGENET_MEAN).view(1, 3, 1, 1, 1), persistent=False)
|
||||||
self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD).view(1, 3, 1, 1, 1), persistent=False)
|
self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD).view(1, 3, 1, 1, 1), persistent=False)
|
||||||
|
|
||||||
|
def _required_encoder_names(self) -> list[str]:
|
||||||
|
names = ["encoder.conv_in.weight", "encoder.conv_in.bias", "encoder.norm_out.weight", "encoder.norm_out.bias", "encoder.conv_out.weight", "encoder.conv_out.bias"]
|
||||||
|
for i in range(len(CH_MULT)):
|
||||||
|
for b in range(NUM_RES_BLOCKS):
|
||||||
|
base = f"encoder.down.{i}.block.{b}."
|
||||||
|
names += [
|
||||||
|
base + "conv1.weight", base + "conv1.bias",
|
||||||
|
base + "conv2.weight", base + "conv2.bias",
|
||||||
|
base + "norm1.weight", base + "norm1.bias",
|
||||||
|
base + "norm2.weight", base + "norm2.bias",
|
||||||
|
]
|
||||||
|
if b == 0 and i in NIN_LEVELS:
|
||||||
|
names += [base + "nin_shortcut.weight", base + "nin_shortcut.bias"]
|
||||||
|
if i in DOWNSAMPLE_LEVELS:
|
||||||
|
names += [f"encoder.down.{i}.downsample.conv.weight", f"encoder.down.{i}.downsample.conv.bias"]
|
||||||
|
return names
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_safetensors(cls, path: str | Path, *, device: str | torch.device = "cuda", tiling: bool = True) -> "MiniMaxH3VideoVAEEncoder":
|
def from_safetensors(cls, path: str | Path, *, device: str | torch.device = "cuda", tiling: bool = True) -> "MiniMaxH3VideoVAEEncoder":
|
||||||
model = cls(device="meta", tiling=tiling)
|
model = cls(tiling=tiling)
|
||||||
expected = model.state_dict()
|
names = model._required_encoder_names()
|
||||||
if os.getenv("H3_FAST_SAFETENSORS", "").lower() in {"1", "true", "yes", "on"}:
|
with safe_open(str(path), framework="pt", device=str(device)) as ck:
|
||||||
from fastsafetensors import fastsafe_open
|
available = set(ck.keys())
|
||||||
|
missing = [n for n in names if n not in available]
|
||||||
fast_device = "cuda:0" if str(device) == "cuda" else str(device)
|
|
||||||
available = set()
|
|
||||||
with fastsafe_open(filenames=[str(path)], nogds=True, device=fast_device) as checkpoint:
|
|
||||||
available = set(checkpoint.keys())
|
|
||||||
missing = sorted(set(expected) - available)
|
|
||||||
weights = {name: checkpoint.get_tensor(name).clone().detach() for name in expected}
|
|
||||||
if missing:
|
if missing:
|
||||||
raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(missing)}")
|
raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(sorted(missing)[:16])}")
|
||||||
weights = {name: t.to(device=device, dtype=torch.float32) for name, t in weights.items()}
|
W = {n: ck.get_tensor(n).to(dtype=torch.float32).to(device) for n in names}
|
||||||
elif os.getenv("H3_DISABLE_MMAP", "").lower() in {"1", "true", "yes", "on"}:
|
|
||||||
from safetensors.torch import load
|
|
||||||
|
|
||||||
with open(path, "rb") as file:
|
# Build the structured params dict.
|
||||||
available_weights = load(file.read())
|
down = []
|
||||||
available = set(available_weights)
|
for i in range(len(CH_MULT)):
|
||||||
missing = sorted(set(expected) - available)
|
mid = CH * CH_MULT[i]
|
||||||
if missing:
|
blocks = []
|
||||||
raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(missing)}")
|
for b in range(NUM_RES_BLOCKS):
|
||||||
weights = {name: available_weights[name].to(device=device, dtype=torch.float32) for name in expected}
|
base = f"encoder.down.{i}.block.{b}."
|
||||||
else:
|
blk = {
|
||||||
with safe_open(str(path), framework="pt", device=str(device)) as checkpoint:
|
"conv1_w": W[base + "conv1.weight"], "conv1_b": W[base + "conv1.bias"],
|
||||||
available = set(checkpoint.keys())
|
"conv2_w": W[base + "conv2.weight"], "conv2_b": W[base + "conv2.bias"],
|
||||||
shape_errors = [(name, tuple(expected[name].shape), tuple(checkpoint.get_slice(name).get_shape())) for name in expected if name in available and tuple(expected[name].shape) != tuple(checkpoint.get_slice(name).get_shape())]
|
"norm1_w": W[base + "norm1.weight"], "norm1_b": W[base + "norm1.bias"],
|
||||||
if shape_errors:
|
"norm2_w": W[base + "norm2.weight"], "norm2_b": W[base + "norm2.bias"],
|
||||||
raise ValueError(f"incompatible H3 VAE checkpoint; shape mismatch: {shape_errors}")
|
"nin": None,
|
||||||
weights = {name: checkpoint.get_tensor(name).to(dtype=torch.float32) for name in expected}
|
}
|
||||||
missing = sorted(set(expected) - available)
|
if base + "nin_shortcut.weight" in W:
|
||||||
if missing:
|
blk["nin"] = (W[base + "nin_shortcut.weight"], W[base + "nin_shortcut.bias"])
|
||||||
raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(missing)}")
|
blocks.append(blk)
|
||||||
model.load_state_dict(weights, strict=True, assign=True)
|
level_down = None
|
||||||
|
if i in DOWNSAMPLE_LEVELS:
|
||||||
|
ds = f"encoder.down.{i}.downsample.conv."
|
||||||
|
level_down = {"w": W[ds + "weight"], "b": W[ds + "bias"], "time": TIME_DOWN[i], "space": SPACE_DOWN[i]}
|
||||||
|
down.append({"blocks": blocks, "down": level_down})
|
||||||
|
|
||||||
|
E = {
|
||||||
|
"conv_in": (W["encoder.conv_in.weight"], W["encoder.conv_in.bias"]),
|
||||||
|
"down": down,
|
||||||
|
"norm_out_w": W["encoder.norm_out.weight"], "norm_out_b": W["encoder.norm_out.bias"],
|
||||||
|
"conv_out": (W["encoder.conv_out.weight"], W["encoder.conv_out.bias"]),
|
||||||
|
}
|
||||||
|
model.W = E
|
||||||
|
model.quant_conv.to(device, torch.float32)
|
||||||
|
for b in ("latents_mean", "latents_std", "pixel_mean", "pixel_std"):
|
||||||
|
getattr(model, b).to(device)
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
@torch.inference_mode()
|
||||||
def _encode_moments(self, x: torch.Tensor) -> torch.Tensor:
|
def _encode_moments(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
return self.quant_conv(self.encoder(x))
|
return F.conv3d(_encoder_run(x.to(torch.float32), self.W), self.quant_conv.weight, self.quant_conv.bias)
|
||||||
|
|
||||||
def _adaptive_encode(self, x: torch.Tensor) -> torch.Tensor:
|
def _adaptive_encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
if self.tiling:
|
if self.tiling:
|
||||||
|
|
@ -259,7 +232,7 @@ class MiniMaxH3VideoVAEEncoder(nn.Module):
|
||||||
if i > 0:
|
if i > 0:
|
||||||
tile = self.blend(rows[i - 1][j], tile, latent_y_overlap[i - 1], dim=-2)
|
tile = self.blend(rows[i - 1][j], tile, latent_y_overlap[i - 1], dim=-2)
|
||||||
if j > 0:
|
if j > 0:
|
||||||
tile = self.blend(row[j - 1], tile, latent_x_overlap[j - 1], dim=-1)
|
tile = self.blend(row[j - 1], tile, latent_x_overlap[j], dim=-1)
|
||||||
if i < len(rows) - 1:
|
if i < len(rows) - 1:
|
||||||
tile = tile[..., :-latent_y_overlap[i], :]
|
tile = tile[..., :-latent_y_overlap[i], :]
|
||||||
if j < len(row) - 1:
|
if j < len(row) - 1:
|
||||||
|
|
@ -273,14 +246,14 @@ class MiniMaxH3VideoVAEEncoder(nn.Module):
|
||||||
pad_size = (-x.shape[2]) % self.clip_length
|
pad_size = (-x.shape[2]) % self.clip_length
|
||||||
x = torch.cat([x, x[:, :, -1:].repeat(1, 1, pad_size, 1, 1)], dim=2)
|
x = torch.cat([x, x[:, :, -1:].repeat(1, 1, pad_size, 1, 1)], dim=2)
|
||||||
num_chunks = x.shape[2] // self.clip_length
|
num_chunks = x.shape[2] // self.clip_length
|
||||||
z_list = [self._adaptive_encode(x[:, :, i * self.clip_length:(i + 1) * self.clip_length, :, :]) for i in range(num_chunks)]
|
z_list = [self._adaptive_encode(x[:, :, i * self.clip_length : (i + 1) * self.clip_length, :, :]) for i in range(num_chunks)]
|
||||||
z = torch.cat(z_list, dim=2)
|
z = torch.cat(z_list, dim=2)
|
||||||
if self.token_drop > 0:
|
if self.token_drop > 0:
|
||||||
z = z[:, :, :-self.token_drop]
|
z = z[:, :, :-self.token_drop]
|
||||||
return z
|
return z
|
||||||
|
|
||||||
def encode(self, x: torch.Tensor) -> torch.Tensor:
|
def encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
"""``[B, 3, T, H, W]`` pixels in ``[-1, 1]`` -> normalized latents ``[B, 24, T_lat, H//16, W//16]``."""
|
"""``[B,3,H,W]`` or ``[B,3,T,H,W]`` pixels in ``[-1, 1]`` -> normalized latents ``[B,24,T_lat,H//16,W//16]``."""
|
||||||
if x.ndim == 4:
|
if x.ndim == 4:
|
||||||
x = x.unsqueeze(2)
|
x = x.unsqueeze(2)
|
||||||
x = x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x))
|
x = x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x))
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue