Skip to content

vllm.models.glm5next.common.model

Classes:

_LOGIT_SCALE = 1.0 module-attribute

Output logit scale. A GLM-5.3-Flash trained value that neither the checkpoint nor Glm5NextTextConfig carries.

_MHC_POST_MULT_VALUE = 2.0 module-attribute

mHC post-multiplier. A GLM-5.3-Flash trained value that neither the checkpoint nor Glm5NextTextConfig carries.

_MHC_TAU = 0.05 module-attribute

mHC routing temperature. A GLM-5.3-Flash trained value that neither the checkpoint nor Glm5NextTextConfig carries.

_VISION_RMS_NORM_EPS = 1e-06 module-attribute

Vision tower RMSNorm epsilon.

GLM-5.3-Flash checkpoints ship vision_config.rms_norm_eps = 1e-5, but the vision tower was trained with 1e-6. Serving with 1e-5 drifts the RMSNorm and produces repetitive/degraded image descriptions, so force the trained value regardless of the checkpoint field.

Glm5NextModel

Bases: Module, EagleModelMixin

Source code in vllm/models/glm5next/common/model.py
class Glm5NextModel(nn.Module, EagleModelMixin):
    def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
        super().__init__()

        config = vllm_config.model_config.hf_text_config
        _validate_supported_config(config)
        self.config = config

        self.vocab_size = config.vocab_size
        self.device = current_platform.device_type

        self.is_v32 = config.index_topk is not None
        if self.is_v32:
            topk_tokens = config.index_topk
            assert topk_tokens is not None
            # Reserve room for the incomplete pool tail.
            kpool = config.index_kpool
            assert kpool is not None
            buffer_width = topk_tokens + (kpool - 1 if kpool > 1 else 0)
            # Sparse MLA tiles top-k in 128 columns; padded slots remain masked.
            sparse_topk_block_n = 128
            buffer_width = (
                (buffer_width + sparse_topk_block_n - 1) // sparse_topk_block_n
            ) * sparse_topk_block_n
            topk_indices_buffer = torch.empty(
                vllm_config.scheduler_config.max_num_batched_tokens,
                buffer_width,
                dtype=torch.int32,
                device=self.device,
            )
        else:
            # Full-MLA config (no kpool sparse indexer): no topk buffer.
            topk_indices_buffer = None

        if get_pp_group().is_first_rank:
            self.embed_tokens = VocabParallelEmbedding(
                config.vocab_size,
                config.hidden_size,
                prefix=f"{prefix}.embed_tokens",
            )
        else:
            self.embed_tokens = PPMissingLayer()

        def get_layer(prefix: str):
            layer_idx = int(prefix.rsplit(".", 1)[1])
            return Glm5NextDecoderLayer(
                vllm_config=vllm_config,
                config=config,
                layer_idx=layer_idx,
                prefix=prefix,
                topk_indices_buffer=topk_indices_buffer,
            )

        self.start_layer, self.end_layer, self.layers = make_layers(
            config.num_hidden_layers,
            get_layer,
            prefix=f"{prefix}.layers",
        )
        self._aux_post_op = MHCPostOp()
        # The active slice is fixed after construction; cache it so forward
        # doesn't rebuild the slice (a fresh list) every step.
        self._active_layers = self.layers[self.start_layer : self.end_layer]
        self.is_fused_shared_expert_enabled = is_model_fused_shared_expert_compatible(
            self.layers, Glm5NextMoE, "mlp"
        )

        if get_pp_group().is_last_rank:
            self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        else:
            self.norm = PPMissingLayer()

        self.is_sequence_parallel = (
            vllm_config.parallel_config.use_sequence_parallel_moe
        )

        world_size = get_tensor_model_parallel_world_size()
        assert config.num_attention_heads % world_size == 0, (
            "num_attention_heads must be divisible by world_size"
        )

    def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
        return self.embed_tokens(input_ids)

    def _aux_hidden_state(
        self,
        hidden_states: torch.Tensor,
        residual: torch.Tensor | None,
        post: torch.Tensor | None,
        comb: torch.Tensor | None,
    ) -> torch.Tensor:
        """Completed residual stream entering a layer, as one hidden vector.

        mHC layers defer their final ``hc_post`` into the next layer's fused
        pre-op, so ``hidden_states`` holds the raw layer output while the
        widened stream is completed here and contracted back to
        ``hidden_size``. Non-mHC layers already return the summed stream.
        """
        if post is None:
            return hidden_states
        assert residual is not None and comb is not None
        completed = self._aux_post_op(hidden_states, residual, post, comb)
        return hc_contract(completed, self.config.mhc_num_residual_streams)

    def forward(
        self,
        input_ids: torch.Tensor | None,
        positions: torch.Tensor,
        intermediate_tensors: IntermediateTensors | None,
        inputs_embeds: torch.Tensor | None = None,
        **kwargs,
    ) -> torch.Tensor:
        if get_pp_group().is_first_rank:
            if inputs_embeds is not None:
                hidden_states = inputs_embeds
            else:
                hidden_states = self.embed_input_ids(input_ids)
            residual = None
            post = None
            comb = None
        else:
            assert intermediate_tensors is not None
            hidden_states = intermediate_tensors["hidden_states"]
            residual = intermediate_tensors["residual"]
            # post/comb (deferred mHC hc_post state) are not propagated across
            # PP ranks; the receiving rank's first mHC layer uses standalone pre.
            post = None
            comb = None

        full_num_tokens = positions.shape[0]
        if self.is_sequence_parallel:
            hidden_states = sp_shard(hidden_states)

        aux_hidden_states: list[torch.Tensor] = []
        for idx, layer in enumerate(self._active_layers, start=self.start_layer):
            if idx in self.aux_hidden_state_layers:
                aux_hidden_state = self._aux_hidden_state(
                    hidden_states, residual, post, comb
                )
                if self.is_sequence_parallel:
                    aux_hidden_state = sp_all_gather(aux_hidden_state)[:full_num_tokens]
                aux_hidden_states.append(aux_hidden_state)
            hidden_states, residual, post, comb = layer(
                positions, hidden_states, residual, post, comb
            )

        if not get_pp_group().is_last_rank:
            # PP is gated off for GLM-5.3-Flash (no make_empty_intermediate_tensors),
            # so this branch is not exercised. post/comb are the deferred
            # hc_post state of this rank's last mHC layer; a future PP path
            # would need to propagate them, but for now they are dropped (the
            # receiving rank's first layer would fall back to standalone pre).
            return IntermediateTensors(
                {"hidden_states": hidden_states, "residual": residual}
            )

        if self.end_layer in self.aux_hidden_state_layers:
            final_aux = self._aux_hidden_state(hidden_states, residual, post, comb)
            if self.is_sequence_parallel:
                final_aux = sp_all_gather(final_aux)[:full_num_tokens]
            aux_hidden_states.append(final_aux)

        if self.is_sequence_parallel:
            hidden_states = sp_all_gather(hidden_states)[:full_num_tokens]

        hidden_states = self.norm(hidden_states)
        if aux_hidden_states:
            return hidden_states, aux_hidden_states
        return hidden_states

    def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
        stacked_params_mapping = [
            # (param_name, shard_name, shard_id)
            (".gate_up_proj", ".gate_proj", 0),
            (".gate_up_proj", ".up_proj", 1),
            # MLA: fuse q_a_proj and kv_a_proj_with_mqa
            (".fused_qkv_a_proj", ".q_a_proj", 0),
            (".fused_qkv_a_proj", ".kv_a_proj_with_mqa", 1),
            # Indexer: fuse wk and weights_proj
            (".wk_weights_proj", ".wk", 0),
            (".wk_weights_proj", ".weights_proj", 1),
            # KDA: merge q, k, v, b, f_a, g_a projections into one GEMM
            (".in_proj_qkvbfg_a", ".q_proj", 0),
            (".in_proj_qkvbfg_a", ".k_proj", 1),
            (".in_proj_qkvbfg_a", ".v_proj", 2),
            (".in_proj_qkvbfg_a", ".b_proj", 3),
            (".in_proj_qkvbfg_a", ".f_a_proj", 4),
            (".in_proj_qkvbfg_a", ".g_a_proj", 5),
        ]
        if _is_moe(self.config):
            # Params for weights, fp8 weight scales, fp8 activation scales
            # (param_name, weight_name, expert_id, shard_id)
            # EPLB: the mapping enumerates physical experts, so it must cover
            # the redundant replicas or their slots are never loaded.
            num_redundant_experts = next(
                (
                    layer.mlp.n_redundant_experts
                    for layer in self.layers
                    if isinstance(layer, Glm5NextDecoderLayer)
                    and isinstance(layer.mlp, Glm5NextMoE)
                ),
                0,
            )
            num_fused_shared = _num_fused_shared_experts(
                self.config.n_shared_experts, self.is_fused_shared_expert_enabled
            )
            expert_params_mapping = fused_moe_make_expert_params_mapping(
                self,
                ckpt_gate_proj_name="gate_proj",
                ckpt_down_proj_name="down_proj",
                ckpt_up_proj_name="up_proj",
                num_experts=self.config.n_routed_experts + num_fused_shared,
                num_redundant_experts=num_redundant_experts,
            )
        else:
            expert_params_mapping = []
        params_dict = dict(self.named_parameters())
        loaded_params: set[str] = set()

        # GLM-5.3-Flash NoPE checkpoints omit the RoPE rows from
        # ``kv_a_proj_with_mqa``; pad them with zeros for the model shape.
        kv_a_pad_size = 0
        if self.config.mla_use_nope and self.config.qk_rope_head_dim > 0:
            kv_a_pad_size = self.config.qk_rope_head_dim

        _pending_wk_fp8: dict = {}

        for args in weights:
            name, loaded_weight = args[:2]
            kwargs: dict = args[2] if len(args) > 2 else {}
            if "rotary_emb.inv_freq" in name:
                continue

            spec_layer = get_spec_layer_idx_from_weight_name(self.config, name)
            if spec_layer is not None:
                continue  # skip spec decode layers for main model
            if "rotary_emb.cos_cached" in name or "rotary_emb.sin_cached" in name:
                # Models trained using ColossalAI may include these tensors in
                # the checkpoint. Skip them.
                continue
            if self.is_fused_shared_expert_enabled:
                name = _fused_shared_expert_name(name, self.config.n_routed_experts)

            # Handle FP8 indexer WK: dequantize to BF16 for fusion with
            # weights_proj into wk_weights_proj.
            if _try_load_fp8_indexer_wk(
                name,
                loaded_weight,
                _pending_wk_fp8,
                params_dict,
                loaded_params,
            ):
                continue

            # FP8 checkpoint: dequantize BF16-kept MLA projections
            # (q_a_proj / kv_a_proj_with_mqa / o_proj) to BF16.
            if _try_load_fp8_attn_proj(
                name,
                loaded_weight,
                _pending_wk_fp8,
                params_dict,
                loaded_params,
                kv_a_pad_size,
            ):
                continue

            # Pad kv_a_proj_with_mqa for NoPE models
            if kv_a_pad_size > 0 and ".kv_a_proj_with_mqa." in name:
                pad = torch.zeros(
                    kv_a_pad_size,
                    *loaded_weight.shape[1:],
                    dtype=loaded_weight.dtype,
                    device=loaded_weight.device,
                )
                loaded_weight = torch.cat([loaded_weight, pad], dim=0)

            for param_name, weight_name, shard_id in stacked_params_mapping:
                if weight_name not in name:
                    continue
                # We have mlp.experts[0].gate_proj in the checkpoint.
                # Since we handle the experts below in expert_params_mapping,
                # we need to skip here BEFORE we update the name, otherwise
                # name will be updated to mlp.experts[0].gate_up_proj, which
                # will then be updated below in expert_params_mapping
                # for mlp.experts[0].gate_gate_up_proj, which breaks load.
                if ("mlp.experts." in name) and name not in params_dict:
                    continue
                name_mapped = name.replace(weight_name, param_name)
                # QKV fusion: skip if fused module doesn't exist in model
                if param_name == ".fused_qkv_a_proj" and name_mapped not in params_dict:
                    continue
                name = name_mapped
                # Skip loading extra bias for GPTQ models.
                if name.endswith(".bias") and name not in params_dict:
                    continue
                if is_pp_missing_parameter(name, self):
                    continue
                param = params_dict[name]
                weight_loader = param.weight_loader
                weight_loader(param, loaded_weight, shard_id)
                break
            else:
                is_expert_weight = False
                for (
                    param_name,
                    weight_name,
                    expert_id,
                    expert_shard_id,
                ) in expert_params_mapping:
                    if weight_name not in name:
                        continue
                    # A checkpoint expert may map to several physical replicas
                    # under EPLB; keep `name` intact and try the next entry
                    # when this physical expert is not local to the rank.
                    is_expert_weight = True
                    name_mapped = name.replace(weight_name, param_name)
                    if is_pp_missing_parameter(name_mapped, self):
                        continue
                    param = params_dict[name_mapped]
                    weight_loader = param.weight_loader
                    success = weight_loader(
                        param,
                        loaded_weight,
                        name_mapped,
                        expert_id=expert_id,
                        shard_id=expert_shard_id,
                        return_success=True,
                    )
                    if success:
                        name = name_mapped
                        break
                else:
                    if is_expert_weight:
                        continue
                    # Skip loading extra bias for GPTQ models.
                    if (
                        name.endswith(".bias")
                        and name not in params_dict
                        and not _is_linear_attn(self.config)
                    ):  # noqa: E501
                        continue
                    # Remapping the name of FP8 kv-scale.
                    remapped_name = maybe_remap_kv_scale_name(name, params_dict)
                    if remapped_name is None:
                        continue
                    name = remapped_name
                    if is_pp_missing_parameter(name, self):
                        continue

                    param = params_dict[name]
                    weight_loader = getattr(
                        param, "weight_loader", default_weight_loader
                    )
                    weight_loader(param, loaded_weight, **kwargs)
            loaded_params.add(name)
        return loaded_params

_aux_hidden_state(hidden_states, residual, post, comb)

Completed residual stream entering a layer, as one hidden vector.

mHC layers defer their final hc_post into the next layer's fused pre-op, so hidden_states holds the raw layer output while the widened stream is completed here and contracted back to hidden_size. Non-mHC layers already return the summed stream.

Source code in vllm/models/glm5next/common/model.py
def _aux_hidden_state(
    self,
    hidden_states: torch.Tensor,
    residual: torch.Tensor | None,
    post: torch.Tensor | None,
    comb: torch.Tensor | None,
) -> torch.Tensor:
    """Completed residual stream entering a layer, as one hidden vector.

    mHC layers defer their final ``hc_post`` into the next layer's fused
    pre-op, so ``hidden_states`` holds the raw layer output while the
    widened stream is completed here and contracted back to
    ``hidden_size``. Non-mHC layers already return the summed stream.
    """
    if post is None:
        return hidden_states
    assert residual is not None and comb is not None
    completed = self._aux_post_op(hidden_states, residual, post, comb)
    return hc_contract(completed, self.config.mhc_num_residual_streams)

_dequant_fp8_block(weight_fp8, scale_inv, block_size=128)

Dequantize a block-FP8 (e4m3) weight with per-block scale to BF16.

Unlike scaled_dequantize this tolerates a non-divisible (partial last block) shape by zero-padding to a multiple of block_size before the scale broadcast and trimming back afterwards (e.g. kv_a_proj_with_mqa is 576 rows = 4*128 + 64).

Source code in vllm/models/glm5next/common/model.py
def _dequant_fp8_block(
    weight_fp8: torch.Tensor,
    scale_inv: torch.Tensor,
    block_size: int = 128,
) -> torch.Tensor:
    """Dequantize a block-FP8 (e4m3) weight with per-block scale to BF16.

    Unlike ``scaled_dequantize`` this tolerates a non-divisible (partial last
    block) shape by zero-padding to a multiple of ``block_size`` before the
    scale broadcast and trimming back afterwards (e.g. kv_a_proj_with_mqa is
    576 rows = 4*128 + 64).
    """
    out_dim, in_dim = weight_fp8.shape
    pad_out = (-out_dim) % block_size
    pad_in = (-in_dim) % block_size
    w = weight_fp8
    if pad_out or pad_in:
        w = torch.nn.functional.pad(w, (0, pad_in, 0, pad_out))
    # scale_inv is (ceil(out/block), ceil(in/block)); broadcast to (out, in).
    s = scale_inv.to(torch.float32)
    s_full = s.repeat_interleave(block_size, dim=0).repeat_interleave(block_size, dim=1)
    out = (w.to(torch.float32) * s_full).to(torch.bfloat16)
    return out[:out_dim, :in_dim].contiguous()

_fused_shared_expert_name(name, n_routed_experts)

Point a checkpoint mlp.shared_experts.* tensor at the fused MoE's shared-expert slot, which follows the routed experts; other names are returned unchanged.

Source code in vllm/models/glm5next/common/model.py
def _fused_shared_expert_name(name: str, n_routed_experts: int) -> str:
    """Point a checkpoint ``mlp.shared_experts.*`` tensor at the fused MoE's
    shared-expert slot, which follows the routed experts; other names are
    returned unchanged."""
    return name.replace("mlp.shared_experts.", f"mlp.experts.{n_routed_experts}.", 1)

_fused_shared_experts_tuned(parallel_config)

AITER has fused-MoE configs tuned for the fused shared-expert shape (one more expert and one more top-k slot than the routed MoE) only on gfx950, with every expert on each rank and its weights split by TP4 or TP8. Data, prefill context and expert parallelism change that split, so any other GPU or parallel layout would run untuned fallback kernels.

Source code in vllm/models/glm5next/common/model.py
def _fused_shared_experts_tuned(parallel_config: ParallelConfig) -> bool:
    """AITER has fused-MoE configs tuned for the fused shared-expert shape
    (one more expert and one more top-k slot than the routed MoE) only on
    gfx950, with every expert on each rank and its weights split by TP4 or
    TP8. Data, prefill context and expert parallelism change that split, so
    any other GPU or parallel layout would run untuned fallback kernels."""
    from vllm.platforms.rocm import on_gfx950

    reasons: list[str] = []
    if not on_gfx950():
        reasons.append("the GPU is not gfx950")
    if parallel_config.tensor_parallel_size not in (4, 8):
        reasons.append(
            f"tensor_parallel_size is {parallel_config.tensor_parallel_size}"
        )
    if parallel_config.data_parallel_size != 1:
        reasons.append(f"data_parallel_size is {parallel_config.data_parallel_size}")
    if parallel_config.prefill_context_parallel_size != 1:
        reasons.append(
            "prefill_context_parallel_size is "
            f"{parallel_config.prefill_context_parallel_size}"
        )
    if parallel_config.enable_expert_parallel:
        reasons.append("expert parallelism is enabled")

    if not reasons:
        return True
    logger.warning_once(
        "VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS is ignored for GLM-5.3-Flash: "
        "%s. AITER has tuned configs for its fused shared-expert MoE only on "
        "gfx950 at TP4 and TP8, without data, prefill context or expert "
        "parallelism. Running the shared experts as a separate MLP.",
        "; ".join(reasons),
    )
    return False

_num_fused_shared_experts(n_shared_experts, enabled)

Expert slots the fused MoE appends for the shared expert; must match the num_fused_shared_experts that FusedMoE allocates.

Source code in vllm/models/glm5next/common/model.py
def _num_fused_shared_experts(n_shared_experts: int | None, enabled: bool) -> int:
    """Expert slots the fused MoE appends for the shared expert; must match the
    ``num_fused_shared_experts`` that ``FusedMoE`` allocates."""
    if not enabled or n_shared_experts is None:
        return 0
    if n_shared_experts > 1:
        raise NotImplementedError(
            "Fused shared-expert loading supports only 1 shared expert per "
            f"layer, but config.n_shared_experts is {n_shared_experts}. Set "
            "VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS=0 to run the shared "
            "experts as a separate MLP."
        )
    return n_shared_experts

_try_load_fp8_attn_proj(name, tensor, buf, params_dict, loaded_params, kv_a_pad_size)

Dequantize FP8 q_a_proj / kv_a_proj_with_mqa / o_proj to BF16 on load.

The FP8 checkpoint stores these as block-FP8 (weight + weight_scale_inv), but the model holds them in BF16 (fused_qkv_a_proj is always BF16 via DeepSeekV2FusedQkvAProjLinear; o_proj is excluded by modules_to_not_convert). When the model target is BF16 (no weight_scale_inv param) we dequantize; otherwise we return False so the normal stacked/direct path loads the FP8 tensor as-is.

Source code in vllm/models/glm5next/common/model.py
def _try_load_fp8_attn_proj(
    name,
    tensor,
    buf,
    params_dict,
    loaded_params,
    kv_a_pad_size: int,
) -> bool:
    """Dequantize FP8 q_a_proj / kv_a_proj_with_mqa / o_proj to BF16 on load.

    The FP8 checkpoint stores these as block-FP8 (weight + weight_scale_inv),
    but the model holds them in BF16 (``fused_qkv_a_proj`` is always BF16 via
    DeepSeekV2FusedQkvAProjLinear; ``o_proj`` is excluded by
    modules_to_not_convert). When the model target is BF16 (no
    ``weight_scale_inv`` param) we dequantize; otherwise we return False so the
    normal stacked/direct path loads the FP8 tensor as-is.
    """
    matched = None
    for suffix, info in _FP8_ATTN_PROJS.items():
        if suffix in name:
            matched = (suffix, info)
            break
    if matched is None:
        return False
    suffix, (key, target_base, shard_id, is_kva) = matched
    is_weight = name.endswith(".weight") and tensor.dtype == torch.float8_e4m3fn
    # Need to accept both the DeepSeek-native ``weight_scale_inv`` and the Quark
    # ``weight_scale`` names before feeding the shared block dequant below.
    is_scale = "weight_scale_inv" in name or name.endswith(".weight_scale")
    if not is_weight and not is_scale:
        return False

    layer_prefix = name.rsplit(suffix, 1)[0]
    target_w = f"{layer_prefix}.{target_base}.weight"
    target_s = f"{layer_prefix}.{target_base}.weight_scale_inv"
    # If the model actually kept this projection in FP8, let the normal path
    # handle it (it has a weight_scale_inv param).
    if target_s in params_dict:
        return False

    entry = buf.setdefault(layer_prefix, {}).setdefault(key, {})
    entry["weight" if is_weight else "scale"] = tensor
    if "weight" not in entry or "scale" not in entry:
        return True

    weight_fp8, scale_inv = entry["weight"], entry["scale"]
    buf[layer_prefix].pop(key, None)
    block_size = weight_fp8.shape[1] // scale_inv.shape[1]
    weight_bf16 = _dequant_fp8_block(weight_fp8, scale_inv, block_size)
    # NoPE: pad kv_a rope portion (kv_lora_rank -> kv_lora_rank + qk_rope_head_dim).
    if is_kva and kv_a_pad_size > 0:
        pad = torch.zeros(
            kv_a_pad_size,
            weight_bf16.shape[1],
            dtype=weight_bf16.dtype,
            device=weight_bf16.device,
        )
        weight_bf16 = torch.cat([weight_bf16, pad], dim=0)

    param = params_dict[target_w]
    if shard_id is None:
        param.weight_loader(param, weight_bf16)
    else:
        param.weight_loader(param, weight_bf16, shard_id)
    loaded_params.add(target_w)
    return True

_validate_supported_config(config)

Reject checkpoints using config options this implementation lacks.

The kpool indexer kernels always keep the incomplete trailing pool, so a checkpoint asking otherwise would be served silently wrong.

Source code in vllm/models/glm5next/common/model.py
def _validate_supported_config(config: Glm5NextTextConfig) -> None:
    """Reject checkpoints using config options this implementation lacks.

    The kpool indexer kernels always keep the incomplete trailing pool, so a
    checkpoint asking otherwise would be served silently wrong.
    """
    if config.index_topk is not None and not config.index_kpool_always_select_tail:
        raise NotImplementedError(
            "GLM-5.3 sparse indexer requires index_kpool_always_select_tail=True"
        )