Skip to content

kandinsky6_sr

Kandinsky6 video super-resolution (SR) DiT.

A text-free Kandinsky6 DiT: no text tower and no cross-attention; a learned pooled_bias stands in for the pooled empty-caption embedding the model was trained with. The blocks reuse the Kandinsky6 building blocks (models/dits/kandinsky6.py). Input [B, T, H, W, 2C + 1] (noised latent | anchor | anchor mask), output [B, T, H, W, out_visual_dim] (n_grid * C for the pi-Flow checkpoint).

The residual stream stays in the model's compute_dtype between blocks, as in the Diffusers model. Modulation, norm and time-embedding maths happen in fp32 internally and are cast back to compute_dtype right after each gated residual add. Sequence parallelism is not supported (dense LocalAttention), and NABLA sparse attention is not wired.

Classes

fastvideo.models.dits.kandinsky6_sr.Kandinsky6SRDecoderBlock

Kandinsky6SRDecoderBlock(model_dim: int, time_dim: int, ff_dim: int, head_dim: int, supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None, prefix: str = '', quant_config: QuantizationConfig | None = None)

Bases: Module

Text-free decoder block: modulated self-attention + modulated feed-forward (6 modulation params).

x = (x + gate * attention(LN(x) * (1 + scale) + shift)).to(compute_dtype) per sub-layer: the gated residual add happens in fp32 (matching the norm/modulation maths) and is cast back to compute_dtype immediately after, so the residual stream tracks the model's working dtype between blocks instead of staying fp32 forever after the first block.

Source code in fastvideo/models/dits/kandinsky6_sr.py
def __init__(self,
             model_dim: int,
             time_dim: int,
             ff_dim: int,
             head_dim: int,
             supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
             prefix: str = "",
             quant_config: QuantizationConfig | None = None):
    super().__init__()
    self.visual_modulation = Kandinsky6Modulation(time_dim, model_dim, 6)
    self.self_attention_norm = LayerNormScaleShift(model_dim,
                                                   norm_type="layer",
                                                   eps=1e-5,
                                                   elementwise_affine=False,
                                                   dtype=torch.float32,
                                                   compute_dtype=torch.float32)
    self.self_attention = Kandinsky6Attention(model_dim,
                                              head_dim,
                                              supported_attention_backends=supported_attention_backends,
                                              prefix=f"{prefix}.self_attention",
                                              quant_config=quant_config)
    self.feed_forward_norm = LayerNormScaleShift(model_dim,
                                                 norm_type="layer",
                                                 eps=1e-5,
                                                 elementwise_affine=False,
                                                 dtype=torch.float32,
                                                 compute_dtype=torch.float32)
    self.feed_forward = Kandinsky6FeedForward(model_dim, ff_dim, prefix=f"{prefix}.feed_forward",
                                              quant_config=quant_config)

fastvideo.models.dits.kandinsky6_sr.Kandinsky6SROutLayer

Kandinsky6SROutLayer(model_dim: int, time_dim: int, visual_dim: int, patch_size: tuple[int, int, int])

Bases: Kandinsky6OutLayer

Kandinsky6OutLayer for an fp32 residual stream: modulate in fp32, cast to the parameter dtype, project.

Source code in fastvideo/models/dits/kandinsky6.py
def __init__(self, model_dim: int, time_dim: int, visual_dim: int, patch_size: tuple[int, int, int]):
    super().__init__()
    self.patch_size = patch_size
    self.modulation = Kandinsky6Modulation(time_dim, model_dim, 2)
    self.norm = nn.LayerNorm(model_dim, eps=1e-5, elementwise_affine=False)
    self.out_layer = ReplicatedLinear(model_dim, math.prod(patch_size) * visual_dim, bias=True)

fastvideo.models.dits.kandinsky6_sr.Kandinsky6SRTransformer3DModel

Kandinsky6SRTransformer3DModel(config: Kandinsky6SRConfig, hf_config: dict[str, Any])

Bases: BaseDiT

forward(hidden_states [B, T, H, W, 2C + 1], timestep [B], visual_rope_pos, scale_factor) -> [B, T, H, W, out_visual_dim].

Source code in fastvideo/models/dits/kandinsky6_sr.py
def __init__(self, config: Kandinsky6SRConfig, hf_config: dict[str, Any]) -> None:
    super().__init__(config=config, hf_config=hf_config)
    arch = config.arch_config
    assert isinstance(arch, Kandinsky6SRArchConfig)
    self._reject_unknown_config_keys(arch, hf_config)
    quant_config = config.quant_config
    self.quant_config = quant_config

    head_dim = sum(arch.axes_dims)
    self.in_visual_dim = arch.in_visual_dim
    self.model_dim = arch.model_dim
    self.patch_size = arch.patch_size
    self.input_channels = arch.input_channels

    self.time_embeddings = Kandinsky6TimeEmbeddings(arch.model_dim, arch.time_dim)
    self.pooled_bias = nn.Parameter(torch.zeros(arch.time_dim))

    self.visual_embeddings = Kandinsky6VisualEmbeddings(self.input_channels, arch.model_dim, arch.patch_size)
    self.visual_rope_embeddings = Kandinsky6RoPE3D(arch.axes_dims)
    self.visual_transformer_blocks = nn.ModuleList([
        Kandinsky6SRDecoderBlock(arch.model_dim,
                                 arch.time_dim,
                                 arch.ff_dim,
                                 head_dim,
                                 self._supported_attention_backends,
                                 prefix=f"{config.prefix}.visual_transformer_blocks.{i}",
                                 quant_config=quant_config) for i in range(arch.num_visual_blocks)
    ])
    self.out_layer = Kandinsky6SROutLayer(arch.model_dim, arch.time_dim, arch.out_visual_dim, arch.patch_size)
    if arch.attention_params is not None:
        self._warn_if_sparse_attention_requested(arch.attention_params)
    if arch.nabla_threshold is not None:
        logger.warning(
            "Kandinsky6SR checkpoint requests NABLA sparse attention via nabla_threshold=%s; the SR DiT runs "
            "dense attention (sparse NABLA is not wired), so results differ slightly from the reference.",
            arch.nabla_threshold)

    self.gradient_checkpointing = False
    self.hidden_size = arch.hidden_size
    self.num_attention_heads = arch.num_attention_heads
    self.num_channels_latents = arch.num_channels_latents
    self.__post_init__()

Attributes

fastvideo.models.dits.kandinsky6_sr.Kandinsky6SRTransformer3DModel.compute_dtype property
compute_dtype: dtype

Parameter dtype of the linear stacks (the dtype the DiT input is cast to).

Methods:

fastvideo.models.dits.kandinsky6_sr.Kandinsky6SRTransformer3DModel.forward
forward(hidden_states: Tensor, timestep: Tensor, visual_rope_pos: tuple[Tensor, Tensor, Tensor] | list[Tensor], scale_factor: tuple[float, ...] = (1.0, 1.0, 1.0), return_dict: bool = False, **kwargs: Any) -> Tensor

Run the DiT.

Parameters:

Name Type Description Default
hidden_states Tensor

[B, T, H, W, 2C + 1] noised latent | anchor | anchor mask.

required
timestep Tensor

[B] scheduler timestep (sigma * 1000).

required
visual_rope_pos tuple[Tensor, Tensor, Tensor] | list[Tensor]

(arange(T), arange(H // ph), arange(W // pw)) position indices (shared by the batch).

required
scale_factor tuple[float, ...]

RoPE frequency scale per axis (t, h, w).

(1.0, 1.0, 1.0)
Source code in fastvideo/models/dits/kandinsky6_sr.py
def forward(
    self,
    hidden_states: torch.Tensor,
    timestep: torch.Tensor,
    visual_rope_pos: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | list[torch.Tensor],
    scale_factor: tuple[float, ...] = (1.0, 1.0, 1.0),
    return_dict: bool = False,
    **kwargs: Any,
) -> torch.Tensor:
    """Run the DiT.

    Args:
        hidden_states: ``[B, T, H, W, 2C + 1]`` noised latent | anchor | anchor mask.
        timestep: ``[B]`` scheduler timestep (sigma * 1000).
        visual_rope_pos: ``(arange(T), arange(H // ph), arange(W // pw))`` position indices (shared by the batch).
        scale_factor: RoPE frequency scale per axis (t, h, w).
    """
    if kwargs:
        raise TypeError(f"Kandinsky6SRTransformer3DModel.forward got unsupported arguments {sorted(kwargs)}; the "
                        "SR DiT has no text / sparse-attention / LQ inputs.")
    if hidden_states.ndim != 5:
        raise ValueError(f"hidden_states must be [B,T,H,W,C], got {tuple(hidden_states.shape)}")
    if hidden_states.shape[-1] != self.input_channels:
        raise ValueError(
            f"hidden_states has {hidden_states.shape[-1]} channels but the trained input layer expects "
            f"{self.input_channels} (noised latent | anchor | anchor mask).")

    time_embed = self.time_embeddings(timestep)
    time_embed = time_embed + self.pooled_bias.to(time_embed.dtype)

    visual_embed = self.visual_embeddings(hidden_states.to(self.compute_dtype).contiguous())
    batch_size, duration, height, width, dim = visual_embed.shape
    visual_rope = self.visual_rope_embeddings((batch_size, duration, height, width), visual_rope_pos, scale_factor)
    visual_embed = visual_embed.flatten(1, 3)
    visual_rope = visual_rope.flatten(1, 3)

    compute_dtype = self.compute_dtype
    for block in self.visual_transformer_blocks:
        if torch.is_grad_enabled() and self.gradient_checkpointing:
            visual_embed = torch.utils.checkpoint.checkpoint(block, visual_embed, time_embed, visual_rope,
                                                             compute_dtype, use_reentrant=False)
        else:
            visual_embed = block(visual_embed, time_embed, visual_rope, compute_dtype)

    visual_embed = visual_embed.reshape(batch_size, duration, height, width, dim)
    return self.out_layer(visual_embed, time_embed)

Functions: