def __init__(self, config: Kandinsky5VideoConfig, 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)
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}",
quant_config=quant_config) 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",
quant_config=quant_config) 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__()