Skip to content

kandinsky6

Native FastVideo implementation of the Kandinsky6 T2VA/IT2VA transformer.

Kandinsky6 extends Kandinsky5's DiT (fastvideo/models/dits/kandinsky5.py) with a second, parallel "audio" tower and a fused video<->audio decoder block, so this file mirrors kandinsky5.py's structure closely for every shared piece (time/text/visual embeddings, RoPE, modulation, feed-forward, NABLA sparse attention) and only adds what's new: the audio tower and Kandinsky6FusedTransformerDecoderBlock's cross-modal attention, ported from the reference diffusers.models.transformers.transformer_kandinsky6 module.

Video and audio latents are ordinary batched tensors -- [B, T, H, W, C] for video and [B, A, D] for audio -- consistent with Kandinsky5 and the rest of FastVideo's pipeline/stage system. Kandinsky6RoPE3D.forward takes an explicit shape tuple rather than deriving T/H/W from the position tensors' lengths.

Classes

fastvideo.models.dits.kandinsky6.Kandinsky6Attention

Kandinsky6Attention(num_channels: int, head_dim: int, supported_attention_backends: tuple[AttentionBackendEnum, ...] | None, prefix: str = '', kv_dim: int | None = None, use_nabla: bool = False, quant_config: QuantizationConfig | None = None)

Bases: Module

Self- or cross-attention. kv_dim lets K/V come from a differently-sized stream (the video<->audio cross-modal attentions).

Source code in fastvideo/models/dits/kandinsky6.py
def __init__(
    self,
    num_channels: int,
    head_dim: int,
    supported_attention_backends: tuple[AttentionBackendEnum, ...] | None,
    prefix: str = "",
    kv_dim: int | None = None,
    use_nabla: bool = False,
    quant_config: QuantizationConfig | None = None,
):
    super().__init__()
    assert num_channels % head_dim == 0
    self.num_heads = num_channels // head_dim
    kv_dim = kv_dim or num_channels

    self.to_query = ReplicatedLinear(num_channels,
                                     num_channels,
                                     bias=True,
                                     quant_config=quant_config,
                                     prefix=f"{prefix}.to_query")
    self.to_key = ReplicatedLinear(kv_dim,
                                   num_channels,
                                   bias=True,
                                   quant_config=quant_config,
                                   prefix=f"{prefix}.to_key")
    self.to_value = ReplicatedLinear(kv_dim,
                                     num_channels,
                                     bias=True,
                                     quant_config=quant_config,
                                     prefix=f"{prefix}.to_value")
    self.query_norm = nn.RMSNorm(head_dim)
    self.key_norm = nn.RMSNorm(head_dim)
    self.out_layer = ReplicatedLinear(num_channels,
                                      num_channels,
                                      bias=True,
                                      quant_config=quant_config,
                                      prefix=f"{prefix}.out_layer")
    self.local_attention = LocalAttention(
        num_heads=self.num_heads,
        head_size=head_dim,
        causal=False,
        supported_attention_backends=supported_attention_backends,
    )
    # Only the video self-attention gets a second NABLA-backed attention
    # layer; audio self-attention and every cross-attention are dense.
    self.nabla_attention = None
    if use_nabla:
        self.nabla_attention = LocalAttention(
            num_heads=self.num_heads,
            head_size=head_dim,
            causal=False,
            supported_attention_backends=supported_attention_backends,
            default_backend=AttentionBackendEnum.NABLA_ATTN,
        )

fastvideo.models.dits.kandinsky6.Kandinsky6FusedTransformerDecoderBlock

Kandinsky6FusedTransformerDecoderBlock(model_dim: int, time_dim: int, ff_dim: int, head_dim: int, model_dim_a: int, time_dim_a: int, ff_dim_a: int, head_dim_a: int, supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None, prefix: str = '', use_nabla: bool = False, ca_rope: bool = False, cross_gates: bool = False, fix_modulation: bool = False, quant_config: QuantizationConfig | None = None)

Bases: Module

Joint video+audio decoder block: per-modality self/text-cross attention plus a dedicated bidirectional video<->audio cross-attention, all independently AdaLN-modulated. Ported from diffusers.models.transformers.transformer_kandinsky6.Kandinsky6FusedTransformerDecoderBlock.

Source code in fastvideo/models/dits/kandinsky6.py
def __init__(
    self,
    model_dim: int,
    time_dim: int,
    ff_dim: int,
    head_dim: int,
    model_dim_a: int,
    time_dim_a: int,
    ff_dim_a: int,
    head_dim_a: int,
    supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
    prefix: str = "",
    use_nabla: bool = False,
    ca_rope: bool = False,
    cross_gates: bool = False,
    fix_modulation: bool = False,
    quant_config: QuantizationConfig | None = None,
):
    super().__init__()
    self.videoT = Kandinsky6TransformerDecoderBlock(model_dim,
                                                    time_dim,
                                                    ff_dim,
                                                    head_dim,
                                                    supported_attention_backends,
                                                    prefix=f"{prefix}.videoT",
                                                    use_nabla=use_nabla,
                                                    quant_config=quant_config)
    self.audioT = Kandinsky6TransformerDecoderBlock(model_dim_a,
                                                    time_dim_a,
                                                    ff_dim_a,
                                                    head_dim_a,
                                                    supported_attention_backends,
                                                    prefix=f"{prefix}.audioT",
                                                    use_nabla=False,
                                                    quant_config=quant_config)
    self.va_cross_attention = Kandinsky6Attention(model_dim,
                                                  head_dim,
                                                  supported_attention_backends=supported_attention_backends,
                                                  prefix=f"{prefix}.va_cross_attention",
                                                  kv_dim=model_dim_a,
                                                  quant_config=quant_config)
    self.av_cross_attention = Kandinsky6Attention(model_dim_a,
                                                  head_dim_a,
                                                  supported_attention_backends=supported_attention_backends,
                                                  prefix=f"{prefix}.av_cross_attention",
                                                  kv_dim=model_dim,
                                                  quant_config=quant_config)
    self.va_modulation = Kandinsky6Modulation(
        time_dim,
        model_dim if not cross_gates else model_dim * 2 + model_dim_a,
        1 if cross_gates else 3,
    )
    self.av_modulation = Kandinsky6Modulation(
        time_dim_a,
        model_dim_a if not cross_gates else model_dim_a * 2 + model_dim,
        1 if cross_gates else 3,
    )
    self.va_normalization = nn.LayerNorm(model_dim, elementwise_affine=False)
    self.av_normalization = nn.LayerNorm(model_dim_a, elementwise_affine=False)
    self.ca_rope = ca_rope
    self.cross_gates = cross_gates
    self.fix_modulation = fix_modulation
    self.model_dim = model_dim
    self.model_dim_a = model_dim_a

Methods:

fastvideo.models.dits.kandinsky6.Kandinsky6FusedTransformerDecoderBlock.forward
forward(vis: Tensor | None, aud: Tensor | None, text_v: Tensor, text_a: Tensor | None, time_embed: tuple[Tensor, Tensor | None], vis_rope: Tensor | None, aud_rope: Tensor | None, sparse_params: dict[str, Any] | None, va_gate_scale: float = 1.0, av_gate_scale: float = 1.0) -> tuple[Tensor | None, Tensor | None]

vis/aud may each be None (matches the diffusers reference's guarded Kandinsky6FusedTransformerDecoderBlock): every stage below only runs for a present modality, and the video<->audio cross-modal mixing only runs when both are present.

Source code in fastvideo/models/dits/kandinsky6.py
def forward(
    self,
    vis: torch.Tensor | None,
    aud: torch.Tensor | None,
    text_v: torch.Tensor,
    text_a: torch.Tensor | None,
    time_embed: tuple[torch.Tensor, torch.Tensor | None],
    vis_rope: torch.Tensor | None,
    aud_rope: torch.Tensor | None,
    sparse_params: dict[str, Any] | None,
    va_gate_scale: float = 1.0,
    av_gate_scale: float = 1.0,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
    """``vis``/``aud`` may each be ``None`` (matches the diffusers reference's guarded
    ``Kandinsky6FusedTransformerDecoderBlock``): every stage below only runs for a present
    modality, and the video<->audio cross-modal mixing only runs when both are present."""
    t_v, t_a = time_embed
    gate_v = vis_out_t = None

    if vis is not None:
        sa_p, ca_p, ff_p = torch.chunk(self.videoT.visual_modulation(t_v).unsqueeze(dim=1), 3, dim=-1)
        shift, scale, gate = torch.chunk(sa_p, 3, dim=-1)
        vis = _apply_gate_sum(
            vis,
            self.videoT.self_attention(
                self.videoT.self_attention_norm(vis.float(), shift=shift, scale=scale,
                                                convert_modulation_dtype=True).type_as(vis),
                rotary_emb=vis_rope,
                sparse_params=sparse_params,
            ),
            gate,
        )
        shift, scale, gate_v = torch.chunk(ca_p, 3, dim=-1)
        vis_pre_ca = self.videoT.cross_attention_norm(vis.float(), shift=shift, scale=scale,
                                                      convert_modulation_dtype=True).type_as(vis)
        vis_out_t = self.videoT.cross_attention(vis_pre_ca, encoder_hidden_states=text_v)

    if aud is not None:
        sa_p, ca_p, ff_p_a = torch.chunk(self.audioT.visual_modulation(t_a).unsqueeze(dim=1), 3, dim=-1)
        shift, scale, gate = torch.chunk(sa_p, 3, dim=-1)
        aud = _apply_gate_sum(
            aud,
            self.audioT.self_attention(
                self.audioT.self_attention_norm(aud.float(), shift=shift, scale=scale,
                                                convert_modulation_dtype=True).type_as(aud),
                rotary_emb=aud_rope,
            ),
            gate,
        )
        shift, scale, gate_a = torch.chunk(ca_p, 3, dim=-1)
        aud_pre_ca = self.audioT.cross_attention_norm(aud.float(), shift=shift, scale=scale,
                                                       convert_modulation_dtype=True).type_as(aud)
        aud_out_t = self.audioT.cross_attention(aud_pre_ca, encoder_hidden_states=text_a)
        aud = _apply_gate_sum(aud, aud_out_t, gate_a)

        if vis is not None:
            t_va_mod = t_a if not self.fix_modulation else t_v
            t_av_mod = t_v if not self.fix_modulation else t_a
            va_params = self.va_modulation(t_va_mod).unsqueeze(dim=1)
            av_params = self.av_modulation(t_av_mod).unsqueeze(dim=1)
            if self.cross_gates:
                va_shift, va_scale, va_gate = torch.split(
                    va_params, [self.model_dim, self.model_dim, self.model_dim_a], dim=-1)
                av_shift, av_scale, av_gate = torch.split(
                    av_params, [self.model_dim_a, self.model_dim_a, self.model_dim], dim=-1)
            else:
                va_shift, va_scale, va_gate = torch.chunk(va_params, 3, dim=-1)
                av_shift, av_scale, av_gate = torch.chunk(av_params, 3, dim=-1)

            vis = _apply_gate_sum(vis, vis_out_t, gate_v)
            vis_for_va = _apply_scale_shift(self.va_normalization, vis, va_scale, va_shift)
            aud_for_av = _apply_scale_shift(self.av_normalization, aud, av_scale, av_shift)
            rq_v = vis_rope if self.ca_rope else None
            rk_a = aud_rope if self.ca_rope else None
            vis_from_aud = self.va_cross_attention(vis_for_va, encoder_hidden_states=aud_pre_ca, rotary_emb=rq_v,
                                                   rotary_emb_kv=rk_a)
            aud_from_vis = self.av_cross_attention(aud_for_av, encoder_hidden_states=vis_pre_ca, rotary_emb=rk_a,
                                                   rotary_emb_kv=rq_v)
            vis = _apply_gate_sum(vis, vis_from_aud, (va_gate if not self.cross_gates else av_gate) * va_gate_scale)
            aud = _apply_gate_sum(aud, aud_from_vis, (av_gate if not self.cross_gates else va_gate) * av_gate_scale)
    elif vis is not None:
        vis = _apply_gate_sum(vis, vis_out_t, gate_v)

    if vis is not None:
        shift, scale, gate = torch.chunk(ff_p, 3, dim=-1)
        vis = _apply_gate_sum(
            vis,
            self.videoT.feed_forward(
                self.videoT.feed_forward_norm(vis.float(), shift=shift, scale=scale,
                                              convert_modulation_dtype=True).type_as(vis)),
            gate,
        )
    if aud is not None:
        shift, scale, gate = torch.chunk(ff_p_a, 3, dim=-1)
        aud = _apply_gate_sum(
            aud,
            self.audioT.feed_forward(
                self.audioT.feed_forward_norm(aud.float(), shift=shift, scale=scale,
                                              convert_modulation_dtype=True).type_as(aud)),
            gate,
        )
    return vis, aud

fastvideo.models.dits.kandinsky6.Kandinsky6OutLayer

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

Bases: Module

Projects visual hidden states back to packed latent patches.

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.Kandinsky6OutLayerAudio

Kandinsky6OutLayerAudio(model_dim: int, time_dim: int, audio_dim: int)

Bases: Module

Projects audio hidden states back to audio latent channels.

Source code in fastvideo/models/dits/kandinsky6.py
def __init__(self, model_dim: int, time_dim: int, audio_dim: int):
    super().__init__()
    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, audio_dim, bias=True)

fastvideo.models.dits.kandinsky6.Kandinsky6RoPE1D

Kandinsky6RoPE1D(dim: int, max_pos: int = 2048, max_period: float = 10000.0, freqs_scaling: float = 1.0)

Bases: Module

1D rotary embedding for text and audio sequences.

Source code in fastvideo/models/dits/kandinsky6.py
def __init__(self, dim: int, max_pos: int = 2048, max_period: float = 10000.0, freqs_scaling: float = 1.0):
    super().__init__()
    self.max_period = max_period
    self.dim = dim
    self.max_pos = max_pos
    self.freqs_scaling = freqs_scaling
    freq = _build_rotary_freqs(dim // 2, max_period) * freqs_scaling
    pos = torch.arange(max_pos, dtype=freq.dtype)
    self.register_buffer("args", torch.outer(pos, freq), persistent=False)

fastvideo.models.dits.kandinsky6.Kandinsky6RoPE3D

Kandinsky6RoPE3D(axes_dims: tuple[int, int, int], max_pos: tuple[int, int, int] = (128, 128, 128), max_period: float = 10000.0)

Bases: Module

3D rotary embedding for video spatial-temporal (T, H, W) tokens.

Source code in fastvideo/models/dits/kandinsky6.py
def __init__(self,
             axes_dims: tuple[int, int, int],
             max_pos: tuple[int, int, int] = (128, 128, 128),
             max_period: float = 10000.0):
    super().__init__()
    self.axes_dims = axes_dims
    self.max_pos = max_pos
    self.max_period = max_period

    for i, (axes_dim, ax_max_pos) in enumerate(zip(axes_dims, max_pos, strict=True)):
        freq = _build_rotary_freqs(axes_dim // 2, max_period)
        pos = torch.arange(ax_max_pos, dtype=freq.dtype)
        self.register_buffer(f"args_{i}", torch.outer(pos, freq), persistent=False)

fastvideo.models.dits.kandinsky6.Kandinsky6TextEmbeddings

Kandinsky6TextEmbeddings(in_dim: int, model_dim: int)

Bases: Module

Linear + LayerNorm projection, reused for text tokens and audio latents.

Source code in fastvideo/models/dits/kandinsky6.py
def __init__(self, in_dim: int, model_dim: int):
    super().__init__()
    self.in_layer = ReplicatedLinear(in_dim, model_dim, bias=True)
    self.norm = nn.LayerNorm(model_dim, elementwise_affine=True)

fastvideo.models.dits.kandinsky6.Kandinsky6Transformer3DModel

Kandinsky6Transformer3DModel(config: Kandinsky6VideoAudioConfig, hf_config: dict[str, Any])

Bases: BaseDiT

Native FastVideo implementation of the Kandinsky6 T2VA/IT2VA transformer.

Source code in fastvideo/models/dits/kandinsky6.py
def __init__(self, config: Kandinsky6VideoAudioConfig, hf_config: dict[str, Any]) -> None:
    super().__init__(config=config, hf_config=hf_config)
    arch = config.arch_config
    quant_config = config.quant_config
    self.quant_config = quant_config

    head_dim = sum(arch.axes_dims)
    head_dim_a = sum(arch.axes_dims_a)
    self.in_visual_dim = arch.in_visual_dim
    self.in_audio_dim = arch.in_audio_dim
    self.model_dim = arch.model_dim
    self.patch_size = arch.patch_size
    self.visual_cond = arch.visual_cond
    self.attention_engine = arch.attention_engine
    self.visual_token_type_num_embeddings = arch.visual_token_type_num_embeddings
    self.scale_factor = tuple(float(v) for v in arch.scale_factor)

    visual_embed_dim = (2 * arch.in_visual_dim + 1) if arch.visual_cond else arch.in_visual_dim

    self.visual_embeddings = Kandinsky6VisualEmbeddings(visual_embed_dim, arch.model_dim, arch.patch_size)
    if self.visual_token_type_num_embeddings > 0:
        self.visual_token_type_embeddings = nn.Embedding(self.visual_token_type_num_embeddings, arch.model_dim)
    self.visual_rope_embeddings = Kandinsky6RoPE3D(arch.axes_dims)
    self.out_layer = Kandinsky6OutLayer(arch.model_dim, arch.time_dim, arch.out_visual_dim, arch.patch_size)

    use_nabla = arch.attention_engine == "nabla"

    # The model is always the joint video+audio DiT; a video-only call passes
    # hidden_states_audio=None instead of using a different model shape.

    # Registered before the (much smaller) text towers below: enable_layerwise_offload
    # (fastvideo/hooks/layerwise_offload.py) hooks only the first top-level nn.ModuleList in
    # registration order, and with dit_layerwise_offload=True (the default) that has to be these
    # blocks rather than the 4-entry video_text_transformer_blocks.
    self.visual_transformer_blocks = nn.ModuleList([
        Kandinsky6FusedTransformerDecoderBlock(
            arch.model_dim,
            arch.time_dim,
            arch.ff_dim,
            head_dim,
            arch.model_dim_a,
            arch.time_dim_a,
            arch.ff_dim_a,
            head_dim_a,
            self._supported_attention_backends,
            prefix=f"{config.prefix}.visual_transformer_blocks.{i}",
            use_nabla=use_nabla,
            ca_rope=arch.ca_rope,
            cross_gates=arch.cross_gates,
            fix_modulation=arch.fix_modulation,
            quant_config=quant_config) for i in range(arch.num_visual_blocks)
    ])

    self.audio_embeddings = Kandinsky6TextEmbeddings(arch.in_audio_dim, arch.model_dim_a)
    self.audio_rope_embeddings = Kandinsky6RoPE1D(head_dim_a, freqs_scaling=arch.audio_freqs_scaling)
    self.audio_out_layer = Kandinsky6OutLayerAudio(arch.model_dim_a, arch.time_dim_a, arch.out_audio_dim or arch.in_audio_dim)

    for tower_prefix, model_dim, time_dim, hd in (
        ("video", arch.model_dim, arch.time_dim, head_dim),
        ("audio", arch.model_dim_a, arch.time_dim_a, head_dim_a),
    ):
        setattr(self, f"{tower_prefix}_time_embeddings", Kandinsky6TimeEmbeddings(model_dim, time_dim))
        setattr(self, f"{tower_prefix}_text_embeddings", Kandinsky6TextEmbeddings(arch.in_text_dim, model_dim))
        setattr(self, f"{tower_prefix}_pooled_text_embeddings",
                Kandinsky6TextEmbeddings(arch.in_text_dim2, time_dim))
        setattr(self, f"{tower_prefix}_text_rope_embeddings", Kandinsky6RoPE1D(hd))
        setattr(
            self, f"{tower_prefix}_text_transformer_blocks",
            nn.ModuleList([
                Kandinsky6TransformerEncoderBlock(
                    model_dim,
                    time_dim,
                    arch.ff_dim if tower_prefix == "video" else arch.ff_dim_a,
                    hd,
                    self._supported_attention_backends,
                    prefix=f"{config.prefix}.{tower_prefix}_text_transformer_blocks.{i}",
                    quant_config=quant_config) for i in range(arch.num_text_blocks)
            ]))

    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__()

fastvideo.models.dits.kandinsky6.Kandinsky6TransformerDecoderBlock

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

Bases: Module

Self-attention + text cross-attention + feed-forward block.

Used standalone for plain (non-multimodal) T2V/I2V-parity checkpoints, and as the videoT/audioT sub-block inside Kandinsky6FusedTransformerDecoderBlock for T2VA/IT2VA.

Source code in fastvideo/models/dits/kandinsky6.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 = "",
             use_nabla: bool = False,
             quant_config: QuantizationConfig | None = None):
    super().__init__()
    self.visual_modulation = Kandinsky6Modulation(time_dim, model_dim, 9)

    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",
                                              use_nabla=use_nabla,
                                              quant_config=quant_config)

    self.cross_attention_norm = LayerNormScaleShift(model_dim,
                                                    norm_type="layer",
                                                    eps=1e-5,
                                                    elementwise_affine=False,
                                                    dtype=torch.float32,
                                                    compute_dtype=torch.float32)
    self.cross_attention = Kandinsky6Attention(model_dim,
                                               head_dim,
                                               supported_attention_backends=supported_attention_backends,
                                               prefix=f"{prefix}.cross_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.Kandinsky6TransformerEncoderBlock

Kandinsky6TransformerEncoderBlock(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-only self-attention + feed-forward block (video/audio text towers).

Source code in fastvideo/models/dits/kandinsky6.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.text_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)

Functions: