Skip to content

kandinsky5

Classes

fastvideo.models.dits.kandinsky5.Kandinsky5Transformer3DModel

Kandinsky5Transformer3DModel(config: Kandinsky5VideoConfig, hf_config: dict[str, Any])

Bases: BaseDiT

Native FastVideo implementation of Kandinsky5 Transformer.

Source code in fastvideo/models/dits/kandinsky5.py
def __init__(self, config: Kandinsky5VideoConfig,
             hf_config: dict[str, Any]) -> None:
    super().__init__(config=config, hf_config=hf_config)
    arch = config.arch_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.visual_cond = arch.visual_cond
    self.attention_type = arch.attention_type

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

    self.time_embeddings = Kandinsky5TimeEmbeddings(
        arch.model_dim, arch.time_dim)
    self.text_embeddings = Kandinsky5TextEmbeddings(
        arch.in_text_dim, arch.model_dim)
    self.pooled_text_embeddings = Kandinsky5TextEmbeddings(
        arch.in_text_dim2, arch.time_dim)
    self.visual_embeddings = Kandinsky5VisualEmbeddings(
        visual_embed_dim, arch.model_dim, arch.patch_size)

    self.text_rope_embeddings = Kandinsky5RoPE1D(head_dim)
    self.visual_rope_embeddings = Kandinsky5RoPE3D(arch.axes_dims)

    self.text_transformer_blocks = nn.ModuleList([
        Kandinsky5TransformerEncoderBlock(arch.model_dim, arch.time_dim,
                                          arch.ff_dim,
                                          head_dim,
                                          self._supported_attention_backends,
                                          prefix=f"{config.prefix}.text_transformer_blocks.{i}")
        for i in range(arch.num_text_blocks)
    ])
    self.visual_transformer_blocks = nn.ModuleList([
        Kandinsky5TransformerDecoderBlock(arch.model_dim, arch.time_dim,
                                          arch.ff_dim,
                                          head_dim,
                                          self._supported_attention_backends,
                                          prefix=f"{config.prefix}.visual_transformer_blocks.{i}",
                                          use_nabla=arch.attention_type == "nabla")
        for i in range(arch.num_visual_blocks)
    ])

    self.out_layer = Kandinsky5OutLayer(arch.model_dim, arch.time_dim,
                                        arch.out_visual_dim,
                                        arch.patch_size)
    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__()

Functions: