Skip to content

qwen3

Qwen3 causal LM text encoder for FastVideo diffusion models (e.g. Flux2 Klein).

Classes

fastvideo.models.encoders.qwen3.Qwen3Attention

Qwen3Attention(config: Qwen3TextConfig, hidden_size: int, num_heads: int, num_kv_heads: int, rope_theta: float = 1000000.0, rope_scaling: dict[str, Any] | None = None, max_position_embeddings: int = 40960, quant_config: QuantizationConfig | None = None, bias: bool = False, prefix: str = '')

Bases: Module

Qwen3 attention with QK-Norm and tensor parallelism.

Key difference from LLaMA: RMSNorm is applied to Q and K before attention.

Source code in fastvideo/models/encoders/qwen3.py
def __init__(
    self,
    config: Qwen3TextConfig,
    hidden_size: int,
    num_heads: int,
    num_kv_heads: int,
    rope_theta: float = 1000000.0,
    rope_scaling: dict[str, Any] | None = None,
    max_position_embeddings: int = 40960,
    quant_config: QuantizationConfig | None = None,
    bias: bool = False,
    prefix: str = "",
) -> None:
    super().__init__()
    self.hidden_size = hidden_size
    tp_size = get_tp_world_size()
    self.total_num_heads = num_heads
    assert self.total_num_heads % tp_size == 0
    self.num_heads = self.total_num_heads // tp_size
    self.total_num_kv_heads = num_kv_heads
    if self.total_num_kv_heads >= tp_size:
        assert self.total_num_kv_heads % tp_size == 0
    else:
        assert tp_size % self.total_num_kv_heads == 0
    self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)

    self.head_dim = getattr(
        config, "head_dim", self.hidden_size // self.total_num_heads
    )
    self.rotary_dim = self.head_dim
    self.q_size = self.num_heads * self.head_dim
    self.kv_size = self.num_kv_heads * self.head_dim
    self.scaling = self.head_dim**-0.5
    self.rope_theta = rope_theta
    self.max_position_embeddings = max_position_embeddings

    self.qkv_proj = QKVParallelLinear(
        hidden_size=hidden_size,
        head_size=self.head_dim,
        total_num_heads=self.total_num_heads,
        total_num_kv_heads=self.total_num_kv_heads,
        bias=bias,
        quant_config=quant_config,
        prefix=f"{prefix}.qkv_proj",
    )
    self.o_proj = RowParallelLinear(
        input_size=self.total_num_heads * self.head_dim,
        output_size=hidden_size,
        bias=bias,
        quant_config=quant_config,
        prefix=f"{prefix}.o_proj",
    )

    rms_norm_eps = getattr(config, "rms_norm_eps", 1e-6)
    self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)
    self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)

    self.rotary_emb = get_rope(
        self.head_dim,
        rotary_dim=self.rotary_dim,
        max_position=max_position_embeddings,
        base=int(rope_theta),
        rope_scaling=rope_scaling,
        is_neox_style=True,
    )

    self.attn = LocalAttention(
        self.num_heads,
        self.head_dim,
        self.num_kv_heads,
        softmax_scale=self.scaling,
        causal=True,
        supported_attention_backends=config._supported_attention_backends,
    )

fastvideo.models.encoders.qwen3.Qwen3DecoderLayer

Qwen3DecoderLayer(config: Qwen3TextConfig, quant_config: QuantizationConfig | None = None, prefix: str = '')

Bases: Module

Qwen3 transformer decoder layer.

Source code in fastvideo/models/encoders/qwen3.py
def __init__(
    self,
    config: Qwen3TextConfig,
    quant_config: QuantizationConfig | None = None,
    prefix: str = "",
) -> None:
    super().__init__()
    self.hidden_size = config.hidden_size
    rope_theta = getattr(config, "rope_theta", 1000000.0)
    rope_scaling = getattr(config, "rope_scaling", None)
    max_position_embeddings = getattr(config, "max_position_embeddings", 40960)
    attention_bias = getattr(config, "attention_bias", False)

    self.self_attn = Qwen3Attention(
        config=config,
        hidden_size=self.hidden_size,
        num_heads=config.num_attention_heads,
        num_kv_heads=getattr(
            config, "num_key_value_heads", config.num_attention_heads
        ),
        rope_theta=rope_theta,
        rope_scaling=rope_scaling,
        max_position_embeddings=max_position_embeddings,
        quant_config=quant_config,
        bias=attention_bias,
        prefix=f"{prefix}.self_attn",
    )
    self.mlp = Qwen3MLP(
        hidden_size=self.hidden_size,
        intermediate_size=config.intermediate_size,
        hidden_act=config.hidden_act,
        quant_config=quant_config,
        bias=getattr(config, "mlp_bias", False),
        prefix=f"{prefix}.mlp",
    )
    self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
    self.post_attention_layernorm = RMSNorm(
        config.hidden_size, eps=config.rms_norm_eps
    )

fastvideo.models.encoders.qwen3.Qwen3ForCausalLM

Qwen3ForCausalLM(config: Qwen3TextConfig)

Bases: TextEncoder

Qwen3 causal language model for text encoding in diffusion models (e.g. Flux2 Klein).

Features: - Tensor parallelism support - FlashAttention/SDPA support via LocalAttention - QK-Norm for better training stability - output_hidden_states for Klein (layers 9, 18, 27)

Source code in fastvideo/models/encoders/qwen3.py
def __init__(self, config: Qwen3TextConfig) -> None:
    super().__init__(config)

    self.config = config
    self.quant_config = getattr(config, "quant_config", None)

    if getattr(config, "lora_config", None) is not None:
        max_loras = getattr(config.lora_config, "max_loras", 1)
        lora_vocab_size = getattr(config.lora_config, "lora_extra_vocab_size", 1)
        lora_vocab = lora_vocab_size * max_loras
    else:
        lora_vocab = 0
    self.vocab_size = config.vocab_size + lora_vocab
    self.org_vocab_size = config.vocab_size

    self.embed_tokens = VocabParallelEmbedding(
        self.vocab_size,
        config.hidden_size,
        org_num_embeddings=config.vocab_size,
        quant_config=self.quant_config,
    )

    self.layers = nn.ModuleList(
        [
            Qwen3DecoderLayer(
                config=config,
                quant_config=self.quant_config,
                prefix=f"{config.prefix}.layers.{i}",
            )
            for i in range(config.num_hidden_layers)
        ]
    )

    self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)

fastvideo.models.encoders.qwen3.Qwen3MLP

Qwen3MLP(hidden_size: int, intermediate_size: int, hidden_act: str, quant_config: QuantizationConfig | None = None, bias: bool = False, prefix: str = '')

Bases: Module

Qwen3 MLP with SwiGLU activation and tensor parallelism.

Source code in fastvideo/models/encoders/qwen3.py
def __init__(
    self,
    hidden_size: int,
    intermediate_size: int,
    hidden_act: str,
    quant_config: QuantizationConfig | None = None,
    bias: bool = False,
    prefix: str = "",
) -> None:
    super().__init__()
    self.gate_up_proj = MergedColumnParallelLinear(
        input_size=hidden_size,
        output_sizes=[intermediate_size] * 2,
        bias=bias,
        quant_config=quant_config,
        prefix=f"{prefix}.gate_up_proj",
    )
    self.down_proj = RowParallelLinear(
        input_size=intermediate_size,
        output_size=hidden_size,
        bias=bias,
        quant_config=quant_config,
        prefix=f"{prefix}.down_proj",
    )
    if hidden_act != "silu":
        raise ValueError(
            f"Unsupported activation: {hidden_act}. Only silu is supported."
        )
    self.act_fn = SiluAndMul()

Functions: