← back to Exo

.typings/mflux/models/flux/model/flux_transformer/transformer.pyi

59 lines

"""
This type stub file was generated by pyright.
"""

import mlx.core as mx
from mlx import nn
from mflux.models.common.config.config import Config
from mflux.models.common.config.model_config import ModelConfig
from mflux.models.flux.model.flux_transformer.embed_nd import EmbedND
from mflux.models.flux.model.flux_transformer.time_text_embed import TimeTextEmbed
from mflux.models.flux.model.flux_transformer.joint_transformer_block import (
    JointTransformerBlock,
)
from mflux.models.flux.model.flux_transformer.single_transformer_block import (
    SingleTransformerBlock,
)

class Transformer(nn.Module):
    transformer_blocks: list[JointTransformerBlock]
    single_transformer_blocks: list[SingleTransformerBlock]
    x_embedder: nn.Linear
    pos_embed: EmbedND
    time_text_embed: TimeTextEmbed
    norm_out: nn.LayerNorm
    proj_out: nn.Linear
    context_embedder: nn.Linear

    def __init__(
        self,
        model_config: ModelConfig,
        num_transformer_blocks: int = ...,
        num_single_transformer_blocks: int = ...,
    ) -> None: ...
    def __call__(
        self,
        t: int,
        config: Config,
        hidden_states: mx.array,
        prompt_embeds: mx.array,
        pooled_prompt_embeds: mx.array,
        controlnet_block_samples: list[mx.array] | None = ...,
        controlnet_single_block_samples: list[mx.array] | None = ...,
        kontext_image_ids: mx.array | None = ...,
    ) -> mx.array: ...
    @staticmethod
    def compute_rotary_embeddings(
        prompt_embeds: mx.array,
        pos_embed: EmbedND,
        config: Config,
        kontext_image_ids: mx.array | None = ...,
    ) -> mx.array: ...
    @staticmethod
    def compute_text_embeddings(
        t: int,
        pooled_prompt_embeds: mx.array,
        time_text_embed: TimeTextEmbed,
        config: Config,
    ) -> mx.array: ...