Skip to content

kandinsky6_sr

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

CausalConv3d(chan_in: int, chan_out: int, kernel_size: int | tuple[int, int, int], stride: tuple[int, int, int] = (1, 1, 1))

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
encode(x: Tensor) -> tuple[Tensor, list[int]]

[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)