Skip to content

sd3

Classes

fastvideo.models.dits.sd3.SD3Attention

SD3Attention(query_dim: int, heads: int, dim_head: int, out_dim: int, added_kv_proj_dim: int | None = None, context_pre_only: bool | None = None, qk_norm: str | None = None, eps: float = 1e-06, supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None)

Bases: Module

Joint self-attention used in SD3 blocks.

Source code in fastvideo/models/dits/sd3.py
def __init__(
    self,
    query_dim: int,
    heads: int,
    dim_head: int,
    out_dim: int,
    added_kv_proj_dim: int | None = None,
    context_pre_only: bool | None = None,
    qk_norm: str | None = None,
    eps: float = 1e-6,
    supported_attention_backends: tuple[AttentionBackendEnum,
                                        ...] | None = None,
) -> None:
    super().__init__()

    self.heads = heads
    self.head_dim = dim_head
    self.inner_dim = out_dim
    self.context_pre_only = context_pre_only

    self.to_q = ReplicatedLinear(query_dim, self.inner_dim, bias=True)
    self.to_k = ReplicatedLinear(query_dim, self.inner_dim, bias=True)
    self.to_v = ReplicatedLinear(query_dim, self.inner_dim, bias=True)

    self.norm_q, self.norm_k = _build_qk_norm(qk_norm, dim_head, eps)

    self.add_q_proj: ReplicatedLinear | None = None
    self.add_k_proj: ReplicatedLinear | None = None
    self.add_v_proj: ReplicatedLinear | None = None
    self.norm_added_q: nn.Module | None = None
    self.norm_added_k: nn.Module | None = None
    if added_kv_proj_dim is not None:
        self.add_q_proj = ReplicatedLinear(
            added_kv_proj_dim,
            self.inner_dim,
            bias=True,
        )
        self.add_k_proj = ReplicatedLinear(
            added_kv_proj_dim,
            self.inner_dim,
            bias=True,
        )
        self.add_v_proj = ReplicatedLinear(
            added_kv_proj_dim,
            self.inner_dim,
            bias=True,
        )
        self.norm_added_q, self.norm_added_k = _build_qk_norm(
            qk_norm,
            dim_head,
            eps,
        )

    self.to_out = nn.ModuleList([
        ReplicatedLinear(self.inner_dim, out_dim, bias=True),
        nn.Dropout(0.0),
    ])

    if context_pre_only is not None and not context_pre_only:
        self.to_add_out = ReplicatedLinear(self.inner_dim, out_dim,
                                           bias=True)
    else:
        self.to_add_out = None

    self.attn = DistributedAttention(
        num_heads=heads,
        head_size=dim_head,
        causal=False,
        supported_attention_backends=supported_attention_backends,
    )

fastvideo.models.dits.sd3.SD3PatchEmbed

SD3PatchEmbed(height: int = 224, width: int = 224, patch_size: int = 16, in_channels: int = 3, embed_dim: int = 768, layer_norm: bool = False, flatten: bool = True, bias: bool = True, interpolation_scale: float = 1.0, pos_embed_type: str | None = 'sincos', pos_embed_max_size: int | None = None)

Bases: Module

2D patch embedding with SD3 positional embedding cropping behavior.

Source code in fastvideo/models/dits/sd3.py
def __init__(
    self,
    height: int = 224,
    width: int = 224,
    patch_size: int = 16,
    in_channels: int = 3,
    embed_dim: int = 768,
    layer_norm: bool = False,
    flatten: bool = True,
    bias: bool = True,
    interpolation_scale: float = 1.0,
    pos_embed_type: str | None = "sincos",
    pos_embed_max_size: int | None = None,
) -> None:
    super().__init__()

    num_patches = (height // patch_size) * (width // patch_size)
    self.flatten = flatten
    self.layer_norm = layer_norm
    self.pos_embed_max_size = pos_embed_max_size

    self.proj = nn.Conv2d(
        in_channels,
        embed_dim,
        kernel_size=(patch_size, patch_size),
        stride=patch_size,
        bias=bias,
    )
    if layer_norm:
        self.norm = nn.LayerNorm(embed_dim, elementwise_affine=False,
                                 eps=1e-6)
    else:
        self.norm = None

    self.patch_size = patch_size
    self.height = height // patch_size
    self.width = width // patch_size
    self.base_size = height // patch_size
    self.interpolation_scale = interpolation_scale

    grid_size = pos_embed_max_size or int(num_patches ** 0.5)

    if pos_embed_type is None:
        self.pos_embed = None
    elif pos_embed_type == "sincos":
        pos_embed = _get_2d_sincos_pos_embed(
            embed_dim,
            grid_size,
            base_size=self.base_size,
            interpolation_scale=self.interpolation_scale,
        )
        persistent = True if pos_embed_max_size else False
        self.register_buffer(
            "pos_embed",
            pos_embed.float().unsqueeze(0),
            persistent=persistent,
        )
    else:
        raise ValueError(f"Unsupported pos_embed_type: {pos_embed_type}")

fastvideo.models.dits.sd3.SD3Transformer2DModel

SD3Transformer2DModel(config: DiTConfig, hf_config: dict[str, Any], **kwargs)

Bases: BaseDiT

FastVideo-native SD3 Transformer2DModel.

Source code in fastvideo/models/dits/sd3.py
def __init__(self, config: DiTConfig, hf_config: dict[str, Any], **kwargs):
    del kwargs
    super().__init__(config=config, hf_config=hf_config)

    self.fastvideo_config = config
    self.hf_config = hf_config

    arch = config.arch_config

    self.out_channels = arch.out_channels
    self.inner_dim = arch.num_attention_heads * arch.attention_head_dim
    self.patch_size = arch.patch_size
    self.num_layers = arch.num_layers

    self.hidden_size = self.inner_dim
    self.num_attention_heads = arch.num_attention_heads
    self.num_channels_latents = arch.in_channels

    self.pos_embed = SD3PatchEmbed(
        height=arch.sample_size,
        width=arch.sample_size,
        patch_size=arch.patch_size,
        in_channels=arch.in_channels,
        embed_dim=self.inner_dim,
        pos_embed_max_size=arch.pos_embed_max_size,
    )
    self.time_text_embed = CombinedTimestepTextProjEmbeddings(
        embedding_dim=self.inner_dim,
        pooled_projection_dim=arch.pooled_projection_dim,
    )
    self.context_embedder = ReplicatedLinear(arch.joint_attention_dim,
                                             arch.caption_projection_dim)

    dual_layers = getattr(arch, "dual_attention_layers", ())
    dual_layers = tuple(dual_layers) if isinstance(dual_layers,
                                                   list) else dual_layers

    self.transformer_blocks = nn.ModuleList([
        SD3JointTransformerBlock(
            dim=self.inner_dim,
            num_attention_heads=arch.num_attention_heads,
            attention_head_dim=arch.attention_head_dim,
            context_pre_only=i == arch.num_layers - 1,
            qk_norm=arch.qk_norm,
            use_dual_attention=i in dual_layers,
            supported_attention_backends=self._supported_attention_backends,
        ) for i in range(arch.num_layers)
    ])

    self.norm_out = SD3AdaLayerNormContinuous(
        self.inner_dim,
        self.inner_dim,
        elementwise_affine=False,
        eps=1e-6,
        bias=True,
        norm_type="layer_norm",
    )
    self.proj_out = ReplicatedLinear(
        self.inner_dim,
        arch.patch_size * arch.patch_size * self.out_channels,
        bias=True,
    )

    self.gradient_checkpointing = False
    self.__post_init__()

Functions: