Kandinsky6 SR causal video VAE (KVAE): spatial x16, temporal x4 (causal, 1 + 4k frames).
Clips are encoded / decoded in temporal segments (16 pixel frames, the first one 17; 4 latent frames, the first one 5). Every causal conv carries its last input frames to the next segment, so memory is bounded by one segment whatever the clip length, and the segmentation is part of the numerics: it fixes which frames the first-frame special cases see.
Conventions: normalize_data maps [0, 255] pixels to x / 128 - 1; encode returns (latent, split_list) where latent is the raw (unscaled) posterior mean and split_list the pixel-frame segment sizes; decode returns an object whose .sample is in the normalized pixel range. State-dict keys are the checkpoint's own encoder.* / decoder.* keys, so the loader loads strictly without renaming.
Classes
fastvideo.models.vaes.kandinsky6_sr.CausalConv3d
Bases: Module
Conv3d that is causal in time: the first segment is left-padded with copies of its first frame, later segments with the tail of the previous one; height / width are zero padded.
Source code in fastvideo/models/vaes/kandinsky6_sr.py
| def __init__(self, chan_in: int, chan_out: int, kernel_size: int | tuple[int, int, int],
stride: tuple[int, int, int] = (1, 1, 1)) -> None:
super().__init__()
if isinstance(kernel_size, int):
kernel_size = (kernel_size, ) * 3
time_kernel, height_kernel, width_kernel = kernel_size
self.height_pad = height_kernel // 2
self.width_pad = width_kernel // 2
self.time_pad = time_kernel - 1
self.time_kernel = time_kernel
self.time_stride = stride[0]
self.conv = _ChunkedConv3d(chan_in, chan_out, kernel_size, stride=stride)
|
fastvideo.models.vaes.kandinsky6_sr.Kandinsky6SRVAE
Kandinsky6SRVAE(config: Kandinsky6SRVAEConfig)
Bases: Module
Source code in fastvideo/models/vaes/kandinsky6_sr.py
| def __init__(self, config: Kandinsky6SRVAEConfig) -> None:
super().__init__()
self.config = config
arch = config.arch_config
assert isinstance(arch, Kandinsky6SRVAEArchConfig)
if not arch.encoder_config or not arch.decoder_config:
raise ValueError("Kandinsky6SRVAE needs `encoder_config` and `decoder_config` (KVAE architecture) in the "
"component config.json.")
enc = _arch_kwargs(dict(arch.encoder_config), _REQUIRED_ENCODER_KNOBS, "in_channels", "encoder")
dec = _arch_kwargs(dict(arch.decoder_config), _REQUIRED_DECODER_KNOBS, "out_ch", "decoder")
self.temporal_compression = enc["temporal_compress_times"]
self.encoder = Encoder3D(**enc)
self.decoder = Decoder3D(**dec)
|
Methods:
fastvideo.models.vaes.kandinsky6_sr.Kandinsky6SRVAE.decode
decode(z: Tensor) -> DecoderOutput
[B, C, T', h, w] raw latent -> .sample [B, 3, T, H, W] in the normalized pixel range.
Source code in fastvideo/models/vaes/kandinsky6_sr.py
| def decode(self, z: torch.Tensor) -> DecoderOutput:
"""``[B, C, T', h, w]`` raw latent -> ``.sample`` ``[B, 3, T, H, W]`` in the normalized pixel range."""
segment = _SEGMENT_FRAMES // self.temporal_compression
num_frames = z.size(2)
if num_frames == 1:
split_list = [1]
else:
split_list = [segment] * ((num_frames - 1) // segment)
if (num_frames - 1) % segment:
split_list.append((num_frames - 1) % segment)
split_list[0] += 1
cache = _SegmentCache()
samples = []
for chunk in torch.split(z, split_list, dim=2):
samples.append(self.decoder(chunk, cache))
cache.first = False
return DecoderOutput(sample=torch.cat(samples, dim=2))
|
fastvideo.models.vaes.kandinsky6_sr.Kandinsky6SRVAE.encode
[B, C, T, H, W] normalized pixels -> (raw latent mean, pixel-frame segment sizes).
Source code in fastvideo/models/vaes/kandinsky6_sr.py
| def encode(self, x: torch.Tensor) -> tuple[torch.Tensor, list[int]]:
"""``[B, C, T, H, W]`` normalized pixels -> ``(raw latent mean, pixel-frame segment sizes)``."""
split_list = [_SEGMENT_FRAMES + 1]
remaining = x.size(2) - split_list[0]
while remaining > 0:
split_list.append(_SEGMENT_FRAMES)
remaining -= _SEGMENT_FRAMES
split_list[-1] += remaining
cache = _SegmentCache()
latents = []
for segment in torch.split(x, split_list, dim=2):
moments = self.encoder(segment, cache)
cache.first = False
latents.append(moments.chunk(2, dim=1)[0]) # deterministic posterior: the mean
return torch.cat(latents, dim=2), split_list
|
fastvideo.models.vaes.kandinsky6_sr.PXSDownsample
PXSDownsample(in_channels: int, compress_time: bool)
Bases: Module
x2 spatial (strided conv + channel-averaged pixel-unshuffle) and optional x2 causal temporal (two causal convs + average pooling) downsample; doubles the channels.
Source code in fastvideo/models/vaes/kandinsky6_sr.py
| def __init__(self, in_channels: int, compress_time: bool) -> None:
super().__init__()
out_channels = 2 * in_channels
self.spatial_conv = _ChunkedConv3d(in_channels, out_channels, kernel_size=(1, 3, 3), stride=(1, 2, 2),
padding=(0, 1, 1))
self.compress_time = compress_time
if compress_time:
self.temporal_conv = nn.Sequential(
CausalConv3d(out_channels, out_channels, kernel_size=(2, 1, 1)),
CausalConv3d(out_channels, out_channels, kernel_size=(2, 1, 1), stride=(2, 1, 1)),
)
self.linear = _ChunkedConv3d(out_channels, out_channels, kernel_size=1)
|
fastvideo.models.vaes.kandinsky6_sr.PXSUpsample
PXSUpsample(in_channels: int, compress_time: bool)
Bases: Module
Optional x2 causal temporal (frame repeat + causal conv residual) then x2 spatial (nearest + conv residual) upsample.
Source code in fastvideo/models/vaes/kandinsky6_sr.py
| def __init__(self, in_channels: int, compress_time: bool) -> None:
super().__init__()
self.spatial_conv = _ChunkedConv3d(in_channels, in_channels, kernel_size=(1, 3, 3), padding=(0, 1, 1))
self.compress_time = compress_time
if compress_time:
self.temporal_conv = CausalConv3d(in_channels, in_channels, kernel_size=(3, 1, 1))
self.linear = _ChunkedConv3d(in_channels, in_channels, kernel_size=1)
|
fastvideo.models.vaes.kandinsky6_sr.ResnetBlock3D
ResnetBlock3D(in_channels: int, out_channels: int, zq_channels: int | None = None)
Bases: Module
norm -> SiLU -> causal conv, twice, plus a 1x1 shortcut when the width changes. Decoder blocks use the latent-conditioned SpatialNorm3D (zq_channels set).
Source code in fastvideo/models/vaes/kandinsky6_sr.py
| def __init__(self, in_channels: int, out_channels: int, zq_channels: int | None = None) -> None:
super().__init__()
if zq_channels is None:
self.norm1: nn.Module = _K6RMSNorm(in_channels, images=False)
self.norm2: nn.Module = _K6RMSNorm(out_channels, images=False)
else:
self.norm1 = SpatialNorm3D(in_channels, zq_channels)
self.norm2 = SpatialNorm3D(out_channels, zq_channels)
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3)
self.conv2 = CausalConv3d(out_channels, out_channels, kernel_size=3)
if in_channels != out_channels:
self.nin_shortcut = _ChunkedConv3d(in_channels, out_channels, kernel_size=1)
|
fastvideo.models.vaes.kandinsky6_sr.SpatialNorm3D
SpatialNorm3D(f_channels: int, zq_channels: int)
Bases: Module
RMS norm modulated by the latent: norm(f) * conv_y(zq) + conv_b(zq) with zq nearest-resized to f.
Source code in fastvideo/models/vaes/kandinsky6_sr.py
| def __init__(self, f_channels: int, zq_channels: int) -> None:
super().__init__()
self.norm_layer = _K6RMSNorm(f_channels, images=False)
self.conv_y = _ChunkedConv3d(zq_channels, f_channels, kernel_size=1)
self.conv_b = _ChunkedConv3d(zq_channels, f_channels, kernel_size=1)
|