← back to Exo
.typings/mflux/models/qwen/model/qwen_transformer/qwen_transformer_block.pyi
37 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.qwen.model.qwen_transformer.qwen_attention import QwenAttention
class QwenTransformerBlock(nn.Module):
attn: QwenAttention
img_mod_linear: nn.Linear
img_mod_silu: nn.SiLU
txt_mod_linear: nn.Linear
txt_mod_silu: nn.SiLU
img_norm1: nn.RMSNorm
txt_norm1: nn.RMSNorm
img_norm2: nn.RMSNorm
txt_norm2: nn.RMSNorm
img_ff: Any
txt_ff: Any
def __init__(
self, dim: int = ..., num_heads: int = ..., head_dim: int = ...
) -> None: ...
def __call__(
self,
hidden_states: mx.array,
encoder_hidden_states: mx.array,
encoder_hidden_states_mask: mx.array | None,
text_embeddings: mx.array,
image_rotary_emb: tuple[mx.array, mx.array],
block_idx: int | None = ...,
) -> tuple[mx.array, mx.array]: ...
@staticmethod
def _modulate(x: mx.array, mod_params: mx.array) -> tuple[mx.array, mx.array]: ...