← back to Exo

.typings/mflux/models/flux/variants/controlnet/transformer_controlnet.pyi

26 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

class TransformerControlnet(nn.Module):
    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_condition: mx.array,
    ) -> tuple[list[mx.array], list[mx.array]]: ...