← back to Exo
.typings/mflux/models/qwen/model/qwen_transformer/qwen_transformer.pyi
53 lines
"""
This type stub file was generated by pyright.
"""
import mlx.core as mx
from mlx import nn
from typing import Any
from mflux.models.common.config.config import Config
from mflux.models.qwen.model.qwen_transformer.qwen_transformer_block import (
QwenTransformerBlock,
)
class QwenTransformer(nn.Module):
transformer_blocks: list[QwenTransformerBlock]
inner_dim: int
img_in: nn.Linear
txt_in: nn.Linear
txt_norm: nn.RMSNorm
time_text_embed: Any
pos_embed: Any
norm_out: nn.Module
proj_out: nn.Linear
def __init__(
self,
in_channels: int = ...,
out_channels: int = ...,
num_layers: int = ...,
attention_head_dim: int = ...,
num_attention_heads: int = ...,
joint_attention_dim: int = ...,
patch_size: int = ...,
) -> None: ...
def __call__(
self,
t: int,
config: Config,
hidden_states: mx.array,
encoder_hidden_states: mx.array,
encoder_hidden_states_mask: mx.array,
qwen_image_ids: mx.array | None = ...,
cond_image_grid: tuple[int, int, int] | None = ...,
) -> mx.array: ...
@staticmethod
def _compute_timestep(t: int | float, config: Config) -> mx.array: ...
@staticmethod
def _compute_rotary_embeddings(
encoder_hidden_states_mask: mx.array,
pos_embed: Any,
config: Config,
cond_image_grid: tuple[int, int, int] | list[tuple[int, int, int]] | None = ...,
) -> tuple[mx.array, mx.array]: ...