Files
project_6/qwen3_6_scripts/qwen3_5.py
Claude 768d89c31a fix(pybind): add py::arg + defaults to corex_gdn_chunk_recurrent
Python calls: _chunk_fn(q,k,v,g,beta, initial_state=, output_final_state=, use_qk_l2norm_in_kernel=)
C++ had: positional-only (query,key,value,g,beta,chunk_size,initial_state,output_final_state,use_qk_l2norm)

Fix: py::arg() naming + chunk_size=64 default (matches Python fallback).
Re-enable _HAS_COREX_GDN_CHUNK flag.

Rebuild on real machine:
  VLLM_ROOT=/usr/local/corex/lib64/python3/dist-packages/vllm
  bash build_corex_gdn_chunk_recurrent.sh $VLLM_ROOT
Then copy .so to prebuilt/
2026-08-14 02:02:51 +00:00

2660 lines
114 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# Inference-only Qwen3.6-35B-A3B (Qwen3_5 MoE architecture) for Iluvatar BI-V100.
# Pure-PyTorch DeltaNet (no fla / causal_conv1d dependency).
# Includes the native Qwen3.6 vision tower; MTP remains unsupported.
from functools import lru_cache, partial
import hashlib
import os
import sys
import time
from typing import (Any, Dict, Iterable, List, Literal, Mapping, Optional,
Tuple, TypedDict, Union)
def _bi100_model_trace(message: str) -> None:
if os.getenv("BI100_EXECUTOR_STARTUP_DEBUG") == "1":
stamp = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
rank = os.getenv("RANK", os.getenv("LOCAL_RANK", "?"))
print(f"[BI100 STARTUP] {stamp} pid={os.getpid()} rank={rank} {message}",
file=sys.stderr, flush=True)
_bi100_model_trace("qwen3_5 stdlib imports complete; importing torch and vLLM")
import torch
import torch.nn.functional as F
from torch import nn
from PIL import Image
from transformers.image_utils import (ChannelDimension, get_image_size,
infer_channel_dimension_format,
to_numpy_array)
from transformers.models.qwen2_vl import (
image_processing_qwen2_vl as _qwen2_vl_image_processing)
from transformers.models.qwen2_vl.image_processing_qwen2_vl import (
Qwen2VLImageProcessor, smart_resize)
def _compat_make_batched_images(images):
return images if isinstance(images, list) else [images]
def _compat_make_batched_videos(videos):
if isinstance(videos, list) and videos and isinstance(videos[0], list):
return videos
return [videos]
# The CoreX image pins transformers 4.55.3, while its vLLM Qwen2-VL module
# imports helpers introduced by another transformers build.
if not hasattr(_qwen2_vl_image_processing, "make_batched_images"):
_qwen2_vl_image_processing.make_batched_images = \
_compat_make_batched_images
if not hasattr(_qwen2_vl_image_processing, "make_batched_videos"):
_qwen2_vl_image_processing.make_batched_videos = \
_compat_make_batched_videos
from vllm.attention import Attention, AttentionMetadata
from vllm.config import (CacheConfig, LoRAConfig, MultiModalConfig,
SchedulerConfig)
from vllm.distributed import (get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce)
from vllm.model_executor.layers.activation import SiluAndMul
from vllm.model_executor.layers.layernorm import GemmaRMSNorm
from vllm.model_executor.layers.linear import (ColumnParallelLinear,
MergedColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear)
from vllm.model_executor.layers.fused_moe import FusedMoE
from vllm.model_executor.layers.logits_processor import LogitsProcessor
from vllm.model_executor.layers.quantization import QuantizationConfig
from vllm.model_executor.layers.rotary_embedding import (
MRotaryEmbedding, _apply_rotary_emb)
from vllm.model_executor.layers.sampler import Sampler, SamplerOutput
from vllm.model_executor.layers.vocab_parallel_embedding import (
ParallelLMHead, VocabParallelEmbedding)
from vllm.model_executor.model_loader.weight_utils import (
default_weight_loader, sharded_weight_loader)
from vllm.model_executor.models.mamba_cache import MambaCacheManager
from vllm.model_executor.models.qwen2_vl import (Qwen2VisionAttention,
Qwen2VisionRotaryEmbedding)
from vllm.model_executor.sampling_metadata import SamplingMetadata
from vllm.model_executor.utils import set_weight_attrs
from vllm.inputs import INPUT_REGISTRY, InputContext, LLMInputs
from vllm.multimodal import (MULTIMODAL_REGISTRY, MultiModalDataDict,
MultiModalInputs)
from vllm.multimodal.base import MultiModalData
from vllm.sequence import IntermediateTensors, SequenceData
from vllm.transformers_utils.tokenizer import get_tokenizer
from vllm.worker.model_runner import (_BATCH_SIZES_TO_CAPTURE,
_get_graph_batch_size)
from vllm.logger import init_logger
from vllm.bi100_env import env_bool, env_int
from vllm.bi100_profile import (bi100_profile_event_enabled,
bi100_profile_flush,
bi100_profile_transaction, bi100_timer)
try:
from vllm import corex_gdn_causal_conv as _corex_gdn_causal_conv
except ImportError:
_corex_gdn_causal_conv = None
try:
from vllm import corex_gdn_gated_norm as _corex_gdn_gated_norm
except ImportError:
_corex_gdn_gated_norm = None
try:
from vllm import corex_gdn_beta_decay as _corex_gdn_beta_decay
except ImportError:
_corex_gdn_beta_decay = None
try:
from vllm import corex_gdn_qk_map as _corex_gdn_qk_map
except ImportError:
_corex_gdn_qk_map = None
try:
from vllm import corex_gdn_packed_decode as _corex_gdn_packed_decode
except ImportError:
_corex_gdn_packed_decode = None
try:
from vllm import corex_attn_head_rms_norm as _corex_attn_head_rms_norm
except ImportError:
_corex_attn_head_rms_norm = None
try:
from vllm import corex_moe_exact_reduce as _corex_moe_exact_reduce
except ImportError:
_corex_moe_exact_reduce = None
try:
from vllm import corex_moe_weight_gather as _corex_moe_weight_gather
except ImportError:
_corex_moe_weight_gather = None
try:
from vllm import corex_moe_direct_routed as _corex_moe_direct_routed
except ImportError:
_corex_moe_direct_routed = None
try:
from vllm import corex_moe_topk_softmax as _corex_moe_topk_softmax
except ImportError:
_corex_moe_topk_softmax = None
try:
from vllm import corex_moe_index_combine as _corex_moe_index_combine
except ImportError:
_corex_moe_index_combine = None
try:
from vllm import corex_gdn_chunk_recurrent as _corex_gdn_chunk_recurrent
except ImportError:
_corex_gdn_chunk_recurrent = None
_HAS_COREX_GDN_CHUNK = _corex_gdn_chunk_recurrent is not None
from vllm.model_executor.models.interfaces import (HasInnerState, SupportsLoRA,
SupportsMultiModal)
logger = init_logger(__name__)
_bi100_model_trace("qwen3_5 runtime imports complete")
_ALLOW_GDN_NAN_ZERO = env_bool("BI100_GDN_ALLOW_NAN_ZERO", False)
_GDN_FINITE_CHECK = (env_bool("BI100_GDN_FINITE_CHECK", False)
or _ALLOW_GDN_NAN_ZERO)
_DNN_CHUNK_SIZE = env_int("BI100_DNN_CHUNK", 4096, 64, 65536)
_USE_COREX_GDN_CAUSAL_CONV = (
_corex_gdn_causal_conv is not None
and env_bool("BI100_GDN_COREX_CAUSAL_CONV", True))
_USE_COREX_GDN_GATED_NORM = (
_corex_gdn_gated_norm is not None
and env_bool("BI100_GDN_COREX_GATED_NORM", True))
_USE_COREX_GDN_BETA_DECAY = (
_corex_gdn_beta_decay is not None
and env_bool("BI100_GDN_COREX_BETA_DECAY", True))
_USE_COREX_GDN_QK_MAP = (
_corex_gdn_qk_map is not None
and env_bool("BI100_GDN_COREX_QK_MAP", True))
_USE_COREX_GDN_COMBINED_QK_NORM = (
_USE_COREX_GDN_QK_MAP
and env_bool("BI100_GDN_COMBINED_QK_NORM", False))
_USE_COREX_GDN_PACKED_DECODE = (
_corex_gdn_packed_decode is not None
and env_bool("BI100_GDN_COREX_PACKED_DECODE", False))
_USE_COREX_ATTN_HEAD_RMS_NORM = (
_corex_attn_head_rms_norm is not None
and env_bool("BI100_ATTN_COREX_HEAD_RMS_NORM", True))
_USE_COREX_MOE_EXACT_REDUCE = (
_corex_moe_exact_reduce is not None
and env_bool("BI100_MOE_COREX_EXACT_REDUCE", True))
_USE_COREX_MOE_WEIGHT_GATHER = (
_corex_moe_weight_gather is not None
and env_bool("BI100_MOE_COREX_WEIGHT_GATHER", True))
_USE_COREX_MOE_DIRECT_ROUTED = (
_corex_moe_direct_routed is not None
and env_bool("BI100_MOE_COREX_DIRECT_ROUTED", False))
_USE_COREX_MOE_TOPK_SOFTMAX = (
_corex_moe_topk_softmax is not None
and env_bool("BI100_MOE_COREX_TOPK_SOFTMAX", True))
_USE_COREX_MOE_INDEX_COMBINE = (
_corex_moe_index_combine is not None
and env_bool("BI100_MOE_COREX_INDEX_COMBINE", True))
_USE_FUSED_MOE_ACTIVATION = env_bool("BI100_MOE_FUSED_ACTIVATION", True)
# ---------------------------------------------------------------------------
# Qwen3.6 vision tower and vLLM 0.6 multimodal input integration
# ---------------------------------------------------------------------------
_MAX_IMAGE_TOKENS = 1280
@lru_cache(maxsize=None)
def _cached_get_qwen36_image_processor(model_path: str):
# The fast processor in transformers 4.55 calls torch.compiler APIs that
# are absent from the evaluator's torch 2.1 CoreX build.
return Qwen2VLImageProcessor.from_pretrained(model_path)
@lru_cache(maxsize=None)
def _cached_get_qwen36_tokenizer(model_path: str, trust_remote_code: bool):
return get_tokenizer(model_path, trust_remote_code=trust_remote_code)
def _image_cache_marker_tokens(image, tokenizer) -> List[int]:
array = to_numpy_array(image)
digest = hashlib.sha256()
digest.update(str(array.shape).encode("ascii"))
digest.update(str(array.dtype).encode("ascii"))
digest.update(array.tobytes())
marker = f"[image-cache-key:{digest.hexdigest()[:16]}]"
return tokenizer.encode(marker, add_special_tokens=False)
def _make_batched_images(images):
if isinstance(images, list):
if images and isinstance(images[0], list):
return [image for batch in images for image in batch]
return images
return [images]
class Qwen3_5ImagePixelInputs(TypedDict):
type: Literal["pixel_values"]
data: torch.Tensor
image_grid_thw: torch.Tensor
class Qwen3_5ImageEmbeddingInputs(TypedDict):
type: Literal["image_embeds"]
data: torch.Tensor
Qwen3_5ImageInputs = Union[Qwen3_5ImagePixelInputs,
Qwen3_5ImageEmbeddingInputs]
def _vision_pos_embed_interpolate(
embed_weight: torch.Tensor,
t: int,
h: int,
w: int,
num_grid_per_side: int,
merge_size: int,
dtype: torch.dtype,
) -> torch.Tensor:
if h % merge_size or w % merge_size:
raise ValueError(
f"vision grid {(t, h, w)} is not divisible by merge_size="
f"{merge_size}")
hidden_dim = embed_weight.shape[1]
device = embed_weight.device
h_idxs = torch.linspace(0, num_grid_per_side - 1, h,
dtype=torch.float32, device=device)
w_idxs = torch.linspace(0, num_grid_per_side - 1, w,
dtype=torch.float32, device=device)
h_floor = h_idxs.long()
w_floor = w_idxs.long()
h_ceil = torch.clamp(h_floor + 1, max=num_grid_per_side - 1)
w_ceil = torch.clamp(w_floor + 1, max=num_grid_per_side - 1)
dh = h_idxs - h_floor
dw = w_idxs - w_floor
dh_grid, dw_grid = torch.meshgrid(dh, dw, indexing="ij")
hf_grid, wf_grid = torch.meshgrid(h_floor, w_floor, indexing="ij")
hc_grid, wc_grid = torch.meshgrid(h_ceil, w_ceil, indexing="ij")
w11 = dh_grid * dw_grid
w10 = dh_grid - w11
w01 = dw_grid - w11
w00 = 1 - dh_grid - w01
h_grid = torch.stack([hf_grid, hf_grid, hc_grid, hc_grid])
w_grid = torch.stack([wf_grid, wc_grid, wf_grid, wc_grid])
indices = (h_grid * num_grid_per_side + w_grid).reshape(4, -1)
weights = torch.stack([w00, w01, w10, w11], dim=0)
weights = weights.reshape(4, -1, 1).to(dtype=dtype)
combined = (embed_weight[indices] * weights).sum(dim=0)
combined = combined.reshape(
h // merge_size, merge_size,
w // merge_size, merge_size, hidden_dim)
combined = combined.permute(0, 2, 1, 3, 4).reshape(1, -1, hidden_dim)
return combined.expand(t, -1, -1).reshape(-1, hidden_dim).to(dtype)
class Qwen3_5VisionPatchEmbed(nn.Module):
def __init__(self, vision_config) -> None:
super().__init__()
self.patch_size = vision_config.patch_size
self.temporal_patch_size = vision_config.temporal_patch_size
self.hidden_size = vision_config.hidden_size
kernel = (self.temporal_patch_size, self.patch_size, self.patch_size)
self.proj = nn.Conv3d(
vision_config.in_channels,
self.hidden_size,
kernel_size=kernel,
stride=kernel,
bias=True,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
length = x.shape[0]
x = x.view(length, -1, self.temporal_patch_size,
self.patch_size, self.patch_size)
return self.proj(x).view(length, self.hidden_size)
class Qwen3_5VisionMLP(nn.Module):
def __init__(self, vision_config,
quant_config: Optional[QuantizationConfig] = None) -> None:
super().__init__()
self.linear_fc1 = ColumnParallelLinear(
vision_config.hidden_size,
vision_config.intermediate_size,
bias=True,
quant_config=quant_config,
)
self.linear_fc2 = RowParallelLinear(
vision_config.intermediate_size,
vision_config.hidden_size,
bias=True,
quant_config=quant_config,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x, _ = self.linear_fc1(x)
x = F.gelu(x, approximate="tanh")
x, _ = self.linear_fc2(x)
return x
class Qwen3_5VisionBlock(nn.Module):
def __init__(self, vision_config,
quant_config: Optional[QuantizationConfig] = None) -> None:
super().__init__()
dim = vision_config.hidden_size
self.norm1 = nn.LayerNorm(dim, eps=1e-6)
self.norm2 = nn.LayerNorm(dim, eps=1e-6)
self.attn = Qwen2VisionAttention(
embed_dim=dim,
num_heads=vision_config.num_heads,
projection_size=dim,
quant_config=quant_config,
)
self.mlp = Qwen3_5VisionMLP(vision_config, quant_config)
def forward(self, x: torch.Tensor, cu_seqlens: torch.Tensor,
rotary_pos_emb: torch.Tensor) -> torch.Tensor:
x = x + self.attn(
self.norm1(x),
cu_seqlens=cu_seqlens,
rotary_pos_emb=rotary_pos_emb,
)
return x + self.mlp(self.norm2(x))
class Qwen3_5VisionPatchMerger(nn.Module):
def __init__(self, vision_config,
quant_config: Optional[QuantizationConfig] = None) -> None:
super().__init__()
self.hidden_size = (vision_config.hidden_size
* vision_config.spatial_merge_size ** 2)
self.norm = nn.LayerNorm(vision_config.hidden_size, eps=1e-6)
self.linear_fc1 = ColumnParallelLinear(
self.hidden_size,
self.hidden_size,
bias=True,
quant_config=quant_config,
)
self.linear_fc2 = RowParallelLinear(
self.hidden_size,
vision_config.out_hidden_size,
bias=True,
quant_config=quant_config,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.norm(x).view(-1, self.hidden_size)
x, _ = self.linear_fc1(x)
x = F.gelu(x)
x, _ = self.linear_fc2(x)
return x
class Qwen3_5VisionTransformer(nn.Module):
def __init__(self, vision_config,
quant_config: Optional[QuantizationConfig] = None) -> None:
super().__init__()
self.hidden_size = vision_config.hidden_size
self.num_heads = vision_config.num_heads
self.spatial_merge_size = vision_config.spatial_merge_size
self.num_grid_per_side = int(vision_config.num_position_embeddings ** .5)
self.patch_embed = Qwen3_5VisionPatchEmbed(vision_config)
self.pos_embed = nn.Embedding(
vision_config.num_position_embeddings, self.hidden_size)
head_dim = self.hidden_size // self.num_heads
self.rotary_pos_emb = Qwen2VisionRotaryEmbedding(head_dim // 2)
self.blocks = nn.ModuleList([
Qwen3_5VisionBlock(vision_config, quant_config)
for _ in range(vision_config.depth)
])
self.merger = Qwen3_5VisionPatchMerger(vision_config, quant_config)
@property
def dtype(self) -> torch.dtype:
return self.patch_embed.proj.weight.dtype
@property
def device(self) -> torch.device:
return self.patch_embed.proj.weight.device
def _rot_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor:
pos_ids = []
for t, h, w in grid_thw.tolist():
h_ids = torch.arange(h).unsqueeze(1).expand(-1, w)
w_ids = torch.arange(w).unsqueeze(0).expand(h, -1)
h_ids = h_ids.reshape(
h // self.spatial_merge_size, self.spatial_merge_size,
w // self.spatial_merge_size, self.spatial_merge_size,
).permute(0, 2, 1, 3).flatten()
w_ids = w_ids.reshape(
h // self.spatial_merge_size, self.spatial_merge_size,
w // self.spatial_merge_size, self.spatial_merge_size,
).permute(0, 2, 1, 3).flatten()
pos_ids.append(torch.stack([h_ids, w_ids], dim=-1).repeat(t, 1))
pos_ids_t = torch.cat(pos_ids, dim=0).to(self.device)
max_grid_size = int(grid_thw[:, 1:].max().item())
return self.rotary_pos_emb(max_grid_size)[pos_ids_t].flatten(1)
def _absolute_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor:
return torch.cat([
_vision_pos_embed_interpolate(
self.pos_embed.weight, int(t), int(h), int(w),
self.num_grid_per_side, self.spatial_merge_size, self.dtype)
for t, h, w in grid_thw.tolist()
], dim=0)
def forward(self, x: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor:
x = x.to(device=self.device, dtype=self.dtype)
grid_thw = grid_thw.to(device=self.device)
x = self.patch_embed(x)
x = x + self._absolute_pos_emb(grid_thw)
rotary_pos_emb = self._rot_pos_emb(grid_thw)
cu_seqlens = torch.repeat_interleave(
grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0],
).cumsum(dim=0, dtype=torch.int32)
cu_seqlens = F.pad(cu_seqlens, (1, 0), "constant", 0)
x = x.unsqueeze(1)
for block in self.blocks:
x = block(x, cu_seqlens, rotary_pos_emb)
return self.merger(x)
class Qwen3_5InterleavedMRotaryEmbedding(MRotaryEmbedding):
"""Qwen3.5 frequency-interleaved T/H/W rotary embedding."""
def forward(self, positions: torch.Tensor, query: torch.Tensor,
key: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
if positions.ndim not in (1, 2):
raise ValueError(f"invalid MRoPE positions shape {positions.shape}")
num_tokens = positions.shape[-1]
cos_sin = self.cos_sin_cache[positions]
cos_all, sin_all = cos_sin.chunk(2, dim=-1)
if positions.ndim == 2:
if not self.mrope_section:
raise ValueError("mrope_section is required")
cos = cos_all[0].clone()
sin = sin_all[0].clone()
for dim, offset in enumerate((1, 2), start=1):
stop = self.mrope_section[dim] * 3
cos[..., offset:stop:3] = cos_all[dim, ..., offset:stop:3]
sin[..., offset:stop:3] = sin_all[dim, ..., offset:stop:3]
else:
cos, sin = cos_all, sin_all
query_shape = query.shape
query = query.view(num_tokens, -1, self.head_size)
query_rot = _apply_rotary_emb(
query[..., :self.rotary_dim], cos, sin, self.is_neox_style)
query = torch.cat((query_rot, query[..., self.rotary_dim:]), dim=-1)
key_shape = key.shape
key = key.view(num_tokens, -1, self.head_size)
key_rot = _apply_rotary_emb(
key[..., :self.rotary_dim], cos, sin, self.is_neox_style)
key = torch.cat((key_rot, key[..., self.rotary_dim:]), dim=-1)
return query.reshape(query_shape), key.reshape(key_shape)
def _qwen36_pixel_limits(image_processor) -> Tuple[int, int]:
min_pixels = 256 * 256
configured_max = 4096 * 4096
runtime_max = _MAX_IMAGE_TOKENS * (
image_processor.patch_size * image_processor.merge_size) ** 2
return min_pixels, min(configured_max, runtime_max)
def _qwen36_image_token_count(image, image_processor) -> int:
if isinstance(image, Image.Image):
image = image.convert("RGB")
image_array = to_numpy_array(image)
height, width = get_image_size(
image_array, channel_dim=ChannelDimension.LAST)
min_pixels, max_pixels = _qwen36_pixel_limits(image_processor)
if getattr(image_processor, "do_resize", True):
height, width = smart_resize(
height=height,
width=width,
factor=image_processor.patch_size * image_processor.merge_size,
min_pixels=min_pixels,
max_pixels=max_pixels,
)
return (height // image_processor.patch_size
* width // image_processor.patch_size
// image_processor.merge_size ** 2)
def qwen36_image_input_mapper(
ctx: InputContext,
data: MultiModalData[object],
) -> MultiModalInputs:
if isinstance(data, dict):
return MultiModalInputs({
"image_embeds": data.get("image_embeds"),
"image_grid_thw": data.get("image_grid_thw"),
})
image_processor = _cached_get_qwen36_image_processor(
ctx.model_config.model)
min_pixels, max_pixels = _qwen36_pixel_limits(image_processor)
batch_data = image_processor.preprocess(
images=data,
return_tensors="pt",
size={"shortest_edge": min_pixels, "longest_edge": max_pixels},
do_convert_rgb=True,
input_data_format=ChannelDimension.LAST,
).data
return MultiModalInputs(batch_data)
def get_max_qwen36_image_tokens(_ctx: InputContext) -> int:
return _MAX_IMAGE_TOKENS
def dummy_data_for_qwen36(
ctx: InputContext,
seq_len: int,
mm_counts: Mapping[str, int],
) -> Tuple[SequenceData, Optional[MultiModalDataDict]]:
num_images = mm_counts.get("image", 0)
image_tokens = _MAX_IMAGE_TOKENS * num_images
if seq_len < image_tokens + 2:
raise RuntimeError(
f"Qwen3.6 needs {image_tokens + 2} tokens for {num_images} "
f"max-size image(s), but max_model_len is {seq_len}")
config = ctx.model_config.hf_config
seq_data = SequenceData.from_token_counts(
(config.vision_start_token_id, 1),
(config.image_token_id, image_tokens),
(config.vision_end_token_id, 1),
(0, seq_len - image_tokens - 2),
)
dummy_image = Image.new("RGB", (1280, 1024), color=0)
return seq_data, {
"image": (dummy_image if num_images == 1
else [dummy_image] * num_images)
}
def input_processor_for_qwen36(ctx: InputContext,
llm_inputs: LLMInputs) -> LLMInputs:
multi_modal_data = llm_inputs.get("multi_modal_data")
if not multi_modal_data or "image" not in multi_modal_data:
return llm_inputs
images = multi_modal_data["image"]
prompt_token_ids = llm_inputs.get("prompt_token_ids")
if prompt_token_ids is None:
raise ValueError("Qwen3.6 image requests require tokenized prompt input")
config = ctx.model_config.hf_config
image_processor = _cached_get_qwen36_image_processor(
ctx.model_config.model)
tokenizer = _cached_get_qwen36_tokenizer(
ctx.model_config.tokenizer,
ctx.model_config.trust_remote_code,
)
batched_images = _make_batched_images(images)
image_indices = [
idx for idx, token in enumerate(prompt_token_ids)
if token == config.image_token_id
]
if len(image_indices) != len(batched_images):
raise ValueError(
f"found {len(image_indices)} image placeholders for "
f"{len(batched_images)} image(s)")
expanded = []
previous = 0
for index, image in zip(image_indices, batched_images):
vision_start = index - 1
if (vision_start < previous
or prompt_token_ids[vision_start]
!= config.vision_start_token_id):
raise ValueError("image token is not preceded by vision_start")
expanded.extend(prompt_token_ids[previous:vision_start])
expanded.extend(_image_cache_marker_tokens(image, tokenizer))
expanded.extend(prompt_token_ids[vision_start:index])
expanded.extend([config.image_token_id]
* _qwen36_image_token_count(image, image_processor))
previous = index + 1
expanded.extend(prompt_token_ids[previous:])
return LLMInputs(
prompt_token_ids=expanded,
prompt=llm_inputs["prompt"],
multi_modal_data=multi_modal_data,
)
# ---------------------------------------------------------------------------
# Pure-PyTorch DeltaNet kernels (fallbacks from transformers 5.2.0)
# ---------------------------------------------------------------------------
def _l2norm(x: torch.Tensor, dim: int = -1, eps: float = 1e-6) -> torch.Tensor:
return x * torch.rsqrt((x * x).sum(dim=dim, keepdim=True) + eps)
def _check_gdn_finite(tensor: torch.Tensor, *, layer_idx: int,
stage: str) -> torch.Tensor:
if not _GDN_FINITE_CHECK:
return tensor
if torch.isfinite(tensor).all():
return tensor
bad = (~torch.isfinite(tensor)).float().mean().item()
msg = (
f"non-finite values in {stage} GatedDeltaNet layer {layer_idx} "
f"(frac={bad:.4f})"
)
if not _ALLOW_GDN_NAN_ZERO:
raise RuntimeError(msg)
logger.warning("%s; replacing with zeros because BI100_GDN_ALLOW_NAN_ZERO=1",
msg)
return torch.nan_to_num(tensor, nan=0.0, posinf=0.0, neginf=0.0)
def _gdn_segment_ends(seq_len: int, chunk_size: int,
capture_offsets: Iterable[int]) -> List[int]:
ends = list(range(chunk_size, seq_len, chunk_size))
ends.append(seq_len)
ends.extend(offset for offset in capture_offsets
if 0 < offset < seq_len)
return sorted(set(ends))
def _validate_gdn_prefix_key(key: Any) -> Tuple[int, bytes]:
if (not isinstance(key, tuple) or len(key) != 2
or not isinstance(key[0], int) or key[0] <= 0
or not isinstance(key[1], bytes) or len(key[1]) != 32):
raise RuntimeError(f"invalid GDN prefix key: {key!r}")
return key
def _torch_causal_conv1d_update(
hidden_states: torch.Tensor, # (batch, channels, seq=1)
conv_state: torch.Tensor, # (batch, channels, state_len) modified in-place
weight: torch.Tensor, # (channels, kernel_size)
bias: Optional[torch.Tensor] = None,
activation: Optional[str] = None,
) -> torch.Tensor:
_, channels, seq_len = hidden_states.shape
state_len = conv_state.shape[-1]
cat = torch.cat([conv_state, hidden_states], dim=-1).to(weight.dtype)
conv_state.copy_(cat[:, :, -state_len:])
out = F.conv1d(cat, weight.unsqueeze(1), bias, padding=0, groups=channels)
out = out[:, :, -seq_len:]
if activation is not None:
out = F.silu(out)
return out.to(hidden_states.dtype)
def _torch_chunk_gated_delta_rule(
query: torch.Tensor, # (batch, seq, num_heads, head_k_dim)
key: torch.Tensor,
value: torch.Tensor, # (batch, seq, num_heads, head_v_dim)
g: torch.Tensor, # (batch, seq, num_heads)
beta: torch.Tensor, # (batch, seq, num_heads)
chunk_size: int = 64,
initial_state: Optional[torch.Tensor] = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
if use_qk_l2norm_in_kernel:
query = _l2norm(query)
key = _l2norm(key)
# Transpose to (batch, num_heads, seq, dim)
query, key, value, beta, g = [
x.transpose(1, 2).contiguous().to(torch.float32)
for x in (query, key, value, beta, g)
]
batch, num_heads, seq_len, k_dim = key.shape
v_dim = value.shape[-1]
pad = (chunk_size - seq_len % chunk_size) % chunk_size
query = F.pad(query, (0, 0, 0, pad))
key = F.pad(key, (0, 0, 0, pad))
value = F.pad(value, (0, 0, 0, pad))
beta = F.pad(beta, (0, pad))
g = F.pad(g, (0, pad))
total_len = seq_len + pad
scale = 1.0 / (query.shape[-1] ** 0.5)
query = query * scale
v_beta = value * beta.unsqueeze(-1)
k_beta = key * beta.unsqueeze(-1)
query, key, value, k_beta, v_beta = [
x.reshape(x.shape[0], x.shape[1], -1, chunk_size, x.shape[-1])
for x in (query, key, value, k_beta, v_beta)
]
g = g.reshape(g.shape[0], g.shape[1], -1, chunk_size)
mask_upper = torch.triu(
torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device),
diagonal=0)
g = g.cumsum(dim=-1)
decay_mask = ((g.unsqueeze(-1) - g.unsqueeze(-2)).tril().exp().float()).tril()
attn = -((k_beta @ key.transpose(-1, -2)) * decay_mask).masked_fill(mask_upper, 0)
for i in range(1, chunk_size):
row = attn[..., i, :i].clone()
sub = attn[..., :i, :i].clone()
attn[..., i, :i] = row + (row.unsqueeze(-1) * sub).sum(-2)
attn = attn + torch.eye(chunk_size, dtype=attn.dtype, device=attn.device)
value = attn @ v_beta
k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1))
last_state = (
torch.zeros(batch, num_heads, k_dim, v_dim, dtype=value.dtype, device=value.device)
if initial_state is None
else initial_state.to(value)
)
core_out = torch.zeros_like(value)
mask_upper2 = torch.triu(
torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device),
diagonal=1)
for i in range(total_len // chunk_size):
q_i, k_i, v_i = query[:, :, i], key[:, :, i], value[:, :, i]
attn_i = (q_i @ k_i.transpose(-1, -2) * decay_mask[:, :, i]).masked_fill_(mask_upper2, 0)
v_prime = k_cumdecay[:, :, i] @ last_state
v_new = v_i - v_prime
attn_inter = (q_i * g[:, :, i, :, None].exp()) @ last_state
core_out[:, :, i] = attn_inter + attn_i @ v_new
last_state = (
last_state * g[:, :, i, -1, None, None].exp()
+ (k_i * (g[:, :, i, -1, None] - g[:, :, i]).exp()[..., None])
.transpose(-1, -2) @ v_new
)
if not output_final_state:
last_state = None
core_out = core_out.reshape(batch, num_heads, -1, v_dim)[:, :, :seq_len]
core_out = core_out.transpose(1, 2).contiguous()
return core_out, last_state
def _torch_recurrent_gated_delta_rule(
query: torch.Tensor, # (batch, 1, num_heads, head_k_dim)
key: torch.Tensor,
value: torch.Tensor,
g: torch.Tensor, # (batch, 1, num_heads)
beta: torch.Tensor,
initial_state: Optional[torch.Tensor] = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
if use_qk_l2norm_in_kernel:
query = _l2norm(query)
key = _l2norm(key)
query, key, value, beta, g = [
x.transpose(1, 2).contiguous().to(torch.float32)
for x in (query, key, value, beta, g)
]
batch, num_heads, seq_len, k_dim = key.shape
v_dim = value.shape[-1]
scale = 1.0 / (query.shape[-1] ** 0.5)
query = query * scale
core_out = torch.zeros(batch, num_heads, seq_len, v_dim,
dtype=value.dtype, device=value.device)
last_state = (
torch.zeros(batch, num_heads, k_dim, v_dim,
dtype=value.dtype, device=value.device)
if initial_state is None
else initial_state.to(value)
)
for t in range(seq_len):
q_t = query[:, :, t]
k_t = key[:, :, t]
v_t = value[:, :, t]
g_t = g[:, :, t].exp().unsqueeze(-1).unsqueeze(-1)
beta_t = beta[:, :, t].unsqueeze(-1)
last_state = last_state * g_t
kv_mem = (last_state * k_t.unsqueeze(-1)).sum(dim=-2)
delta = (v_t - kv_mem) * beta_t
last_state = last_state + k_t.unsqueeze(-1) * delta.unsqueeze(-2)
core_out[:, :, t] = (last_state * q_t.unsqueeze(-1)).sum(dim=-2)
if not output_final_state:
last_state = None
core_out = core_out.transpose(1, 2).contiguous()
return core_out, last_state
# ---------------------------------------------------------------------------
# Gated RMSNorm (for DeltaNet output normalisation)
# ---------------------------------------------------------------------------
class Qwen3_5RMSNormGated(nn.Module):
def __init__(self, hidden_size: int, eps: float = 1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size))
self.variance_epsilon = eps
def forward(self, hidden_states: torch.Tensor,
gate: torch.Tensor) -> torch.Tensor:
input_dtype = hidden_states.dtype
hs = hidden_states.to(torch.float32)
variance = hs.pow(2).mean(-1, keepdim=True)
hs = hs * torch.rsqrt(variance + self.variance_epsilon)
hs = self.weight * hs.to(input_dtype)
return (hs * F.silu(gate.to(torch.float32))).to(input_dtype)
def forward_decode(self, hidden_states: torch.Tensor,
gate: torch.Tensor) -> torch.Tensor:
if (_USE_COREX_GDN_GATED_NORM
and hidden_states.dtype == torch.float32
and gate.dtype == torch.float16
and self.weight.dtype == torch.float16
and hidden_states.shape[-1] == 128):
hs = hidden_states.float()
inverse = torch.rsqrt(
hs.pow(2).mean(-1, keepdim=True) + self.variance_epsilon)
return _corex_gdn_gated_norm.apply_inverse(
hs, gate, self.weight, inverse)
return self.forward(hidden_states, gate).to(gate.dtype)
def _load_gdn_projection_weight(params_dict, name: str,
loaded_weight: torch.Tensor,
text_cfg) -> bool:
projections = {
"in_proj_qkv": None,
"in_proj_z": 3,
"in_proj_b": 4,
"in_proj_a": 5,
}
source = next((projection for projection in projections
if f".linear_attn.{projection}." in name), None)
if source is None:
return False
target_name = name.replace(
f".linear_attn.{source}.",
".linear_attn.in_proj_qkvzba.",
)
if target_name not in params_dict:
raise ValueError(f"missing fused GDN projection parameter: {target_name}")
param = params_dict[target_name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
if source == "in_proj_qkv":
key_dim = (text_cfg.linear_num_key_heads
* text_cfg.linear_key_head_dim)
value_dim = (text_cfg.linear_num_value_heads
* text_cfg.linear_value_head_dim)
shard_sizes = (key_dim, key_dim, value_dim)
if loaded_weight.shape[0] != sum(shard_sizes):
raise ValueError(
"unexpected fused QKV output size: "
f"{loaded_weight.shape[0]} != {sum(shard_sizes)}")
for shard_id, shard in enumerate(
torch.split(loaded_weight, shard_sizes, dim=0)):
weight_loader(param, shard, shard_id)
else:
weight_loader(param, loaded_weight, projections[source])
return True
def _load_full_attention_qgkv_weight(params_dict, name: str,
loaded_weight: torch.Tensor,
text_cfg) -> bool:
projections = {"q_proj": 0, "k_proj": 1, "v_proj": 2}
source = next((projection for projection in projections
if f".self_attn.{projection}." in name), None)
if source is None:
return False
target_name = name.replace(
f".self_attn.{source}.", ".self_attn.qgkv_proj.")
if target_name not in params_dict:
return False
tp_size = get_tensor_model_parallel_world_size()
tp_rank = get_tensor_model_parallel_rank()
qg_dim = text_cfg.num_attention_heads * text_cfg.head_dim * 2
if qg_dim % tp_size != 0:
raise ValueError(f"QG output size {qg_dim} is not divisible by TP {tp_size}")
local_qg_dim = qg_dim // tp_size
kv_dim = text_cfg.num_key_value_heads * text_cfg.head_dim
expected_rows = qg_dim if source == "q_proj" else kv_dim
if loaded_weight.shape[0] != expected_rows:
raise ValueError(
f"unexpected full-attention {source} output size: "
f"{loaded_weight.shape[0]} != {expected_rows}")
if source == "q_proj":
loaded_weight = loaded_weight.narrow(
0, tp_rank * local_qg_dim, local_qg_dim)
offset = 0
elif source == "k_proj":
offset = local_qg_dim
else:
offset = local_qg_dim + kv_dim
param = params_dict[target_name]
default_weight_loader(
param[offset:offset + loaded_weight.shape[0]], loaded_weight)
return True
# ---------------------------------------------------------------------------
# Gated DeltaNet (linear_attention layers)
# ---------------------------------------------------------------------------
class GatedDeltaNet(nn.Module):
def __init__(
self,
text_cfg,
layer_idx: int,
quant_config: Optional[QuantizationConfig] = None,
) -> None:
super().__init__()
self.layer_idx = layer_idx
self.hidden_size = text_cfg.hidden_size
self.num_v_heads = text_cfg.linear_num_value_heads # checkpoint: 32
self.num_k_heads = text_cfg.linear_num_key_heads # checkpoint: 16
self.head_k_dim = text_cfg.linear_key_head_dim # 128
self.head_v_dim = text_cfg.linear_value_head_dim # 128
self.key_dim = self.num_k_heads * self.head_k_dim # 2048
self.value_dim = self.num_v_heads * self.head_v_dim # checkpoint: 4096
self.conv_dim = self.key_dim * 2 + self.value_dim # checkpoint: 8192
self.conv_kernel_size = text_cfg.linear_conv_kernel_dim # 4
self.head_expand_ratio = self.num_v_heads // self.num_k_heads # checkpoint: 2
tp_size = get_tensor_model_parallel_world_size()
# Keep each logical projection independently TP-sharded while executing
# one GEMM. Per-rank output order is [q, k, v, z, beta, decay].
self.in_proj_qkvzba = MergedColumnParallelLinear(
self.hidden_size,
[self.key_dim, self.key_dim, self.value_dim, self.value_dim,
self.num_v_heads, self.num_v_heads],
bias=False, quant_config=quant_config)
self.out_proj = RowParallelLinear(
self.value_dim, self.hidden_size,
bias=False, quant_config=quant_config)
# Depthwise conv weight — sharded along channel dim (dim 0)
local_conv_dim = self.conv_dim // tp_size
self.conv1d_weight = nn.Parameter(
torch.empty(local_conv_dim, 1, self.conv_kernel_size))
set_weight_attrs(self.conv1d_weight, {
"weight_loader": self._conv1d_weight_loader})
# Per-head scalar parameters — sharded along dim 0
local_num_v = self.num_v_heads // tp_size
self.A_log = nn.Parameter(torch.zeros(local_num_v))
self.dt_bias = nn.Parameter(torch.zeros(local_num_v))
set_weight_attrs(self.A_log, {"weight_loader": sharded_weight_loader(0)})
set_weight_attrs(self.dt_bias, {"weight_loader": sharded_weight_loader(0)})
# Gated RMSNorm on head_v_dim — replicated (head_v_dim=128 is small)
self.norm = Qwen3_5RMSNormGated(self.head_v_dim,
eps=text_cfg.rms_norm_eps)
self.captured_conv_states: Dict[int, torch.Tensor] = {}
self.captured_temporal_states: Dict[int, torch.Tensor] = {}
def _conv1d_weight_loader(self, param: torch.Tensor,
loaded_weight: torch.Tensor) -> None:
# loaded_weight is ordered as [q, k, v] along its channel dimension.
# Must gather channels in the same non-contiguous pattern that
# MergedColumnParallelLinear uses for in_proj_qkv, so that each rank's
# conv1d_weight[i] applies to the correct in_proj_qkv output channel.
tp_rank = get_tensor_model_parallel_rank()
tp_size = get_tensor_model_parallel_world_size()
key_local = self.key_dim // tp_size # 512 with TP=4
val_local = self.value_dim // tp_size # 1024 with TP=4
q_s = loaded_weight[tp_rank * key_local : (tp_rank + 1) * key_local]
k_s = loaded_weight[self.key_dim + tp_rank * key_local :
self.key_dim + (tp_rank + 1) * key_local]
v_s = loaded_weight[2 * self.key_dim + tp_rank * val_local :
2 * self.key_dim + (tp_rank + 1) * val_local]
param.data.copy_(torch.cat([q_s, k_s, v_s], dim=0))
def forward(
self,
hidden_states: torch.Tensor, # (total_tokens, hidden_size)
attn_metadata: AttentionMetadata,
conv_state: torch.Tensor, # (batch, local_conv_dim, kernel-1) in-place
temporal_state: torch.Tensor, # (batch, local_v_heads, k_dim, v_dim) in-place
capture_offsets: Optional[Iterable[int]] = None,
segment_offsets: Optional[Iterable[int]] = None,
) -> torch.Tensor:
tp_size = get_tensor_model_parallel_world_size()
local_key_dim = self.key_dim // tp_size
local_val_dim = self.value_dim // tp_size
local_num_v = self.num_v_heads // tp_size
local_num_k = self.num_k_heads // tp_size
local_conv_dim = self.conv_dim // tp_size
self.captured_conv_states = {}
self.captured_temporal_states = {}
is_prefill = attn_metadata.num_prefill_tokens > 0
projected, _ = self.in_proj_qkvzba(hidden_states)
mixed_qkv_all, z_all, b_all, a_all = torch.split(
projected,
[local_conv_dim, local_val_dim, local_num_v, local_num_v],
dim=-1,
)
if is_prefill:
seq_starts = attn_metadata.query_start_loc.tolist()
outputs = []
state_len = self.conv_kernel_size - 1
weight_2d = self.conv1d_weight.squeeze(1) # (local_conv_dim, kernel)
for si in range(len(seq_starts) - 1):
s, e = int(seq_starts[si]), int(seq_starts[si + 1])
seq_len = e - s
# Shape: (1, local_conv_dim, seq_len)
mixed_qkv = (mixed_qkv_all[s:e]
.transpose(0, 1).unsqueeze(0)
.to(weight_2d.dtype))
# Load prev conv state BEFORE overwriting (needed for causal conv padding).
# For first prefill of a request: mamba_cache is zeros → correct.
# For chunked prefill chunk 2+: carries last state_len tokens from prev chunk.
prev_conv = conv_state[si:si + 1].clone().to(weight_2d.dtype) # [1, local_conv_dim, state_len]
# Save conv state (last state_len positions)
if seq_len >= state_len:
conv_state[si].copy_(mixed_qkv[0, :, -state_len:])
else:
conv_state[si, :, state_len - seq_len:].copy_(
mixed_qkv[0])
conv_state[si, :, :state_len - seq_len] = 0
# Causal conv: left-pad with previous conv state (not zeros).
padded = torch.cat([prev_conv, mixed_qkv], dim=2)
seq_capture_offsets = (set(capture_offsets or ())
if si == 0 else set())
seq_segment_offsets = (set(segment_offsets or ())
if si == 0 else set())
for capture_offset in seq_capture_offsets:
if 0 < capture_offset < seq_len:
self.captured_conv_states[capture_offset] = padded[
0, :, capture_offset:
capture_offset + state_len].clone()
mixed_qkv_conv = F.conv1d(
padded, self.conv1d_weight,
bias=None, padding=0, groups=local_conv_dim)
mixed_qkv_conv = F.silu(mixed_qkv_conv)
# (1, seq_len, local_conv_dim)
mixed_qkv_conv = mixed_qkv_conv.squeeze(0).transpose(0, 1).unsqueeze(0)
q, k, v = torch.split(
mixed_qkv_conv,
[local_key_dim, local_key_dim, local_val_dim], dim=-1)
q = q.reshape(1, seq_len, local_num_k, self.head_k_dim)
k = k.reshape(1, seq_len, local_num_k, self.head_k_dim)
v = v.reshape(1, seq_len, local_num_v, self.head_v_dim)
beta = b_all[s:e].sigmoid().unsqueeze(0) # (1, seq_len, local_num_v)
g = (-self.A_log.float().exp()
* F.softplus(a_all[s:e].float() + self.dt_bias)
).unsqueeze(0) # (1, seq_len, local_num_v)
# Expand k/q to match num_v_heads
q = q.repeat_interleave(self.head_expand_ratio, dim=2)
k = k.repeat_interleave(self.head_expand_ratio, dim=2)
# Sub-sequence chunking: call _torch_chunk_gated_delta_rule
# on _DNN_CHUNK tokens at a time to cap peak memory.
# Full 18K: tensors [1,6,282,64,64]=220 MB each → ~990 MB/call.
# With _DNN_CHUNK=4096: [1,6,64,64,64]=6 MB each → ~137 MB/call.
# State is chained via initial_state / output_final_state.
cur_state = temporal_state[si:si + 1].clone()
core_out_parts = []
segment_ends = _gdn_segment_ends(
seq_len, _DNN_CHUNK_SIZE,
seq_capture_offsets | seq_segment_offsets)
sc_start = 0
_chunk_fn = (
_corex_gdn_chunk_recurrent.torch_chunk_gated_delta_rule
if _HAS_COREX_GDN_CHUNK
else _torch_chunk_gated_delta_rule
)
with bi100_timer(f"L{self.layer_idx}.gdn.prefill"):
for sc_end in segment_ends:
c_out, cur_state = _chunk_fn(
q[:, sc_start:sc_end],
k[:, sc_start:sc_end],
v[:, sc_start:sc_end],
g[:, sc_start:sc_end],
beta[:, sc_start:sc_end],
initial_state=cur_state,
output_final_state=True,
use_qk_l2norm_in_kernel=True,
)
core_out_parts.append(c_out)
if sc_end in seq_capture_offsets:
self.captured_temporal_states[sc_end] = (
cur_state[0].clone())
sc_start = sc_end
if cur_state is not None:
temporal_state[si].copy_(cur_state[0])
# [1, seq_len, num_v_heads, head_v_dim]
core_out = torch.cat(core_out_parts, dim=1)
# Gate + norm + output proj
z = z_all[s:e].reshape(seq_len, local_num_v, self.head_v_dim)
core_out = core_out.reshape(seq_len, local_num_v, self.head_v_dim)
normed = self.norm(
core_out.reshape(-1, self.head_v_dim),
z.reshape(-1, self.head_v_dim))
normed = _check_gdn_finite(
normed, layer_idx=self.layer_idx,
stage="prefill-norm").reshape(seq_len, -1)
normed = normed.to(z_all.dtype)
out, _ = self.out_proj(normed)
outputs.append(out)
result = torch.cat(outputs, dim=0)
return _check_gdn_finite(
result, layer_idx=self.layer_idx, stage="prefill-output")
else:
# Decode: one token per sequence
num_seqs = hidden_states.shape[0]
weight_2d = self.conv1d_weight.squeeze(1)
# (num_seqs, local_conv_dim, 1)
mixed_qkv = (mixed_qkv_all
.to(weight_2d.dtype)
.unsqueeze(-1))
if _USE_COREX_GDN_CAUSAL_CONV:
mixed_qkv_conv = _corex_gdn_causal_conv.causal_conv_update(
conv_state, mixed_qkv, weight_2d)
else:
mixed_qkv_conv = _torch_causal_conv1d_update(
mixed_qkv, conv_state, weight_2d,
bias=None, activation='silu')
# (num_seqs, local_conv_dim, 1) → (num_seqs, 1, local_conv_dim)
mixed_qkv_conv = mixed_qkv_conv.squeeze(-1).unsqueeze(1)
packed_mixed_qkv = mixed_qkv_conv.squeeze(1)
use_corex_packed_decode = (
_USE_COREX_GDN_PACKED_DECODE
and num_seqs == 1
and local_num_k == 4
and local_num_v == 8
and self.head_k_dim == 128
and self.head_v_dim == 128
and packed_mixed_qkv.dtype == torch.float16
and packed_mixed_qkv.shape == (1, 2048)
and packed_mixed_qkv.is_contiguous()
and b_all.dtype == torch.float16
and b_all.shape == (1, 8)
and b_all.is_contiguous()
and a_all.dtype == torch.float16
and a_all.shape == (1, 8)
and a_all.is_contiguous()
and self.A_log.dtype == torch.float16
and self.A_log.shape == (8,)
and self.A_log.is_contiguous()
and self.dt_bias.dtype == torch.float16
and self.dt_bias.shape == (8,)
and self.dt_bias.is_contiguous()
and temporal_state.dtype == torch.float32
and temporal_state.shape == (1, 8, 128, 128)
and temporal_state.is_contiguous())
if use_corex_packed_decode:
with bi100_timer(f"L{self.layer_idx}.gdn.decode"):
core_out = _corex_gdn_packed_decode.packed_decode(
temporal_state, packed_mixed_qkv, b_all, a_all,
self.A_log, self.dt_bias)
else:
q, k, v = torch.split(
mixed_qkv_conv,
[local_key_dim, local_key_dim, local_val_dim], dim=-1)
q = q.reshape(num_seqs, 1, local_num_k, self.head_k_dim)
k = k.reshape(num_seqs, 1, local_num_k, self.head_k_dim)
v = v.reshape(num_seqs, 1, local_num_v, self.head_v_dim)
use_corex_beta_decay = (
_USE_COREX_GDN_BETA_DECAY
and b_all.dtype == torch.float16
and a_all.dtype == torch.float16
and self.A_log.dtype == torch.float16
and self.dt_bias.dtype == torch.float16
and b_all.is_contiguous()
and a_all.is_contiguous())
if use_corex_beta_decay:
beta_decay = _corex_gdn_beta_decay.beta_decay(
b_all, a_all, self.A_log, self.dt_bias)
bt = beta_decay[0]
g_t = beta_decay[1]
else:
beta = b_all.sigmoid()
g = (-self.A_log.float().exp()
* F.softplus(a_all.float() + self.dt_bias))
bt = beta.float()
g_t = g.float().exp_()
# Inlined decode recurrent step (seq_len=1).
# Uses bmm/baddbmm_ to avoid large intermediate tensors.
_scale = self.head_k_dim ** -0.5
q_raw = q.squeeze(1)
k_raw = k.squeeze(1)
use_corex_qk_map = (
_USE_COREX_GDN_QK_MAP
and q_raw.dtype == torch.float16
and k_raw.dtype == torch.float16
and self.head_k_dim == 128
and q_raw.is_contiguous()
and k_raw.is_contiguous())
if use_corex_qk_map:
use_combined_qk_norm = (
_USE_COREX_GDN_COMBINED_QK_NORM
and num_seqs == 1
and local_num_k == 4
and local_num_v == 8
and packed_mixed_qkv.dtype == torch.float16
and packed_mixed_qkv.shape == (1, 2048)
and packed_mixed_qkv.is_contiguous())
if use_combined_qk_norm:
raw_qk = packed_mixed_qkv.narrow(
1, 0, 2 * local_key_dim).view(
num_seqs, 2 * local_num_k,
self.head_k_dim)
normalized_qk = _l2norm(raw_qk)
normalized_q, normalized_k = torch.split(
normalized_qk, local_num_k, dim=1)
else:
normalized_q = _l2norm(q_raw)
normalized_k = _l2norm(k_raw)
qk_mapped = _corex_gdn_qk_map.qk_map(
normalized_q, normalized_k, local_num_v)
q_t = qk_mapped[0]
k_t = qk_mapped[1]
else:
q_expanded = q_raw.repeat_interleave(
self.head_expand_ratio, dim=1)
k_expanded = k_raw.repeat_interleave(
self.head_expand_ratio, dim=1)
q_t = _l2norm(q_expanded).float() * _scale
k_t = _l2norm(k_expanded).float()
v_t = v.squeeze(1).float()
with bi100_timer(f"L{self.layer_idx}.gdn.decode"):
# State shape is (B, H_v, k_dim, v_dim).
temporal_state.mul_(g_t[:, :, None, None])
ts_flat = temporal_state.view(
-1, self.head_k_dim, self.head_v_dim)
BH = ts_flat.shape[0]
kv_mem = torch.bmm(
k_t.view(BH, 1, self.head_k_dim), ts_flat
).view(num_seqs, local_num_v, self.head_v_dim)
delta = (v_t - kv_mem) * bt[:, :, None]
ts_flat.baddbmm_(
k_t.view(BH, self.head_k_dim, 1),
delta.view(BH, 1, self.head_v_dim),
)
core_out = torch.bmm(
q_t.view(BH, 1, self.head_k_dim), ts_flat
).view(num_seqs, local_num_v, self.head_v_dim)
# core_out: (B, H_v, v_dim) = (num_seqs, local_num_v, head_v_dim) already
z = z_all.reshape(num_seqs, local_num_v, self.head_v_dim)
normed = self.norm.forward_decode(
core_out.reshape(-1, self.head_v_dim),
z.reshape(-1, self.head_v_dim))
normed = _check_gdn_finite(
normed, layer_idx=self.layer_idx,
stage="decode-norm").reshape(num_seqs, -1)
out, _ = self.out_proj(normed)
return _check_gdn_finite(
out, layer_idx=self.layer_idx, stage="decode-output")
# ---------------------------------------------------------------------------
# Full Attention (with gated q — unique to Qwen3.5)
# ---------------------------------------------------------------------------
class Qwen3_5AttentionHeadRMSNorm(GemmaRMSNorm):
def forward_cuda(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
):
if (_USE_COREX_ATTN_HEAD_RMS_NORM
and residual is None
and x.dtype == torch.float16
and self.weight.dtype == torch.float16
and x.dim() == 3
and x.shape[0] == 1
and x.shape[-1] == 256
and x.is_contiguous()
and self.weight.is_contiguous()):
original_shape = x.shape
converted, squares = _corex_attn_head_rms_norm.prepare(
x.view(-1, 256))
inverse = torch.rsqrt(
squares.mean(dim=-1, keepdim=True)
+ self.variance_epsilon)
return _corex_attn_head_rms_norm.apply_inverse(
converted, self.weight, inverse).view(original_shape)
return super().forward_cuda(x, residual)
class Qwen3_5FullAttention(nn.Module):
def __init__(
self,
text_cfg,
layer_idx: int,
cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
self.layer_idx = layer_idx
self.hidden_size = text_cfg.hidden_size # 5120
self.num_heads = text_cfg.num_attention_heads # 24
self.num_kv_heads = text_cfg.num_key_value_heads # 4
self.head_dim = text_cfg.head_dim # 256
self.rms_norm_eps = text_cfg.rms_norm_eps
tp_size = get_tensor_model_parallel_world_size()
self.local_num_heads = self.num_heads // tp_size
self.scaling = self.head_dim ** -0.5
self.use_packed_local_qgkv = tp_size > self.num_kv_heads
# When num_kv_heads < tp_size we cannot shard KV further (would give
# fractional heads per rank). Use ReplicatedLinear so every rank holds
# all KV heads; local_num_kv_heads equals the full count.
# When num_kv_heads >= tp_size standard ColumnParallel sharding applies.
if tp_size > self.num_kv_heads:
# GQA-aware TP sharding: ixformer kernel only supports num_kv_heads=1
# per rank. With num_kv_heads=2 < tp_size=4 we cannot shard KV
# evenly, but we CAN assign each rank the ONE KV head that serves
# its Q heads:
# q_per_kv = num_heads // num_kv_heads (e.g. 16//2 = 8)
# Rank r uses KV head r * local_num_heads // q_per_kv
# e.g. ranks 0,1 → KV head 0; ranks 2,3 → KV head 1.
# We replicate all KV heads to every rank and select in forward().
self.proj_kv_heads = self.num_kv_heads # heads available from projection
self.local_num_kv_heads = 1 # heads after rank-local selection
self.q_per_kv_global = self.num_heads // self.num_kv_heads
local_qg_dim = self.local_num_heads * self.head_dim * 2
replicated_kv_dim = self.num_kv_heads * self.head_dim
self.qgkv_proj = ReplicatedLinear(
self.hidden_size, local_qg_dim + 2 * replicated_kv_dim,
bias=False, quant_config=quant_config,
prefix=f"{prefix}.qgkv_proj")
else:
# Standard sharding: each rank gets num_kv_heads // tp_size heads.
self.local_num_kv_heads = self.num_kv_heads // tp_size
self.proj_kv_heads = self.local_num_kv_heads # already sharded
self.q_per_kv_global = None
self.k_proj = ColumnParallelLinear(
self.hidden_size, self.num_kv_heads * self.head_dim,
bias=False, quant_config=quant_config,
prefix=f"{prefix}.k_proj")
self.v_proj = ColumnParallelLinear(
self.hidden_size, self.num_kv_heads * self.head_dim,
bias=False, quant_config=quant_config,
prefix=f"{prefix}.v_proj")
self.local_q_dim = self.local_num_heads * self.head_dim
self.local_kv_dim = self.local_num_kv_heads * self.head_dim
if not self.use_packed_local_qgkv:
# q_proj includes gate: output = num_heads * head_dim * 2
self.q_proj = ColumnParallelLinear(
self.hidden_size, self.num_heads * self.head_dim * 2,
bias=False, quant_config=quant_config,
prefix=f"{prefix}.q_proj")
self.o_proj = RowParallelLinear(
self.num_heads * self.head_dim, self.hidden_size,
bias=False, quant_config=quant_config,
prefix=f"{prefix}.o_proj")
self.q_norm = Qwen3_5AttentionHeadRMSNorm(
self.head_dim, eps=self.rms_norm_eps)
self.k_norm = Qwen3_5AttentionHeadRMSNorm(
self.head_dim, eps=self.rms_norm_eps)
# Partial RoPE: rotary_dim = head_dim * partial_rotary_factor = 256 * 0.25 = 64
rope_params = getattr(text_cfg, "rope_parameters", {}) or {}
rope_theta = rope_params.get("rope_theta", 10_000_000)
partial_factor = rope_params.get("partial_rotary_factor", 0.25)
rotary_dim = int(self.head_dim * partial_factor)
self.rotary_emb = Qwen3_5InterleavedMRotaryEmbedding(
head_size=self.head_dim,
rotary_dim=rotary_dim,
max_position_embeddings=text_cfg.max_position_embeddings,
base=rope_theta,
is_neox_style=True,
dtype=torch.get_default_dtype(),
mrope_section=rope_params.get("mrope_section", [11, 11, 10]),
)
self.attn = Attention(
self.local_num_heads,
self.head_dim,
self.scaling,
num_kv_heads=self.local_num_kv_heads,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"{prefix}.attn",
)
def forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
kv_cache: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
total_tokens = hidden_states.shape[0]
with bi100_timer("full_attn.project_qgkv"):
if self.use_packed_local_qgkv:
projected, _ = self.qgkv_proj(hidden_states)
qg, k, v = torch.split(
projected,
[self.local_num_heads * self.head_dim * 2,
self.proj_kv_heads * self.head_dim,
self.proj_kv_heads * self.head_dim],
dim=-1)
else:
qg, _ = self.q_proj(hidden_states)
k, _ = self.k_proj(hidden_states)
v, _ = self.v_proj(hidden_states)
with bi100_timer("full_attn.norm_rope"):
# q projection output includes gate (dim doubled).
qg = qg.view(total_tokens, self.local_num_heads,
self.head_dim * 2)
q = qg[:, :, :self.head_dim].reshape(total_tokens, -1)
gate = qg[:, :, self.head_dim:].reshape(total_tokens, -1)
q = self.q_norm.forward_cuda(
q.view(total_tokens, self.local_num_heads, self.head_dim)
.contiguous()).view(total_tokens, -1)
# Select the one rank-local KV head before k_norm and RoPE.
if self.q_per_kv_global is not None:
tp_rank = get_tensor_model_parallel_rank()
kv_idx = ((tp_rank * self.local_num_heads)
// self.q_per_kv_global)
k = (k.view(total_tokens, self.proj_kv_heads, self.head_dim)
[:, kv_idx, :].contiguous())
v = (v.view(total_tokens, self.proj_kv_heads, self.head_dim)
[:, kv_idx, :].contiguous())
k = self.k_norm.forward_cuda(
k.view(total_tokens, self.local_num_kv_heads, self.head_dim)
.contiguous()).view(total_tokens, -1)
q, k = self.rotary_emb(positions, q, k)
with bi100_timer("full_attn.attention"):
with bi100_timer(f"L{self.layer_idx}.full_attn"):
attn_out = self.attn(q, k, v, kv_cache, attn_metadata)
with bi100_timer("full_attn.gate"):
attn_out = (attn_out
* torch.sigmoid(gate.float()).to(attn_out.dtype))
with bi100_timer("full_attn.output_proj"):
output, _ = self.o_proj(attn_out)
return output
# ---------------------------------------------------------------------------
# MLP (SwiGLU, same as Qwen2/Qwen3)
# ---------------------------------------------------------------------------
class Qwen3_5MLP(nn.Module):
def __init__(
self,
hidden_size: int,
intermediate_size: int,
hidden_act: str,
quant_config: Optional[QuantizationConfig] = None,
) -> None:
super().__init__()
self.gate_up_proj = MergedColumnParallelLinear(
hidden_size, [intermediate_size] * 2,
bias=False, quant_config=quant_config)
self.down_proj = RowParallelLinear(
intermediate_size, hidden_size,
bias=False, quant_config=quant_config)
if hidden_act != "silu":
raise ValueError(f"Unsupported activation: {hidden_act}")
self.act_fn = SiluAndMul()
def forward(self, x: torch.Tensor) -> torch.Tensor:
gate_up, _ = self.gate_up_proj(x)
x = self.act_fn(gate_up)
x, _ = self.down_proj(x)
return x
# ---------------------------------------------------------------------------
# MoE sparse block (Qwen3.5-MoE / Qwen3.6-35B-A3B)
# ---------------------------------------------------------------------------
class Qwen3_5MoeSparseBlock(nn.Module):
"""Replaces Qwen3_5MLP for qwen3_5_moe_text layers.
FusedMoE is used ONLY for weight storage and loading (create_weights /
weight_loader are pure PyTorch). Its forward kernel is bypassed because
ixformer on BI-V100 lacks vllm_moe_topk_softmax / vllm_invoke_fused_moe_kernel.
Routing and expert computation use a pure-PyTorch loop instead.
Shared expert uses RowParallelLinear(reduce_results=False) so both paths
produce partial (pre-all-reduce) outputs that are combined before a single
all-reduce.
"""
def __init__(
self,
text_cfg,
quant_config: Optional[QuantizationConfig] = None,
) -> None:
super().__init__()
hidden_size = text_cfg.hidden_size
self.num_experts = text_cfg.num_experts
self.top_k = text_cfg.num_experts_per_tok
# Router and scalar shared-expert gate read the same hidden state. Keep
# their checkpoint shards in one replicated weight so forward needs a
# single GEMM for 256 + 1 outputs.
self.router_shared_gate = ReplicatedLinear(
hidden_size, text_cfg.num_experts + 1,
bias=False, quant_config=quant_config)
self.router_shared_gate.weight.weight_loader = \
self._router_shared_gate_weight_loader
# FusedMoE: only used for weight storage + weight_loader.
# Forward is bypassed — see _pure_pytorch_experts().
self.experts = FusedMoE(
num_experts=text_cfg.num_experts,
top_k=text_cfg.num_experts_per_tok,
hidden_size=hidden_size,
intermediate_size=text_cfg.moe_intermediate_size,
reduce_results=False, # we do the all-reduce ourselves below
renormalize=True,
quant_config=quant_config,
)
# Shared expert: defer all-reduce to combine with routed output first
shared_size = text_cfg.shared_expert_intermediate_size
self.shared_expert_gate_up = MergedColumnParallelLinear(
hidden_size, [shared_size] * 2, bias=False,
quant_config=quant_config)
self.shared_expert_down = RowParallelLinear(
shared_size, hidden_size, bias=False, reduce_results=False,
quant_config=quant_config)
self.act_fn = SiluAndMul()
def _router_shared_gate_weight_loader(
self,
param: torch.Tensor,
loaded_weight: torch.Tensor,
shard_id: int,
) -> None:
if shard_id == 0:
offset = 0
rows = self.num_experts
elif shard_id == 1:
offset = self.num_experts
rows = 1
else:
raise ValueError(f"unexpected router/shared gate shard: {shard_id}")
expected = (rows, param.shape[1])
if tuple(loaded_weight.shape) != expected:
raise ValueError(
"unexpected router/shared gate weight shape: "
f"expected {expected}, got {tuple(loaded_weight.shape)}")
param.data.narrow(0, offset, rows).copy_(loaded_weight)
def _pure_pytorch_experts(
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
) -> torch.Tensor:
"""Pure-PyTorch MoE (ixformer has no MoE kernels on BI-V100).
w13_weight: (num_experts, 2*inter_per_partition, hidden) [TP-sharded]
w2_weight: (num_experts, hidden, inter_per_partition) [TP-sharded]
Output is partial (pre-all-reduce), same contract as FusedMoE
with reduce_results=False.
"""
# Fused topk+softmax: single CUB kernel vs 2 PyTorch ops.
# Source: xllm/core/kernels/cuda/moe/moe_topk_softmax_kernels.cuh
if _USE_COREX_MOE_TOPK_SOFTMAX:
topk_weights, topk_ids = _corex_moe_topk_softmax.moe_topk_softmax(
router_logits.float(), self.top_k, True)
topk_ids = topk_ids.to(torch.int64)
topk_weights = topk_weights.to(hidden_states.dtype)
else:
topk_logits, topk_ids = torch.topk(
router_logits.float(), self.top_k, dim=-1) # (T, top_k)
topk_weights = torch.softmax(topk_logits, dim=-1)
topk_weights = topk_weights.to(hidden_states.dtype)
w13 = self.experts.w13_weight # (E, 2*I, H)
w2 = self.experts.w2_weight # (E, H, I)
T = hidden_states.shape[0]
if T == 1:
# Fast path: single token (decode).
# Batched GEMM: replace top_k separate F.linear calls with 2 fused ops.
# gate_up: 1 large GEMM (1,H) × (K*2*I,H)^T → (1, K*2*I)
# down: 1 bmm (K,H,I) @ (K,I,1) → (K,H)
# Total: 3 kernel launches vs previous 16 (top_k*2).
eids = topk_ids[0] # (K,)
ws = topk_weights[0].to(hidden_states.dtype) # (K,)
use_corex_direct = (
_USE_COREX_MOE_DIRECT_ROUTED
and hidden_states.dtype == torch.float16
and w13.dtype == torch.float16
and w2.dtype == torch.float16
and ws.dtype == torch.float16
and hidden_states.is_cuda and w13.is_cuda and w2.is_cuda
and eids.is_cuda and ws.is_cuda
and hidden_states.is_contiguous()
and w13.is_contiguous() and w2.is_contiguous()
and eids.is_contiguous() and ws.is_contiguous()
and hidden_states.shape == (1, 2048)
and w13.shape == (256, 256, 2048)
and w2.shape == (256, 2048, 128)
and eids.shape == (8,) and ws.shape == (8,))
if use_corex_direct:
gate_up = _corex_moe_direct_routed.w13(
hidden_states, w13, eids)
act = self.act_fn(gate_up)
return _corex_moe_direct_routed.w2_reduce(
act, w2, eids, ws)
use_corex_gather = (
_USE_COREX_MOE_WEIGHT_GATHER
and hidden_states.dtype == torch.float16
and w13.dtype == torch.float16
and w2.dtype == torch.float16
and w13.is_cuda and w2.is_cuda and eids.is_cuda
and w13.is_contiguous() and w2.is_contiguous()
and eids.is_contiguous()
and w13.dim() == 3 and w2.dim() == 3
and eids.dim() == 1 and eids.numel() == 8
and w13.shape[0] == w2.shape[0]
and w13.shape[2] == w2.shape[1]
and w13.shape[1] == 2 * w2.shape[2]
and w13.shape[1] * w13.shape[2] % 8 == 0
and w2.shape[1] * w2.shape[2] % 8 == 0)
if use_corex_gather:
w13_sel, w2_sel = _corex_moe_weight_gather.gather(
w13, w2, eids)
else:
w13_sel = w13[eids] # (K, 2*I, H)
w2_sel = w2[eids] # (K, H, I)
H = hidden_states.shape[-1]
gate_up = F.linear(
hidden_states,
w13_sel.reshape(-1, H), # (K*2*I, H) — contiguous after indexing
) # (1, K*2*I)
gate_up = gate_up.view(self.top_k, -1) # (K, 2*I)
if _USE_FUSED_MOE_ACTIVATION:
act = self.act_fn(gate_up) # (K, I)
else:
gate, up = gate_up.chunk(2, dim=-1)
act = F.silu(gate) * up
# bmm: (K,H,I) @ (K,I,1) → (K,H,1) → (K,H)
expert_out = torch.bmm(w2_sel, act.unsqueeze(-1)).squeeze(-1) # (K, H)
if (_USE_COREX_MOE_EXACT_REDUCE
and expert_out.dtype == torch.float16
and ws.dtype == torch.float16
and expert_out.shape[0] == 8):
out = _corex_moe_exact_reduce.serial_float(expert_out, ws)
else:
out = (expert_out * ws.unsqueeze(-1)).sum(
0, keepdim=True).to(hidden_states.dtype) # (1, H)
else:
# General path (prefill / multi-seq): group assignments once.
out = torch.zeros_like(hidden_states)
flat_eids = topk_ids.reshape(-1)
if _USE_COREX_MOE_INDEX_COMBINE:
# Fused CUDA: histogram + prefix_sum + place (11.5x faster)
src_dst, dst_src, expert_sizes = \
_corex_moe_index_combine.moe_compute_index(
flat_eids, w13.shape[0])
sorted_tok_ids = torch.arange(
T, device=topk_ids.device
).repeat_interleave(self.top_k)[dst_src.long()]
sorted_weights = topk_weights.reshape(-1)[dst_src.long()]
expert_counts = expert_sizes.tolist()
else:
order = torch.argsort(flat_eids, stable=True)
sorted_tok_ids = torch.arange(
T, device=topk_ids.device
).repeat_interleave(self.top_k)[order]
sorted_weights = topk_weights.reshape(-1)[order]
expert_counts = torch.bincount(
flat_eids, minlength=w13.shape[0]).tolist()
start = 0
for eid, count in enumerate(expert_counts):
end = start + count
if count == 0:
start = end
continue
tok_ids = sorted_tok_ids[start:end]
tokens = hidden_states[tok_ids] # (n, H)
gate_up = F.linear(tokens, w13[eid]) # (n, 2*I)
gate, up = gate_up.chunk(2, dim=-1)
act = F.silu(gate) * up # (n, I)
expert_out = F.linear(act, w2[eid]) # (n, H)
weights = sorted_weights[start:end].unsqueeze(-1)
out.index_add_(0, tok_ids, (expert_out * weights).to(out.dtype))
start = end
return out # partial, all-reduce done in forward()
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
with bi100_timer("moe.router"):
router_and_shared_gate, _ = self.router_shared_gate(hidden_states)
router_logits = router_and_shared_gate[..., :self.num_experts]
gate_score = router_and_shared_gate[..., self.num_experts:]
with bi100_timer("moe.routed"):
routed_out = self._pure_pytorch_experts(hidden_states, router_logits)
with bi100_timer("moe.shared"):
gate_up, _ = self.shared_expert_gate_up(hidden_states)
shared_out = self.act_fn(gate_up)
shared_out, _ = self.shared_expert_down(shared_out)
shared_out = shared_out * torch.sigmoid(gate_score)
with bi100_timer("moe.combine"):
out = routed_out + shared_out
if self.experts.tp_size > 1:
with bi100_timer("moe.all_reduce"):
out = tensor_model_parallel_all_reduce(out)
return out
# ---------------------------------------------------------------------------
# Decoder layer (dispatches to GatedDeltaNet or Qwen3_5FullAttention)
# ---------------------------------------------------------------------------
class Qwen3_5DecoderLayer(nn.Module):
def __init__(
self,
text_cfg,
layer_idx: int,
layer_type: str,
cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
) -> None:
super().__init__()
self.layer_idx = layer_idx
self.layer_type = layer_type
self._diagnostic_trace_pending = (
os.getenv("BI100_DIAGNOSTIC_LAYER_TRACE") == "1")
self.input_layernorm = GemmaRMSNorm(text_cfg.hidden_size,
eps=text_cfg.rms_norm_eps)
self.post_attention_layernorm = GemmaRMSNorm(text_cfg.hidden_size,
eps=text_cfg.rms_norm_eps)
if layer_type == "linear_attention":
self.linear_attn = GatedDeltaNet(text_cfg, layer_idx,
quant_config=quant_config)
else:
self.self_attn = Qwen3_5FullAttention(
text_cfg, layer_idx,
cache_config=cache_config,
quant_config=quant_config,
prefix=f"layers.{layer_idx}.self_attn",
)
if getattr(text_cfg, 'model_type', '') == 'qwen3_5_moe_text':
self.mlp = Qwen3_5MoeSparseBlock(text_cfg, quant_config=quant_config)
else:
self.mlp = Qwen3_5MLP(
hidden_size=text_cfg.hidden_size,
intermediate_size=text_cfg.intermediate_size,
hidden_act=text_cfg.hidden_act,
quant_config=quant_config,
)
def forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
kv_cache: Optional[torch.Tensor],
attn_metadata: AttentionMetadata,
residual: Optional[torch.Tensor],
# Only for linear_attention layers:
conv_state: Optional[torch.Tensor] = None,
temporal_state: Optional[torch.Tensor] = None,
gdn_capture_offsets: Optional[Iterable[int]] = None,
gdn_segment_offsets: Optional[Iterable[int]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
with bi100_timer("layer.input_norm"):
if residual is None:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
else:
hidden_states, residual = self.input_layernorm(
hidden_states, residual)
if self.layer_type == "linear_attention":
with bi100_timer("layer.gdn"):
hidden_states = self.linear_attn(
hidden_states, attn_metadata, conv_state, temporal_state,
capture_offsets=gdn_capture_offsets,
segment_offsets=gdn_segment_offsets)
else:
with bi100_timer("layer.full_attn"):
hidden_states = self.self_attn(
positions, hidden_states, kv_cache, attn_metadata)
with bi100_timer("layer.post_attn_norm"):
hidden_states, residual = self.post_attention_layernorm(
hidden_states, residual)
with bi100_timer("layer.moe"):
hidden_states = self.mlp(hidden_states)
if self._diagnostic_trace_pending:
self._diagnostic_trace_pending = False
rank = os.getenv("RANK", os.getenv("LOCAL_RANK", "?"))
print(
"[BI100 DIAGNOSTIC] "
f"rank={rank} layer={self.layer_idx} "
f"attention={self.layer_type} "
f"mlp={type(self.mlp).__name__} stage=completed",
file=sys.stderr,
flush=True,
)
return hidden_states, residual
# ---------------------------------------------------------------------------
# Full transformer model
# ---------------------------------------------------------------------------
def _validate_qwen_kv_cache_count(configured_count, kv_caches):
if len(kv_caches) != configured_count:
raise RuntimeError(
"Qwen3.5 allocated KV cache count mismatch: "
f"configured {configured_count}, received {len(kv_caches)}")
class Qwen3_5Model(nn.Module):
def __init__(
self,
text_cfg,
cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
kv_cache_count: Optional[int] = None,
) -> None:
super().__init__()
self.text_cfg = text_cfg
full_attention_count = sum(
layer_type == "full_attention"
for layer_type in text_cfg.layer_types)
if kv_cache_count is None:
kv_cache_count = full_attention_count
if (not isinstance(kv_cache_count, int) or isinstance(kv_cache_count, bool)
or kv_cache_count < full_attention_count):
raise RuntimeError(
"Qwen3.5 configured KV cache count must cover every "
f"full-attention layer: configured {kv_cache_count}, "
f"required {full_attention_count}")
self.kv_cache_count = kv_cache_count
self.embed_tokens = VocabParallelEmbedding(
text_cfg.vocab_size, text_cfg.hidden_size)
self.layers = nn.ModuleList([
Qwen3_5DecoderLayer(
text_cfg, i, text_cfg.layer_types[i],
cache_config=cache_config, quant_config=quant_config)
for i in range(text_cfg.num_hidden_layers)
])
self.norm = GemmaRMSNorm(text_cfg.hidden_size, eps=text_cfg.rms_norm_eps)
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
kv_caches: List[torch.Tensor],
attn_metadata: AttentionMetadata,
conv_states: torch.Tensor, # (num_linear_layers, batch, ...)
temporal_states: torch.Tensor, # (num_linear_layers, batch, ...)
inputs_embeds: Optional[torch.Tensor] = None,
gdn_capture_offsets: Optional[Iterable[int]] = None,
gdn_segment_offsets: Optional[Iterable[int]] = None,
) -> torch.Tensor:
_validate_qwen_kv_cache_count(self.kv_cache_count, kv_caches)
with bi100_timer("model.embed"):
hidden_states = (self.embed_tokens(input_ids)
if inputs_embeds is None else inputs_embeds)
residual = None
attn_idx = 0
linear_idx = 0
capture_offsets = tuple(gdn_capture_offsets or ())
captured_conv_states: Dict[int, List[torch.Tensor]] = {
offset: [] for offset in capture_offsets
}
captured_temporal_states: Dict[int, List[torch.Tensor]] = {
offset: [] for offset in capture_offsets
}
for layer in self.layers:
if layer.layer_type == "linear_attention":
hidden_states, residual = layer(
positions, hidden_states,
kv_cache=None,
attn_metadata=attn_metadata,
residual=residual,
conv_state=conv_states[linear_idx],
temporal_state=temporal_states[linear_idx],
gdn_capture_offsets=capture_offsets,
gdn_segment_offsets=gdn_segment_offsets,
)
for offset in capture_offsets:
captured_conv_states[offset].append(
layer.linear_attn.captured_conv_states[offset])
captured_temporal_states[offset].append(
layer.linear_attn.captured_temporal_states[offset])
linear_idx += 1
else:
kv_cache = kv_caches[attn_idx]
hidden_states, residual = layer(
positions, hidden_states,
kv_cache=kv_cache,
attn_metadata=attn_metadata,
residual=residual,
)
attn_idx += 1
with bi100_timer("model.final_norm"):
hidden_states, _ = self.norm(hidden_states, residual)
self.captured_conv_states = {
offset: torch.stack(states)
for offset, states in captured_conv_states.items()
}
self.captured_temporal_states = {
offset: torch.stack(states)
for offset, states in captured_temporal_states.items()
}
return hidden_states
# ---------------------------------------------------------------------------
# Top-level CausalLM wrapper with MambaCacheManager
# ---------------------------------------------------------------------------
class Qwen3_5ForCausalLM(nn.Module, HasInnerState, SupportsLoRA,
SupportsMultiModal):
has_inner_state = True
supports_lora = True
packed_modules_mapping = {
"gate_up_proj": ["gate_proj", "up_proj"],
}
supported_lora_modules = [
"gate_up_proj",
"down_proj",
"o_proj",
]
embedding_modules = {}
embedding_padding_modules = []
def __init__(
self,
config, # Qwen3_5Config (top-level)
cache_config: Optional[CacheConfig] = None,
quant_config: Optional[QuantizationConfig] = None,
lora_config: Optional[LoRAConfig] = None,
scheduler_config: Optional[SchedulerConfig] = None,
multimodal_config: Optional[MultiModalConfig] = None,
prefix: str = "",
) -> None:
_bi100_model_trace("Qwen3_5ForCausalLM initialization begin")
super().__init__()
self.config = config
self.scheduler_config = scheduler_config
self.multimodal_config = multimodal_config
# The text config holds all architecture parameters
text_cfg = config.text_config
self.text_cfg = text_cfg
rope_parameters = getattr(text_cfg, "rope_parameters", {}) or {}
mrope_sections = rope_parameters.get("mrope_section", [11, 11, 10])
if getattr(config, "rope_scaling", None) is None:
config.rope_scaling = {
"type": "mrope",
"mrope_section": mrope_sections,
}
# Pre-compute counts
self.num_linear_layers = sum(
1 for lt in text_cfg.layer_types if lt == "linear_attention")
self.num_attn_layers = sum(
1 for lt in text_cfg.layer_types if lt == "full_attention")
layers_block_type = getattr(
config, "layers_block_type",
["attention"] * text_cfg.num_hidden_layers)
self.num_kv_cache_layers = sum(
layer_type == "attention" for layer_type in layers_block_type)
if self.num_kv_cache_layers < self.num_attn_layers:
raise RuntimeError(
"Qwen3.5 KV accounting provides fewer caches than "
f"full-attention layers: {self.num_kv_cache_layers} < "
f"{self.num_attn_layers}")
accounting_mode = getattr(
config, "bi100_hybrid_kv_accounting_mode", "legacy40")
accounting_env = os.getenv("BI100_HYBRID_KV_ACCOUNTING", "<unset>")
tp_rank = get_tensor_model_parallel_rank()
full_attention_ordinals = ",".join(
str(index) for index, layer_type in enumerate(text_cfg.layer_types)
if layer_type == "full_attention")
logger.info(
"[BI100] Qwen hybrid KV accounting; tp_rank=%d "
"env_mode=%s config_mode=%s "
"configured_kv_layers=%d full_attention_layers=%d "
"full_attention_ordinals=%s",
tp_rank,
accounting_env,
accounting_mode,
self.num_kv_cache_layers,
self.num_attn_layers,
full_attention_ordinals,
)
# DeltaNet state dimensions (per layer, per sequence, TP-sharded)
tp_size = get_tensor_model_parallel_world_size()
self.conv_dim = (text_cfg.linear_num_key_heads * text_cfg.linear_key_head_dim * 2
+ text_cfg.linear_num_value_heads * text_cfg.linear_value_head_dim)
self.num_v_heads = text_cfg.linear_num_value_heads
self.head_k_dim = text_cfg.linear_key_head_dim
self.head_v_dim = text_cfg.linear_value_head_dim
self.conv_kernel_size = text_cfg.linear_conv_kernel_dim
self.model = Qwen3_5Model(
text_cfg,
cache_config=cache_config,
quant_config=quant_config,
kv_cache_count=self.num_kv_cache_layers,
)
self.visual = Qwen3_5VisionTransformer(
config.vision_config,
quant_config=None,
)
self.lm_head = ParallelLMHead(
text_cfg.vocab_size, text_cfg.hidden_size,
quant_config=quant_config,
)
self.logits_processor = LogitsProcessor(text_cfg.vocab_size)
self.sampler = Sampler()
# Lazy initialised in first forward call
self.mamba_cache: Optional[MambaCacheManager] = None
# Scheduler-owned recurrent prefix states. Keys are stable chained
# content hashes, never recyclable physical KV block ids.
self._gdn_prefix_cache: Dict[
Tuple[int, bytes], Tuple[torch.Tensor, torch.Tensor]] = {}
self._block_size: int = (cache_config.block_size
if cache_config is not None else 16)
self._startup_forward_traced = False
_bi100_model_trace("Qwen3_5ForCausalLM initialization complete")
def _get_mamba_cache_shape(self):
tp_size = get_tensor_model_parallel_world_size()
# Each sequence's state is stored in float32
conv_state_shape = (self.conv_dim // tp_size, self.conv_kernel_size - 1)
temporal_state_shape = (
self.num_v_heads // tp_size, self.head_k_dim, self.head_v_dim)
return conv_state_shape, temporal_state_shape
@staticmethod
def _validate_and_reshape_mm_tensor(
mm_input: Union[torch.Tensor, List[torch.Tensor]],
name: str,
) -> torch.Tensor:
if isinstance(mm_input, list):
return torch.cat(mm_input)
if not isinstance(mm_input, torch.Tensor):
raise ValueError(f"incorrect type for {name}: {type(mm_input)}")
if mm_input.ndim == 2:
return mm_input
if mm_input.ndim == 3:
return torch.cat(list(mm_input))
raise ValueError(
f"{name} must be a 2D tensor or batched 3D tensor, got "
f"shape={tuple(mm_input.shape)}")
def _parse_and_validate_image_input(
self,
**kwargs: object,
) -> Optional[Qwen3_5ImageInputs]:
pixel_values = kwargs.get("pixel_values")
image_embeds = kwargs.get("image_embeds")
image_grid_thw = kwargs.get("image_grid_thw")
if pixel_values is None and image_embeds is None:
return None
if pixel_values is not None:
if image_grid_thw is None:
raise ValueError("image_grid_thw is required with pixel_values")
return Qwen3_5ImagePixelInputs(
type="pixel_values",
data=self._validate_and_reshape_mm_tensor(
pixel_values, "image pixel values"),
image_grid_thw=self._validate_and_reshape_mm_tensor(
image_grid_thw, "image grid_thw"),
)
return Qwen3_5ImageEmbeddingInputs(
type="image_embeds",
data=self._validate_and_reshape_mm_tensor(
image_embeds, "image embeddings"),
)
def _process_image_input(
self,
image_input: Qwen3_5ImageInputs,
) -> torch.Tensor:
if image_input["type"] == "image_embeds":
return image_input["data"].to(dtype=self.visual.dtype,
device=self.visual.device)
return self.visual(
image_input["data"],
grid_thw=image_input["image_grid_thw"],
)
@bi100_profile_transaction
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
kv_caches: List[torch.Tensor],
attn_metadata: AttentionMetadata,
intermediate_tensors: Optional[IntermediateTensors] = None,
**kwargs,
) -> torch.Tensor:
if not self._startup_forward_traced:
self._startup_forward_traced = True
_bi100_model_trace("first model forward entered")
if self.mamba_cache is None:
if self.scheduler_config is not None:
max_batch_size = _get_graph_batch_size(
self.scheduler_config.max_num_seqs)
else:
max_batch_size = max(_BATCH_SIZES_TO_CAPTURE) + 2
self.mamba_cache = MambaCacheManager(
torch.float32,
self.num_linear_layers,
max_batch_size,
*self._get_mamba_cache_shape(),
)
gdn_restore_key = kwargs.pop("gdn_restore_key", None)
gdn_capture_points = kwargs.pop("gdn_capture_points", None) or []
gdn_evict_keys = kwargs.pop("gdn_evict_keys", None) or []
gdn_segment_offsets = kwargs.pop("gdn_segment_offsets", None) or []
mamba_tensors = self.mamba_cache.current_run_tensors(
input_ids, attn_metadata, **kwargs)
# conv_states: (num_linear_layers, batch, local_conv_dim, kernel-1)
# temporal_states: (num_linear_layers, batch, local_num_v, k_dim, v_dim)
conv_states, temporal_states = mamba_tensors
_is_single_seq_prefill = (
attn_metadata is not None
and attn_metadata.num_prefill_tokens > 0
and conv_states.shape[1] == 1 # batch == 1
and getattr(attn_metadata, 'context_lens_tensor', None) is not None
)
has_gdn_actions = (gdn_restore_key is not None
or bool(gdn_capture_points)
or bool(gdn_evict_keys)
or bool(gdn_segment_offsets))
if has_gdn_actions and not _is_single_seq_prefill:
raise RuntimeError(
"GDN prefix-cache actions require a single-sequence prefill")
for evict_key in gdn_evict_keys:
self._gdn_prefix_cache.pop(_validate_gdn_prefix_key(evict_key),
None)
if gdn_restore_key is not None:
restore_key = _validate_gdn_prefix_key(gdn_restore_key)
saved_state = self._gdn_prefix_cache.get(restore_key)
if saved_state is None:
raise RuntimeError(
"scheduler requested a missing GDN prefix state: "
f"blocks={restore_key[0]} digest={restore_key[1].hex()}")
saved_conv, saved_temporal = saved_state
with bi100_timer("gdn_prefix.restore"):
conv_states[:, 0].copy_(
saved_conv.to(device=conv_states.device,
dtype=conv_states.dtype),
non_blocking=True)
temporal_states[:, 0].copy_(
saved_temporal.to(device=temporal_states.device,
dtype=temporal_states.dtype),
non_blocking=True)
query_len = (int(attn_metadata.num_prefill_tokens)
if _is_single_seq_prefill else 0)
capture_keys: Dict[int, Tuple[int, bytes]] = {}
for capture_point in gdn_capture_points:
if not isinstance(capture_point, tuple) or len(capture_point) != 2:
raise RuntimeError(
f"invalid GDN capture point: {capture_point!r}")
offset, capture_key = capture_point
if (not isinstance(offset, int) or offset <= 0
or offset > query_len or offset in capture_keys):
raise RuntimeError(
f"invalid GDN capture offset: {offset!r} "
f"for query_len={query_len}")
capture_keys[offset] = _validate_gdn_prefix_key(capture_key)
if len(capture_keys) > 2:
raise RuntimeError("at most two GDN capture points are supported")
interior_capture_offsets = tuple(
offset for offset in capture_keys if offset < query_len)
segment_offsets = set()
for offset in gdn_segment_offsets:
if (not isinstance(offset, int) or offset <= 0
or offset >= query_len):
raise RuntimeError(
f"invalid GDN segment offset: {offset!r} "
f"for query_len={query_len}")
segment_offsets.add(offset)
if len(segment_offsets) > 128:
raise RuntimeError("at most 128 GDN segment offsets are supported")
interior_segment_offsets = tuple(sorted(segment_offsets))
inputs_embeds = None
image_input = self._parse_and_validate_image_input(**kwargs)
if image_input is not None:
image_mask = input_ids == self.config.image_token_id
num_placeholders = int(image_mask.sum().item())
if num_placeholders:
inputs_embeds = self.model.embed_tokens(input_ids)
image_embeds = self._process_image_input(image_input)
if num_placeholders > image_embeds.shape[0]:
raise ValueError(
f"image token count ({num_placeholders}) exceeds "
f"vision embeddings ({image_embeds.shape[0]})")
# Prefix caching can consume the leading image tokens while
# vLLM 0.6 still supplies the full pixel tensor. The query's
# remaining placeholders always form a suffix of the flattened
# visual token stream.
image_embeds = image_embeds[-num_placeholders:]
inputs_embeds[image_mask, :] = image_embeds.to(
inputs_embeds.dtype)
with bi100_timer("model.forward"):
hidden_states = self.model(
input_ids, positions, kv_caches, attn_metadata,
conv_states, temporal_states,
inputs_embeds=inputs_embeds,
gdn_capture_offsets=interior_capture_offsets,
gdn_segment_offsets=interior_segment_offsets)
for offset, capture_key in capture_keys.items():
if offset == query_len:
captured_conv = conv_states[:, 0]
captured_temporal = temporal_states[:, 0]
else:
captured_conv = self.model.captured_conv_states[offset]
captured_temporal = self.model.captured_temporal_states[offset]
with bi100_timer("gdn_prefix.save"):
self._gdn_prefix_cache[capture_key] = (
captured_conv.detach().cpu().clone(),
captured_temporal.detach().cpu().clone(),
)
if bi100_profile_event_enabled():
profile_prefill_tokens = int(
getattr(attn_metadata, "num_prefill_tokens", 0) or 0)
profile_decode_tokens = int(
getattr(attn_metadata, "num_decode_tokens", 0) or 0)
profile_context_len = 0
if profile_prefill_tokens > 0:
profile_seq_lens = getattr(attn_metadata, "seq_lens", None)
if (not isinstance(profile_seq_lens, list)
or len(profile_seq_lens) != 1
or not isinstance(profile_seq_lens[0], int)):
raise RuntimeError(
"BI100 profile requires one host-visible prefill "
"sequence length")
profile_context_len = (
profile_seq_lens[0] - profile_prefill_tokens)
if profile_context_len < 0:
raise RuntimeError(
"BI100 profile observed a negative prefill context")
bi100_profile_flush(
tp_rank=get_tensor_model_parallel_rank(),
phase=("prefill" if profile_prefill_tokens > 0 else "decode"),
prefill_tokens=profile_prefill_tokens,
decode_tokens=profile_decode_tokens,
context_len=profile_context_len,
gdn_restore=bool(gdn_restore_key is not None),
gdn_capture_points=len(gdn_capture_points),
gdn_evict_keys=len(gdn_evict_keys),
)
return hidden_states
def compute_logits(
self,
hidden_states: torch.Tensor,
sampling_metadata: SamplingMetadata,
) -> Optional[torch.Tensor]:
# All TP ranks must call logits_processor to participate in the NCCL
# gather inside lm_head. Non-driver ranks return None after the gather.
# With chunked prefill, intermediate chunks have seq_groups=None on all
# ranks; _apply_logits_processors is guarded against this in
# logits_processor.py (patched by patch_xformers_sdpa_seq.py).
logits = self.logits_processor(self.lm_head, hidden_states,
sampling_metadata)
return logits
def sample(
self,
logits: torch.Tensor,
sampling_metadata: SamplingMetadata,
) -> Optional[SamplerOutput]:
return self.sampler(logits, sampling_metadata)
def copy_inputs_before_cuda_graphs(self, input_buffers, **kwargs):
return self.mamba_cache.copy_inputs_before_cuda_graphs(
input_buffers, **kwargs)
def get_seqlen_agnostic_capture_inputs(self, batch_size: int):
return self.mamba_cache.get_seqlen_agnostic_capture_inputs(batch_size)
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
_bi100_model_trace("dense load_weights begin")
loaded_count = 0
stacked_params_mapping = [
# (param_name, weight_name, shard_id)
("gate_up_proj", "gate_proj", 0),
("gate_up_proj", "up_proj", 1),
]
params_dict = dict(self.named_parameters())
for name, loaded_weight in weights:
loaded_count += 1
# Skip vision and MTP branches
if (name.startswith("model.visual")
or name.startswith("mtp.")
or name.startswith("model.mtp")):
continue
# Prefix remapping: checkpoint may wrap under language_model
if name.startswith("model.language_model."):
name = "model." + name[len("model.language_model."):]
# Skip positional embedding caches
if "rotary_emb.inv_freq" in name:
continue
if _load_full_attention_qgkv_weight(
params_dict, name, loaded_weight, self.text_cfg):
continue
if _load_gdn_projection_weight(
params_dict, name, loaded_weight, self.text_cfg):
continue
# Remap conv1d.weight → conv1d_weight
# The conv has depth (1) dim in the checkpoint that we handle separately
if ".linear_attn.conv1d.weight" in name:
name = name.replace(".linear_attn.conv1d.weight",
".linear_attn.conv1d_weight")
# Stacked param loading (gate_up_proj)
for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
if name.endswith(".bias") and name not in params_dict:
break
if name not in params_dict:
break
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
break
else:
if name.endswith(".bias") and name not in params_dict:
continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
weight_loader(param, loaded_weight)
_bi100_model_trace(f"dense load_weights complete items={loaded_count}")
# ---------------------------------------------------------------------------
# Qwen3.6-35B-A3B (Qwen3_5-MoE architecture)
# ---------------------------------------------------------------------------
@MULTIMODAL_REGISTRY.register_image_input_mapper(qwen36_image_input_mapper)
@MULTIMODAL_REGISTRY.register_max_image_tokens(get_max_qwen36_image_tokens)
@INPUT_REGISTRY.register_dummy_data(dummy_data_for_qwen36)
@INPUT_REGISTRY.register_input_processor(input_processor_for_qwen36)
class Qwen3_5MoeForCausalLM(Qwen3_5ForCausalLM):
"""Qwen3.6-35B-A3B: same hybrid-attention backbone as 27B, dense MLP
replaced by Qwen3_5MoeSparseBlock (256 routed experts + shared expert).
Only load_weights differs from the dense variant.
"""
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
_bi100_model_trace("MoE load_weights begin")
loaded_count = 0
vision_loaded_count = 0
# Checkpoint key format for this model (transformers Qwen3_5MoeExperts):
# mlp.experts.gate_up_proj shape (num_experts, 2*intermediate, hidden)
# mlp.experts.down_proj shape (num_experts, hidden, intermediate)
# mlp.gate.weight shape (num_experts, hidden) [router]
# mlp.shared_expert_gate.weight shape (1, hidden)
# mlp.shared_expert.{gate,up,down}_proj.weight [shared MLP]
# Our FusedMoE stores:
# mlp.experts.w13_weight shape (num_experts, 2*intermediate//tp, hidden)
# mlp.experts.w2_weight shape (num_experts, hidden, intermediate//tp)
# Our router/shared gate stores both tensors in one (num_experts+1, H)
# replicated weight. Our shared expert stores:
# mlp.shared_expert_gate_up.weight (merged gate+up)
# mlp.shared_expert_down.weight
stacked_params_mapping = [
# (param_name, weight_name, shard_id)
# shared expert
("shared_expert_gate_up", "shared_expert.gate_proj", 0),
("shared_expert_gate_up", "shared_expert.up_proj", 1),
# linear_attention dense proj (same as 27B)
("gate_up_proj", "gate_proj", 0),
("gate_up_proj", "up_proj", 1),
]
params_dict = dict(self.named_parameters())
for name, loaded_weight in weights:
loaded_count += 1
if name.startswith("model.visual."):
name = "visual." + name[len("model.visual."):]
if "attn.qkv.weight" in name:
num_heads = self.config.vision_config.num_heads
hidden_size = self.config.vision_config.hidden_size
head_size = hidden_size // num_heads
loaded_weight = loaded_weight.view(
3, num_heads, head_size, hidden_size)
loaded_weight = loaded_weight.transpose(0, 1).reshape(
-1, hidden_size)
elif "attn.qkv.bias" in name:
num_heads = self.config.vision_config.num_heads
hidden_size = self.config.vision_config.hidden_size
head_size = hidden_size // num_heads
loaded_weight = loaded_weight.view(
3, num_heads, head_size)
loaded_weight = loaded_weight.transpose(0, 1).reshape(-1)
if name not in params_dict:
raise ValueError(f"unexpected Qwen3.6 vision weight: {name}")
param = params_dict[name]
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
weight_loader(param, loaded_weight)
vision_loaded_count += 1
continue
# MTP is not used by the fixed evaluator command.
if (name.startswith("mtp.")
or name.startswith("model.mtp")):
continue
# Prefix remapping for VL checkpoint (Qwen3_5MoeForConditionalGeneration):
# model.language_model.model.{layers,embed_tokens,norm} -> model.{...}
# model.language_model.lm_head -> lm_head
# Prefix remapping: checkpoint may wrap under language_model
if name.startswith("model.language_model."):
name = "model." + name[len("model.language_model."):]
if "rotary_emb.inv_freq" in name:
continue
if _load_full_attention_qgkv_weight(
params_dict, name, loaded_weight, self.text_cfg):
continue
if _load_gdn_projection_weight(
params_dict, name, loaded_weight, self.text_cfg):
continue
if name.endswith(".mlp.gate.weight"):
fused_name = name[:-len("gate.weight")] \
+ "router_shared_gate.weight"
if fused_name not in params_dict:
raise ValueError(
f"missing fused router/shared gate: {fused_name}")
params_dict[fused_name].weight_loader(
params_dict[fused_name], loaded_weight, 0)
continue
if name.endswith(".mlp.shared_expert_gate.weight"):
fused_name = name[:-len("shared_expert_gate.weight")] \
+ "router_shared_gate.weight"
if fused_name not in params_dict:
raise ValueError(
f"missing fused router/shared gate: {fused_name}")
params_dict[fused_name].weight_loader(
params_dict[fused_name], loaded_weight, 1)
continue
if ".linear_attn.conv1d.weight" in name:
name = name.replace(".linear_attn.conv1d.weight",
".linear_attn.conv1d_weight")
# --- Fused routed-expert weights (all experts in one tensor) ---
if "mlp.experts.gate_up_proj" in name:
# loaded_weight: (num_experts, 2*intermediate, hidden)
w13_name = name.replace("mlp.experts.gate_up_proj",
"mlp.experts.w13_weight")
if w13_name not in params_dict:
continue
param = params_dict[w13_name]
n_exp = loaded_weight.shape[0]
inter = loaded_weight.shape[1] // 2
gate_w = loaded_weight[:, :inter, :].contiguous()
up_w = loaded_weight[:, inter:, :].contiguous()
for eid in range(n_exp):
param.weight_loader(param, gate_w[eid], "w1_weight", "w1", eid)
param.weight_loader(param, up_w[eid], "w3_weight", "w3", eid)
continue
if "mlp.experts.down_proj" in name:
# loaded_weight: (num_experts, hidden, intermediate)
w2_name = name.replace("mlp.experts.down_proj",
"mlp.experts.w2_weight")
if w2_name not in params_dict:
continue
param = params_dict[w2_name]
n_exp = loaded_weight.shape[0]
for eid in range(n_exp):
param.weight_loader(param, loaded_weight[eid], "w2_weight", "w2", eid)
continue
# --- Shared expert down_proj rename ---
if "mlp.shared_expert.down_proj" in name:
name = name.replace("mlp.shared_expert.down_proj",
"mlp.shared_expert_down")
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
continue
# --- Individual expert weights (FT checkpoint: experts.{i}.{proj}.weight) ---
# Standard transformers fine-tuning saves each expert separately instead of
# the pre-merged (num_experts, ...) tensors in the original checkpoint.
if ".mlp.experts." in name:
parts = name.split(".mlp.experts.", 1)
expert_rest = parts[1] # e.g. "0.gate_proj.weight"
dot_pos = expert_rest.find(".")
if dot_pos > 0 and expert_rest[:dot_pos].isdigit():
eid = int(expert_rest[:dot_pos])
proj_raw = expert_rest[dot_pos + 1:]
proj = proj_raw[:-7] if proj_raw.endswith(".weight") else proj_raw
prefix = parts[0] # e.g. "model.layers.0"
if proj == "gate_proj":
w13_name = f"{prefix}.mlp.experts.w13_weight"
if w13_name in params_dict:
param = params_dict[w13_name]
param.weight_loader(param, loaded_weight, "w1_weight", "w1", eid)
elif proj == "up_proj":
w13_name = f"{prefix}.mlp.experts.w13_weight"
if w13_name in params_dict:
param = params_dict[w13_name]
param.weight_loader(param, loaded_weight, "w3_weight", "w3", eid)
elif proj == "down_proj":
w2_name = f"{prefix}.mlp.experts.w2_weight"
if w2_name in params_dict:
param = params_dict[w2_name]
param.weight_loader(param, loaded_weight, "w2_weight", "w2", eid)
continue
# --- Stacked / standard weights ---
for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
if name not in params_dict:
break
param = params_dict[name]
param.weight_loader(param, loaded_weight, shard_id)
break
else:
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
_bi100_model_trace(
f"MoE load_weights complete items={loaded_count} "
f"vision_items={vision_loaded_count}")