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