Skip to content

flux

Classes

fastvideo.models.dits.flux.FluxJointAttention

FluxJointAttention(dim: int, num_attention_heads: int, attention_head_dim: int, supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None)

Bases: Module

Joint attention: text tokens precede image tokens (Diffusers order).

Source code in fastvideo/models/dits/flux.py
def __init__(
    self,
    dim: int,
    num_attention_heads: int,
    attention_head_dim: int,
    supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
) -> None:
    super().__init__()
    self.heads = num_attention_heads
    self.head_dim = attention_head_dim
    self.inner_dim = num_attention_heads * attention_head_dim

    self.norm_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
    self.norm_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
    self.norm_added_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
    self.norm_added_k = nn.RMSNorm(attention_head_dim, eps=1e-6)

    self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
    self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
    self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
    self.add_q_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
    self.add_k_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
    self.add_v_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)

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

    self.attn = DistributedAttention(
        num_heads=num_attention_heads,
        head_size=attention_head_dim,
        causal=False,
        supported_attention_backends=supported_attention_backends,
    )

fastvideo.models.dits.flux.FluxPosEmbed

FluxPosEmbed(theta: int, axes_dim: list[int])

Bases: Module

1D RoPE axes concatenated per Diffusers FluxPosEmbed.

Source code in fastvideo/models/dits/flux.py
def __init__(self, theta: int, axes_dim: list[int]) -> None:
    super().__init__()
    self.theta = theta
    self.axes_dim = axes_dim

fastvideo.models.dits.flux.FluxSingleStreamAttention

FluxSingleStreamAttention(dim: int, num_attention_heads: int, attention_head_dim: int, supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None)

Bases: Module

Self-attention on concatenated text+image sequence (single blocks).

Source code in fastvideo/models/dits/flux.py
def __init__(
    self,
    dim: int,
    num_attention_heads: int,
    attention_head_dim: int,
    supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
) -> None:
    super().__init__()
    self.heads = num_attention_heads
    self.head_dim = attention_head_dim
    self.inner_dim = num_attention_heads * attention_head_dim

    self.norm_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
    self.norm_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
    self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
    self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
    self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
    self.attn = DistributedAttention(
        num_heads=num_attention_heads,
        head_size=attention_head_dim,
        causal=False,
        supported_attention_backends=supported_attention_backends,
    )

fastvideo.models.dits.flux.FluxTransformer2DModel

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

Bases: BaseDiT

FastVideo FLUX transformer; load Diffusers FLUX safetensors 1:1.

Source code in fastvideo/models/dits/flux.py
def __init__(self, config: DiTConfig, hf_config: dict[str, Any], **kwargs) -> None:
    del kwargs
    super().__init__(config=config, hf_config=hf_config)
    self.fastvideo_config = config
    self.hf_config = hf_config
    arch = config.arch_config

    out_ch = arch.out_channels
    self.out_channels = out_ch if out_ch is not None else arch.in_channels
    self.inner_dim = arch.num_attention_heads * arch.attention_head_dim
    self.hidden_size = self.inner_dim
    self.num_attention_heads = arch.num_attention_heads
    self.num_channels_latents = arch.in_channels

    axes_list = list(arch.axes_dims_rope)
    self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_list)
    if arch.guidance_embeds:
        self.time_text_embed = FluxCombinedTimestepGuidanceTextProjEmbeddings(
            embedding_dim=self.inner_dim,
            pooled_projection_dim=arch.pooled_projection_dim,
        )
    else:
        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, self.inner_dim)
    self.x_embedder = ReplicatedLinear(arch.in_channels, self.inner_dim)

    self.transformer_blocks = nn.ModuleList(
        [
            FluxTransformerBlock(
                dim=self.inner_dim,
                num_attention_heads=arch.num_attention_heads,
                attention_head_dim=arch.attention_head_dim,
                supported_attention_backends=self._supported_attention_backends,
            )
            for _ in range(arch.num_layers)
        ]
    )
    self.single_transformer_blocks = nn.ModuleList(
        [
            FluxSingleTransformerBlock(
                dim=self.inner_dim,
                num_attention_heads=arch.num_attention_heads,
                attention_head_dim=arch.attention_head_dim,
                supported_attention_backends=self._supported_attention_backends,
            )
            for _ in range(arch.num_single_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: