feat(CRITICAL): 从 GitHub 扫描搬运 ixformer SDK + xllm 完整 GDN/MoE 代码

来源:
  1. Chranos/ixformer (GitHub) → ixformer_sdk/ (230 files, 70K lines)
     - inference/functions/vllm.py: vllm_moe_topk_softmax 完整实现 (2033 lines)
     - inference/functions/moe.py: MoE ops 完整实现 (1380 lines)
     - contrib/vllm_flash_attn/: FA2 Python 接口 (1018 lines)
     - contrib/tgi/fused_moe.py: TGI fused MoE (429 lines)
     - csrc/include/ixformer/: C++ kernel headers + cmake

  2. Deep-Spark/xllm (GitHub) → upstream_ref/xllm_latest/ (+15 files)
     - npu_torch/qwen3_5_decoder_layer_impl.cpp/.h
     - npu_torch/qwen3_5_gated_delta_net.cpp/.h
     - npu_torch/qwen3_next_*.cpp/.h (6 files)
     - npu_torch/attention.cpp/.h + fused_moe.cpp/.h + CMakeLists.txt
     - models/llm/qwen3_5.h + qwen3_5_mtp.h + qwen3_next.h
     - models/vlm/qwen3_5.h

调用链完整性:
  ixformer_sdk/inference/functions/vllm.py
    → ops.infer.moe_topk_softmax() (C++ 层)
    → 这就是 base 镜像 libixformer.so 里的实现

  upstream_ref/xllm_latest/core/layers/ilu/fused_moe.cpp
    → ixformer::infer::topk_softmax() (直接 C++ 调用)
    → ixformer::infer::group_gemm() → 完整 7-step MoE pipeline
This commit is contained in:
project6-dev
2026-08-11 02:31:56 +00:00
parent a8b16da5da
commit 87a19d2d00
250 changed files with 76690 additions and 0 deletions

View File

@@ -0,0 +1,218 @@
/* Copyright 2025 The xLLM Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <cstdint>
#include <string>
#include <unordered_set>
#include <vector>
#include "core/layers/npu_torch/qwen3_5_decoder_layer_impl.h"
#include "models/model_registry.h"
#include "qwen3_next.h"
namespace xllm {
class Qwen3_5ModelImpl : public Qwen3NextModelImpl {
public:
explicit Qwen3_5ModelImpl(const ModelContext& context)
: Qwen3NextModelImpl(context, /*init_decoder_layers=*/false) {
const int32_t n_layers = context.get_model_args().n_layers();
for (int32_t layer_id = 0; layer_id < n_layers; ++layer_id) {
add_decoder_layer(
std::make_shared<layer::Qwen3_5DecoderLayerImpl>(context, layer_id));
}
}
};
TORCH_MODULE(Qwen3_5Model);
class Qwen3_5ForCausalLMImpl : public Qwen3NextForCausalLMImpl {
public:
explicit Qwen3_5ForCausalLMImpl(const ModelContext& context)
: Qwen3NextForCausalLMImpl(context, /*init_model=*/false) {
set_model_module(std::make_shared<Qwen3_5ModelImpl>(context));
}
};
TORCH_MODULE(Qwen3_5ForCausalLM);
#define LOAD_ARG_TEXT_OR_ROOT(arg_name, json_key, default_value) \
LOAD_ARG_OR(arg_name, "text_config." json_key, default_value); \
LOAD_ARG_OR(arg_name, json_key, args->arg_name())
#define LOAD_ARG_TEXT_OR_ROOT_CHAIN(arg_name, json_key, default_value) \
LOAD_ARG_TEXT_OR_ROOT(arg_name, json_key, default_value)
#define LOAD_QWEN3_5_ROPE_ARG(arg_name, default_value) \
LOAD_ARG_OR(arg_name, "text_config." #arg_name, default_value); \
LOAD_ARG_OR(arg_name, #arg_name, args->arg_name()); \
LOAD_ARG_OR( \
arg_name, "text_config.rope_scaling." #arg_name, args->arg_name()); \
LOAD_ARG_OR(arg_name, "rope_scaling." #arg_name, args->arg_name()); \
LOAD_ARG_OR( \
arg_name, "text_config.rope_parameters." #arg_name, args->arg_name()); \
LOAD_ARG_OR(arg_name, "rope_parameters." #arg_name, args->arg_name())
#define LOAD_QWEN3_5_NEXT_COMPAT_ARGS(default_moe_intermediate_size, \
default_num_experts, \
default_num_experts_per_tok, \
default_shared_expert_intermediate_size) \
LOAD_ARG_TEXT_OR_ROOT(attention_bias, "attention_bias", false); \
LOAD_ARG_TEXT_OR_ROOT(attention_dropout, "attention_dropout", 0.0f); \
LOAD_ARG_TEXT_OR_ROOT(bos_token_id, "bos_token_id", 151643); \
LOAD_ARG_TEXT_OR_ROOT(decoder_sparse_step, "decoder_sparse_step", 1); \
LOAD_ARG_TEXT_OR_ROOT(eos_token_id, "eos_token_id", 151645); \
LOAD_ARG_TEXT_OR_ROOT(head_dim, "head_dim", 256); \
LOAD_ARG_TEXT_OR_ROOT(hidden_act, "hidden_act", "silu"); \
LOAD_ARG_TEXT_OR_ROOT(hidden_size, "hidden_size", 2048); \
LOAD_ARG_TEXT_OR_ROOT(initializer_range, "initializer_range", 0.02f); \
LOAD_ARG_TEXT_OR_ROOT(intermediate_size, "intermediate_size", 5120); \
LOAD_ARG_TEXT_OR_ROOT( \
max_position_embeddings, "max_position_embeddings", 262144); \
LOAD_ARG_TEXT_OR_ROOT(max_window_layers, "max_window_layers", 28); \
LOAD_ARG_TEXT_OR_ROOT(moe_intermediate_size, \
"moe_intermediate_size", \
default_moe_intermediate_size); \
LOAD_ARG_TEXT_OR_ROOT(norm_topk_prob, "norm_topk_prob", true); \
LOAD_ARG_TEXT_OR_ROOT(n_heads, "num_attention_heads", 16); \
LOAD_ARG_TEXT_OR_ROOT(num_experts, "num_experts", default_num_experts); \
LOAD_ARG_TEXT_OR_ROOT(num_experts_per_tok, \
"num_experts_per_tok", \
default_num_experts_per_tok); \
LOAD_ARG_TEXT_OR_ROOT(n_layers, "num_hidden_layers", 48); \
LOAD_ARG_OR(n_kv_heads, "text_config.num_key_value_heads", 2); \
LOAD_ARG_OR( \
n_kv_heads, "num_key_value_heads", args->n_kv_heads().value_or(2)); \
LOAD_ARG_TEXT_OR_ROOT(output_router_logits, "output_router_logits", false); \
LOAD_ARG_TEXT_OR_ROOT(rms_norm_eps, "rms_norm_eps", 1e-6); \
LOAD_QWEN3_5_ROPE_ARG(rope_theta, 10000000.0f); \
LOAD_ARG_TEXT_OR_ROOT(router_aux_loss_coef, "router_aux_loss_coef", 0.001f); \
LOAD_ARG_TEXT_OR_ROOT(use_sliding_window, "use_sliding_window", false); \
LOAD_ARG_TEXT_OR_ROOT(sliding_window, "sliding_window", 4096); \
LOAD_ARG_TEXT_OR_ROOT(tie_word_embeddings, "tie_word_embeddings", false); \
LOAD_ARG_TEXT_OR_ROOT(vocab_size, "vocab_size", 151936); \
LOAD_ARG_TEXT_OR_ROOT( \
mlp_only_layers, "mlp_only_layers", std::vector<int32_t>()); \
LOAD_ARG_TEXT_OR_ROOT(attn_output_gate, "attn_output_gate", true); \
LOAD_ARG_TEXT_OR_ROOT( \
full_attention_interval, "full_attention_interval", 4); \
LOAD_ARG_TEXT_OR_ROOT(linear_conv_kernel_dim, "linear_conv_kernel_dim", 4); \
LOAD_ARG_TEXT_OR_ROOT(linear_key_head_dim, "linear_key_head_dim", 128); \
LOAD_ARG_TEXT_OR_ROOT(linear_num_key_heads, "linear_num_key_heads", 16); \
LOAD_ARG_TEXT_OR_ROOT(linear_num_value_heads, "linear_num_value_heads", 32); \
LOAD_ARG_TEXT_OR_ROOT(linear_value_head_dim, "linear_value_head_dim", 128); \
LOAD_QWEN3_5_ROPE_ARG(partial_rotary_factor, 0.25f); \
LOAD_ARG_OR(rope_scaling_mrope_section, \
"text_config.rope_scaling.mrope_section", \
std::vector<int64_t>()); \
LOAD_ARG_OR(rope_scaling_mrope_section, \
"text_config.rope_parameters.mrope_section", \
args->rope_scaling_mrope_section()); \
LOAD_ARG_OR(rope_scaling_mrope_section, \
"rope_parameters.mrope_section", \
args->rope_scaling_mrope_section()); \
LOAD_ARG_OR(rope_scaling_mrope_interleaved, \
"text_config.rope_scaling.mrope_interleaved", \
false); \
LOAD_ARG_OR(rope_scaling_mrope_interleaved, \
"text_config.rope_parameters.mrope_interleaved", \
args->rope_scaling_mrope_interleaved()); \
LOAD_ARG_OR(rope_scaling_mrope_interleaved, \
"rope_parameters.mrope_interleaved", \
args->rope_scaling_mrope_interleaved()); \
LOAD_ARG_TEXT_OR_ROOT(shared_expert_intermediate_size, \
"shared_expert_intermediate_size", \
default_shared_expert_intermediate_size); \
LOAD_ARG_OR( \
num_nextn_predict_layers, "text_config.mtp_num_hidden_layers", 0); \
LOAD_ARG_OR(num_nextn_predict_layers, \
"mtp_num_hidden_layers", \
args->num_nextn_predict_layers()); \
LOAD_ARG_OR(num_nextn_predict_layers, \
"text_config.num_nextn_predict_layers", \
args->num_nextn_predict_layers()); \
LOAD_ARG_OR(num_nextn_predict_layers, \
"num_nextn_predict_layers", \
args->num_nextn_predict_layers()); \
LOAD_ARG_OR( \
layer_types, "text_config.layer_types", std::vector<std::string>()); \
LOAD_ARG_OR(layer_types, "layer_types", args->layer_types()); \
LOAD_ARG_OR( \
layer_types, "text_config.layers_block_type", args->layer_types()); \
LOAD_ARG_OR(layer_types, "layers_block_type", args->layer_types()); \
LOAD_ARG_OR( \
n_routed_experts, "text_config.n_routed_experts", args->num_experts()); \
LOAD_ARG_OR(n_routed_experts, "n_routed_experts", args->num_experts()); \
SET_ARG(n_shared_experts, \
args->shared_expert_intermediate_size() > 0 ? 1 : 0); \
SET_ARG(scoring_func, "softmax"); \
SET_ARG(topk_method, ""); \
SET_ARG(n_group, -1); \
SET_ARG(topk_group, 0); \
SET_ARG(routed_scaling_factor, 1.0f); \
SET_ARG(stop_token_ids, \
std::unordered_set<int32_t>({args->eos_token_id()})); \
LOAD_ARG_TEXT_OR_ROOT(mamba_ssm_dtype, "mamba_ssm_dtype", "float32")
#define LOAD_QWEN3_5_TYPE_AND_DTYPE(default_model_type) \
LOAD_ARG_OR(model_type, "model_type", default_model_type); \
LOAD_ARG_OR(dtype, "text_config.dtype", "bfloat16"); \
LOAD_ARG_OR(dtype, "dtype", args->dtype()); \
LOAD_ARG_OR(dtype, "text_config.torch_dtype", args->dtype()); \
LOAD_ARG_OR(dtype, "torch_dtype", args->dtype())
REGISTER_CAUSAL_MODEL(qwen3_5, Qwen3_5ForCausalLM);
REGISTER_MODEL_ARGS(qwen3_5, [&] {
LOAD_QWEN3_5_TYPE_AND_DTYPE("qwen3_5");
LOAD_QWEN3_5_NEXT_COMPAT_ARGS(/*moe_intermediate_size=*/0,
/*num_experts=*/0,
/*num_experts_per_tok=*/0,
/*shared_expert_intermediate_size=*/0);
});
REGISTER_CAUSAL_MODEL(qwen3_5_text, Qwen3_5ForCausalLM);
REGISTER_MODEL_ARGS(qwen3_5_text, [&] {
LOAD_QWEN3_5_TYPE_AND_DTYPE("qwen3_5_text");
LOAD_QWEN3_5_NEXT_COMPAT_ARGS(/*moe_intermediate_size=*/0,
/*num_experts=*/0,
/*num_experts_per_tok=*/0,
/*shared_expert_intermediate_size=*/0);
});
REGISTER_CAUSAL_MODEL(qwen3_5_moe, Qwen3_5ForCausalLM);
REGISTER_MODEL_ARGS(qwen3_5_moe, [&] {
LOAD_QWEN3_5_TYPE_AND_DTYPE("qwen3_5_moe");
LOAD_QWEN3_5_NEXT_COMPAT_ARGS(/*moe_intermediate_size=*/512,
/*num_experts=*/512,
/*num_experts_per_tok=*/10,
/*shared_expert_intermediate_size=*/512);
});
REGISTER_CAUSAL_MODEL(qwen3_5_moe_text, Qwen3_5ForCausalLM);
REGISTER_MODEL_ARGS(qwen3_5_moe_text, [&] {
LOAD_QWEN3_5_TYPE_AND_DTYPE("qwen3_5_moe_text");
LOAD_QWEN3_5_NEXT_COMPAT_ARGS(/*moe_intermediate_size=*/512,
/*num_experts=*/512,
/*num_experts_per_tok=*/10,
/*shared_expert_intermediate_size=*/512);
});
#undef LOAD_QWEN3_5_TYPE_AND_DTYPE
#undef LOAD_QWEN3_5_NEXT_COMPAT_ARGS
#undef LOAD_QWEN3_5_ROPE_ARG
#undef LOAD_ARG_TEXT_OR_ROOT_CHAIN
#undef LOAD_ARG_TEXT_OR_ROOT
} // namespace xllm

View File

@@ -0,0 +1,280 @@
/* Copyright 2026 The xLLM Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <glog/logging.h>
#include <algorithm>
#include <memory>
#include <string>
#include <vector>
#include "core/framework/model/model_input_params.h"
#include "core/layers/common/linear.h"
#include "models/model_registry.h"
#include "qwen3_5.h"
namespace xllm {
namespace {
StateDict find_lm_head_state_dict(const StateDict& state_dict) {
static const std::vector<std::string> kLmHeadPrefixes = {
"lm_head.",
"model.lm_head.",
"language_model.lm_head.",
"model.language_model.lm_head."};
for (const auto& prefix : kLmHeadPrefixes) {
auto sub_dict = state_dict.get_dict_with_prefix(prefix);
if (sub_dict.get_tensor("weight").defined() ||
sub_dict.get_tensor("qweight").defined()) {
return sub_dict;
}
}
return StateDict({}, "");
}
bool load_qwen3_5_mtp_model_args(const JsonReader& json,
ModelArgs* args,
const std::string& base_model_type,
const std::string& mtp_model_type) {
auto base_loader = ModelRegistry::get_model_args_loader(base_model_type);
if (base_loader == nullptr || base_loader(json, args) == false) {
return false;
}
int32_t mtp_num_layers = args->num_nextn_predict_layers();
if (mtp_num_layers <= 0) {
mtp_num_layers = 1;
}
args->model_type(mtp_model_type);
args->num_nextn_predict_layers(mtp_num_layers);
args->n_layers(mtp_num_layers);
args->layer_types(std::vector<std::string>(
static_cast<size_t>(mtp_num_layers), "full_attention"));
return true;
}
} // namespace
class Qwen3_5MtpModelImpl : public Qwen3HybridModelImplBase {
public:
explicit Qwen3_5MtpModelImpl(const ModelContext& context)
: Qwen3HybridModelImplBase(context) {
const auto& options = context.get_tensor_options();
const int32_t n_layers =
std::max<int32_t>(static_cast<int32_t>(model_args_.n_layers()), 1);
pre_fc_norm_embedding_ = register_module(
"pre_fc_norm_embedding",
layer::Qwen3NextRMSNorm(
model_args_.hidden_size(), model_args_.rms_norm_eps(), options));
pre_fc_norm_hidden_ = register_module(
"pre_fc_norm_hidden",
layer::Qwen3NextRMSNorm(
model_args_.hidden_size(), model_args_.rms_norm_eps(), options));
fc_ = register_module("fc",
layer::ReplicatedLinear(model_args_.hidden_size() * 2,
model_args_.hidden_size(),
/*bias=*/false,
QuantArgs(),
options));
layers_.reserve(n_layers);
for (int32_t layer_id = 0; layer_id < n_layers; ++layer_id) {
add_decoder_layer(
std::make_shared<layer::Qwen3_5DecoderLayerImpl>(context, layer_id));
}
}
ModelOutput forward(torch::Tensor tokens,
torch::Tensor positions,
std::vector<KVCache>& kv_caches,
const ModelInputParams& input_params) override {
torch::NoGradGuard no_grad;
if (dp_size_ > 1 && tokens.sizes() == 0) {
tokens = torch::tensor({1}).to(torch::kInt32).to(device_);
positions = torch::tensor({0}).to(torch::kInt32).to(device_);
}
auto attn_metadata = layer::AttentionMetadataBuilder::build(
input_params,
model_args_.enable_mla(),
build_attention_mask(input_params));
torch::Tensor embedding = embed_tokens_(tokens);
torch::Tensor hidden = input_params.input_embedding;
if (hidden.defined() == false) {
hidden = embedding;
}
embedding = std::get<0>(pre_fc_norm_embedding_->forward(embedding));
hidden = std::get<0>(pre_fc_norm_hidden_->forward(hidden));
torch::Tensor mtp_hidden = fc_(torch::cat({embedding, hidden}, -1));
CHECK_EQ(kv_caches.size(), layers_.size());
std::optional<torch::Tensor> residual = std::nullopt;
for (size_t i = 0; i < layers_.size(); ++i) {
mtp_hidden = layers_[i]->forward(mtp_hidden,
residual,
positions,
attn_metadata,
kv_caches[i],
input_params);
}
auto [new_mtp_hidden, new_res] = norm_->forward(mtp_hidden, residual);
mtp_hidden = new_mtp_hidden;
return ModelOutput(mtp_hidden);
}
void load_state_dict(const StateDict& state_dict) override {
load_shared_embeddings(state_dict);
load_mtp_state_dict(state_dict);
}
void load_shared_embeddings(const StateDict& state_dict) {
auto embedding_state_dict =
state_dict.get_dict_with_prefix("embed_tokens.");
if (embedding_state_dict.get_tensor("weight").defined()) {
shared_embedding_loaded_ = true;
}
embed_tokens_->load_state_dict(embedding_state_dict);
}
void load_mtp_state_dict(const StateDict& state_dict) {
if (state_dict.get_tensor("pre_fc_norm_embedding.weight").defined()) {
pre_fc_norm_embedding_loaded_ = true;
}
if (state_dict.get_tensor("pre_fc_norm_hidden.weight").defined()) {
pre_fc_norm_hidden_loaded_ = true;
}
if (state_dict.get_tensor("fc.weight").defined() ||
state_dict.get_tensor("fc.qweight").defined()) {
fc_loaded_ = true;
}
if (state_dict.get_tensor("norm.weight").defined()) {
norm_loaded_ = true;
}
pre_fc_norm_embedding_->load_state_dict(
state_dict.get_dict_with_prefix("pre_fc_norm_embedding."));
pre_fc_norm_hidden_->load_state_dict(
state_dict.get_dict_with_prefix("pre_fc_norm_hidden."));
fc_->load_state_dict(state_dict.get_dict_with_prefix("fc."));
for (size_t i = 0; i < layers_.size(); ++i) {
layers_[i]->load_state_dict(
state_dict.get_dict_with_prefix("layers." + std::to_string(i) + "."));
}
norm_->load_state_dict(state_dict.get_dict_with_prefix("norm."));
}
void verify_loaded_weights(const std::string& prefix) const override {
CHECK(shared_embedding_loaded_)
<< "Failed to find shared embedding weights for qwen3.5 mtp draft "
"model";
CHECK(pre_fc_norm_embedding_loaded_)
<< "Failed to find mtp pre_fc_norm_embedding weights for qwen3.5 mtp "
"draft model";
CHECK(pre_fc_norm_hidden_loaded_)
<< "Failed to find mtp pre_fc_norm_hidden weights for qwen3.5 mtp "
"draft model";
CHECK(fc_loaded_) << "Failed to find mtp fc weights for qwen3.5 mtp draft "
"model";
CHECK(norm_loaded_)
<< "Failed to find mtp norm weights for qwen3.5 mtp draft model";
for (size_t i = 0; i < layers_.size(); ++i) {
layers_[i]->verify_loaded_weights(prefix + "layers." + std::to_string(i) +
".");
}
}
private:
layer::Qwen3NextRMSNorm pre_fc_norm_embedding_{nullptr};
layer::Qwen3NextRMSNorm pre_fc_norm_hidden_{nullptr};
layer::ReplicatedLinear fc_{nullptr};
bool shared_embedding_loaded_ = false;
bool pre_fc_norm_embedding_loaded_ = false;
bool pre_fc_norm_hidden_loaded_ = false;
bool fc_loaded_ = false;
bool norm_loaded_ = false;
};
class Qwen3_5MtpForCausalLMImpl : public Qwen3HybridForCausalLMImplBase {
public:
explicit Qwen3_5MtpForCausalLMImpl(const ModelContext& context)
: Qwen3HybridForCausalLMImplBase(context) {
mtp_model_ = std::make_shared<Qwen3_5MtpModelImpl>(context);
set_model_module(mtp_model_);
}
void load_model(std::unique_ptr<ModelLoader> loader) {
static const std::vector<std::string> kEmbeddingPrefixes = {
"model.language_model.", "language_model.model.", "model.", ""};
static const std::vector<std::string> kMtpPrefixes = {"mtp.", "model.mtp."};
bool lm_head_loaded = false;
for (const auto& state_dict : loader->get_state_dicts()) {
auto shared_embedding_state_dict =
state_dict->get_dict_with_prefix(kEmbeddingPrefixes);
auto mtp_state_dict = state_dict->get_dict_with_prefix(kMtpPrefixes);
mtp_model_->load_shared_embeddings(shared_embedding_state_dict);
mtp_model_->load_mtp_state_dict(mtp_state_dict);
if (tie_word_embeddings_) {
lm_head_->load_state_dict(
shared_embedding_state_dict.get_dict_with_prefix("embed_tokens."));
if (shared_embedding_state_dict.get_tensor("embed_tokens.weight")
.defined()) {
lm_head_loaded = true;
}
} else {
auto lm_head_state_dict = find_lm_head_state_dict(*state_dict);
lm_head_->load_state_dict(lm_head_state_dict);
if (lm_head_state_dict.get_tensor("weight").defined() ||
lm_head_state_dict.get_tensor("qweight").defined()) {
lm_head_loaded = true;
}
}
}
CHECK(lm_head_loaded)
<< "Failed to find lm_head weights for qwen3.5 mtp draft model";
mtp_model_->verify_loaded_weights("mtp.");
}
private:
std::shared_ptr<Qwen3_5MtpModelImpl> mtp_model_;
};
TORCH_MODULE(Qwen3_5MtpForCausalLM);
REGISTER_CAUSAL_MODEL(qwen3_5_mtp, Qwen3_5MtpForCausalLM);
REGISTER_CAUSAL_MODEL(qwen3_5_moe_mtp, Qwen3_5MtpForCausalLM);
REGISTER_MODEL_ARGS_LOADER(qwen3_5_mtp,
[](const JsonReader& json, ModelArgs* args) {
return load_qwen3_5_mtp_model_args(
json, args, "qwen3_5", "qwen3_5_mtp");
});
REGISTER_MODEL_ARGS_LOADER(qwen3_5_moe_mtp,
[](const JsonReader& json, ModelArgs* args) {
return load_qwen3_5_mtp_model_args(
json, args, "qwen3_5_moe", "qwen3_5_moe_mtp");
});
} // namespace xllm

View File

@@ -0,0 +1,126 @@
/* Copyright 2026 The xLLM Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <string>
#include <unordered_set>
#include <vector>
#include "core/layers/npu_torch/qwen3_next_decoder_layer_impl.h"
#include "models/model_registry.h"
#include "qwen3_next_hybrid_base.h"
namespace xllm {
class Qwen3NextModelImpl : public Qwen3HybridModelImplBase {
public:
explicit Qwen3NextModelImpl(const ModelContext& context)
: Qwen3NextModelImpl(context, /*init_decoder_layers=*/true) {}
protected:
explicit Qwen3NextModelImpl(const ModelContext& context,
bool init_decoder_layers)
: Qwen3HybridModelImplBase(context) {
if (init_decoder_layers) {
const int32_t n_layers = context.get_model_args().n_layers();
for (int32_t layer_id = 0; layer_id < n_layers; ++layer_id) {
add_decoder_layer(std::make_shared<layer::Qwen3NextDecoderLayerImpl>(
context, layer_id));
}
}
}
};
TORCH_MODULE(Qwen3NextModel);
class Qwen3NextForCausalLMImpl : public Qwen3HybridForCausalLMImplBase {
public:
explicit Qwen3NextForCausalLMImpl(const ModelContext& context)
: Qwen3NextForCausalLMImpl(context, /*init_model=*/true) {}
protected:
explicit Qwen3NextForCausalLMImpl(const ModelContext& context,
bool init_model)
: Qwen3HybridForCausalLMImplBase(context) {
if (init_model) {
set_model_module(std::make_shared<Qwen3NextModelImpl>(context));
}
}
};
TORCH_MODULE(Qwen3NextForCausalLM);
// register the causal model
REGISTER_CAUSAL_MODEL(qwen3_next, Qwen3NextForCausalLM);
// register the model args
REGISTER_MODEL_ARGS(qwen3_next, [&] {
LOAD_ARG_OR(model_type, "model_type", "qwen3_next");
LOAD_ARG_OR(dtype, "torch_dtype", "");
LOAD_ARG_OR(attention_bias, "attention_bias", false);
LOAD_ARG_OR(attention_dropout, "attention_dropout", 0.0f);
LOAD_ARG_OR(bos_token_id, "bos_token_id", 151643);
LOAD_ARG_OR(decoder_sparse_step, "decoder_sparse_step", 1);
LOAD_ARG_OR(eos_token_id, "eos_token_id", 151645);
LOAD_ARG_OR(head_dim, "head_dim", 256);
LOAD_ARG_OR(hidden_act, "hidden_act", "silu");
LOAD_ARG_OR(hidden_size, "hidden_size", 2048);
LOAD_ARG_OR(initializer_range, "initializer_range", 0.02f);
LOAD_ARG_OR(intermediate_size, "intermediate_size", 5120);
LOAD_ARG_OR(max_position_embeddings, "max_position_embeddings", 262144);
LOAD_ARG_OR(max_window_layers, "max_window_layers", 28);
LOAD_ARG_OR(moe_intermediate_size, "moe_intermediate_size", 512);
LOAD_ARG_OR(norm_topk_prob, "norm_topk_prob", true);
LOAD_ARG_OR(n_heads, "num_attention_heads", 16);
LOAD_ARG_OR(num_experts, "num_experts", 512);
LOAD_ARG_OR(num_experts_per_tok, "num_experts_per_tok", 10);
LOAD_ARG_OR(n_layers, "num_hidden_layers", 48);
LOAD_ARG_OR(n_kv_heads, "num_key_value_heads", 2);
LOAD_ARG_OR(output_router_logits, "output_router_logits", false);
LOAD_ARG_OR(rms_norm_eps, "rms_norm_eps", 1e-6);
LOAD_ARG_OR(rope_theta, "rope_theta", 10000000.0f);
LOAD_ARG_OR(router_aux_loss_coef, "router_aux_loss_coef", 0.001f);
LOAD_ARG_OR(use_sliding_window, "use_sliding_window", false);
LOAD_ARG_OR(sliding_window, "sliding_window", 4096);
LOAD_ARG_OR(tie_word_embeddings, "tie_word_embeddings", false);
LOAD_ARG_OR(vocab_size, "vocab_size", 151936);
LOAD_ARG_OR(mlp_only_layers, "mlp_only_layers", std::vector<int>());
// Additional parameters for Qwen3-Next architecture
LOAD_ARG_OR(attn_output_gate, "attn_output_gate", true);
LOAD_ARG_OR(full_attention_interval, "full_attention_interval", 4);
LOAD_ARG_OR(linear_conv_kernel_dim, "linear_conv_kernel_dim", 4);
LOAD_ARG_OR(linear_key_head_dim, "linear_key_head_dim", 128);
LOAD_ARG_OR(linear_num_key_heads, "linear_num_key_heads", 16);
LOAD_ARG_OR(linear_num_value_heads, "linear_num_value_heads", 32);
LOAD_ARG_OR(linear_value_head_dim, "linear_value_head_dim", 128);
LOAD_ARG_OR(partial_rotary_factor, "partial_rotary_factor", 0.25f);
LOAD_ARG_OR(
shared_expert_intermediate_size, "shared_expert_intermediate_size", 512);
LOAD_ARG_OR(layer_types, "layer_types", std::vector<std::string>());
// MoE compatibility with fused_moe implementation.
LOAD_ARG_OR(n_routed_experts, "n_routed_experts", args->num_experts());
SET_ARG(n_shared_experts,
args->shared_expert_intermediate_size() > 0 ? 1 : 0);
SET_ARG(scoring_func, "softmax");
SET_ARG(topk_method, "");
SET_ARG(n_group, -1);
SET_ARG(topk_group, 0);
SET_ARG(routed_scaling_factor, 1.0);
SET_ARG(stop_token_ids, std::unordered_set<int32_t>({args->eos_token_id()}));
});
} // namespace xllm

View File

@@ -0,0 +1,323 @@
/* Copyright 2026 The xLLM Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <torch/torch.h>
#include <algorithm>
#include <memory>
#include <string>
#include <vector>
#include "core/common/global_flags.h"
#include "core/framework/kv_cache/kv_cache.h"
#include "core/framework/model/model_input_params.h"
#include "core/framework/model/model_output.h"
#include "core/framework/model_context.h"
#include "core/framework/model_loader.h"
#include "core/layers/common/attention_mask.h"
#include "core/layers/common/attention_metadata_builder.h"
#include "core/layers/common/lm_head.h"
#include "core/layers/common/qwen3_next_rms_norm.h"
#include "core/layers/common/word_embedding.h"
#include "core/layers/npu_torch/qwen3_next_hybrid_decoder_layer_base.h"
namespace xllm {
class Qwen3HybridModelModule : public torch::nn::Module {
public:
virtual ModelOutput forward(torch::Tensor tokens,
torch::Tensor positions,
std::vector<KVCache>& kv_caches,
const ModelInputParams& input_params) = 0;
virtual void load_state_dict(const StateDict& state_dict) = 0;
virtual void verify_loaded_weights(const std::string& prefix) const = 0;
virtual layer::WordEmbedding get_word_embedding() = 0;
virtual void set_word_embedding(layer::WordEmbedding& word_embedding) = 0;
};
using Qwen3HybridModelModulePtr = std::shared_ptr<Qwen3HybridModelModule>;
class Qwen3HybridModelImplBase : public Qwen3HybridModelModule {
public:
explicit Qwen3HybridModelImplBase(const ModelContext& context)
: device_(context.get_tensor_options().device()),
model_args_(context.get_model_args()) {
auto options = context.get_tensor_options();
auto parallel_args = context.get_parallel_args();
blocks_ = register_module("layers", torch::nn::ModuleList());
layers_.reserve(model_args_.n_layers());
device_ = options.device();
dtype_ = options.dtype().toScalarType();
norm_ = register_module(
"norm",
xllm::layer::Qwen3NextRMSNorm(
model_args_.hidden_size(), model_args_.rms_norm_eps(), options));
embed_tokens_ =
register_module("embed_tokens", layer::WordEmbedding(context));
int32_t mask_value = FLAGS_enable_chunked_prefill ? -9984 : 1;
attn_mask_ = layer::AttentionMask(options.device(),
options.dtype().toScalarType(),
/*mask_value=*/mask_value);
dp_size_ = parallel_args.dp_size();
}
// tokens: [num_tokens]
// positions: [num_tokens] token pos in the sequence
ModelOutput forward(torch::Tensor tokens,
torch::Tensor positions,
std::vector<KVCache>& kv_caches,
const ModelInputParams& input_params) override {
// Disable gradient computation to reduce memory usage during inference
torch::NoGradGuard no_grad;
if (dp_size_ > 1) {
if (tokens.sizes() == 0) {
tokens = torch::tensor({1}).to(torch::kInt32).to(device_);
positions = torch::tensor({0}).to(torch::kInt32).to(device_);
}
}
layer::AttentionMetadata attn_metadata =
layer::AttentionMetadataBuilder::build(
input_params,
model_args_.enable_mla(),
build_attention_mask(input_params));
torch::Tensor h = embed_tokens_(tokens);
torch::Tensor mrope_cos_sin;
for (const auto& layer : layers_) {
mrope_cos_sin = layer->build_mrope_cos_sin(positions);
if (mrope_cos_sin.defined()) break;
}
std::optional<torch::Tensor> residual = std::nullopt;
for (size_t i = 0; i < layers_.size(); i++) {
auto& layer = layers_[i];
h = layer->forward(h,
residual,
positions,
attn_metadata,
kv_caches[i],
input_params,
mrope_cos_sin);
}
auto [hidden_states, residual_out] = norm_->forward(h, residual);
h = hidden_states;
return ModelOutput(h);
}
// load the weight from the checkpoint
void load_state_dict(const StateDict& state_dict) override {
embed_tokens_->load_state_dict(
state_dict.get_dict_with_prefix("embed_tokens."));
for (int i = 0; i < static_cast<int>(layers_.size()); i++) {
layers_[i]->load_state_dict(
state_dict.get_dict_with_prefix("layers." + std::to_string(i) + "."));
}
norm_->load_state_dict(state_dict.get_dict_with_prefix("norm."));
}
void verify_loaded_weights(const std::string& prefix) const override {
for (size_t i = 0; i < layers_.size(); ++i) {
layers_[i]->verify_loaded_weights(prefix + "layers." + std::to_string(i) +
".");
}
}
layer::WordEmbedding get_word_embedding() override { return embed_tokens_; }
void set_word_embedding(layer::WordEmbedding& word_embedding) override {
embed_tokens_ = word_embedding;
}
void add_decoder_layer(layer::Qwen3HybridDecoderLayerModulePtr layer) {
layers_.push_back(layer);
blocks_->push_back(layer);
}
int32_t num_hidden_layers() const {
return static_cast<int32_t>(layers_.size());
}
protected:
torch::Tensor build_attention_mask(const ModelInputParams& input_params) {
max_seq_len_ = std::max(input_params.kv_max_seq_len, max_seq_len_);
if (!FLAGS_enable_chunked_prefill) {
return attn_mask_.get_attn_mask(max_seq_len_, dtype_, device_);
}
const int32_t num_sequences = input_params.num_sequences;
if (num_sequences <= 0) {
return attn_mask_.get_attn_mask(max_seq_len_, dtype_, device_);
}
std::vector<torch::Tensor> req_mask_vec;
req_mask_vec.reserve(num_sequences);
for (int32_t j = 0; j < num_sequences; ++j) {
req_mask_vec.emplace_back(
attn_mask_.gen_append_mask(input_params.q_seq_lens_vec[j],
input_params.kv_seq_lens_vec[j],
max_seq_len_,
dtype_,
device_));
}
return torch::cat(req_mask_vec, 0);
}
ModelArgs model_args_;
torch::nn::ModuleList blocks_{nullptr};
std::vector<layer::Qwen3HybridDecoderLayerModulePtr> layers_;
int32_t max_seq_len_ = 0;
int32_t dp_size_ = 1;
torch::Device device_;
torch::ScalarType dtype_ = torch::kFloat;
layer::Qwen3NextRMSNorm norm_{nullptr};
layer::AttentionMask attn_mask_;
layer::WordEmbedding embed_tokens_{nullptr};
};
class Qwen3HybridForCausalLMImplBase : public torch::nn::Module {
public:
explicit Qwen3HybridForCausalLMImplBase(const ModelContext& context) {
tie_word_embeddings_ = context.get_model_args().tie_word_embeddings();
lm_head_ = register_module("lm_head", layer::LmHead(context));
}
// tokens: [num_tokens]
// positions: [num_tokens] token pos in the sequence
// returns: [num_tokens, hidden_size]
ModelOutput forward(const torch::Tensor& tokens,
const torch::Tensor& positions,
std::vector<KVCache>& kv_caches,
const ModelInputParams& input_params) {
return model_->forward(tokens, positions, kv_caches, input_params);
}
// hidden_states: [num_tokens, hidden_size]
// seleted_idxes: [num_tokens]
// returns: [num_tokens, vocab_size]
torch::Tensor logits(const torch::Tensor& hidden_states,
const torch::Tensor& seleted_idxes) {
auto h = hidden_states;
if (seleted_idxes.defined()) {
h = h.index_select(/*dim=*/0, seleted_idxes);
}
return lm_head_(h);
}
// hidden_states: [num_tokens, hidden_size]
// seleted_idxes: [num_tokens]
torch::Tensor pooler(const torch::Tensor& hidden_states,
const torch::Tensor& seleted_idxes) {
auto h = hidden_states;
if (seleted_idxes.defined()) {
h = h.index_select(/*dim=*/0, seleted_idxes);
}
namespace F = torch::nn::functional;
return F::normalize(h, F::NormalizeFuncOptions().p(2).dim(1));
}
void load_model(std::unique_ptr<ModelLoader> loader) {
auto has_model_weights = [](const StateDict& dict) {
return dict.get_tensor("embed_tokens.weight").defined() ||
dict.get_dict_with_prefix("layers.").size() > 0 ||
dict.get_tensor("norm.weight").defined();
};
auto has_lm_head_weights = [](const StateDict& dict) {
return dict.get_tensor("weight").defined() ||
dict.get_tensor("qweight").defined();
};
for (const auto& state_dict : loader->get_state_dicts()) {
auto model_state_dict = state_dict->get_dict_with_prefix("model.");
if (!has_model_weights(model_state_dict)) {
auto language_model_state_dict =
state_dict->get_dict_with_prefix("language_model.model.");
if (has_model_weights(language_model_state_dict)) {
model_state_dict = language_model_state_dict;
} else {
auto wrapped_language_model_state_dict =
state_dict->get_dict_with_prefix("model.language_model.");
if (has_model_weights(wrapped_language_model_state_dict)) {
model_state_dict = wrapped_language_model_state_dict;
}
}
}
model_->load_state_dict(model_state_dict);
auto lm_head_state_dict = state_dict->get_dict_with_prefix("lm_head.");
if (!has_lm_head_weights(lm_head_state_dict)) {
auto language_model_lm_head_state_dict =
state_dict->get_dict_with_prefix("language_model.lm_head.");
if (has_lm_head_weights(language_model_lm_head_state_dict)) {
lm_head_state_dict = language_model_lm_head_state_dict;
} else {
auto wrapped_language_model_lm_head_state_dict =
state_dict->get_dict_with_prefix("model.language_model.lm_head.");
if (has_lm_head_weights(wrapped_language_model_lm_head_state_dict)) {
lm_head_state_dict = wrapped_language_model_lm_head_state_dict;
} else {
auto wrapped_lm_head_state_dict =
state_dict->get_dict_with_prefix("model.lm_head.");
if (has_lm_head_weights(wrapped_lm_head_state_dict)) {
lm_head_state_dict = wrapped_lm_head_state_dict;
}
}
}
}
if (!has_lm_head_weights(lm_head_state_dict) && tie_word_embeddings_) {
auto tied_lm_head_state_dict =
model_state_dict.get_dict_with_prefix("embed_tokens.");
if (has_lm_head_weights(tied_lm_head_state_dict)) {
lm_head_state_dict = tied_lm_head_state_dict;
}
}
lm_head_->load_state_dict(lm_head_state_dict);
}
model_->verify_loaded_weights("model.");
}
virtual void prepare_expert_weight(int32_t layer_id,
const std::vector<int32_t>& expert_ids) {
return;
}
virtual void update_expert_weight(int32_t layer_id) { return; }
layer::LmHead get_lm_head() { return lm_head_; }
void set_lm_head(layer::LmHead& head) { lm_head_ = head; }
layer::WordEmbedding get_word_embedding() {
return model_->get_word_embedding();
}
void set_word_embedding(layer::WordEmbedding& word_embedding) {
model_->set_word_embedding(word_embedding);
}
void set_model_module(Qwen3HybridModelModulePtr model) {
model_ = register_module("model", std::move(model));
}
protected:
bool tie_word_embeddings_{false};
layer::LmHead lm_head_{nullptr};
Qwen3HybridModelModulePtr model_;
};
} // namespace xllm

View File

@@ -0,0 +1,312 @@
/* Copyright 2026 The xLLM Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include "core/framework/model/model_output.h"
#include "core/layers/common/lm_head.h"
#include "core/layers/common/qwen3_next_rms_norm.h"
#include "core/layers/common/rms_norm.h"
#include "core/layers/mlu/qwen3_5_decoder_layer.h"
#include "core/layers/qwen3_vision_layer.h"
#include "models/llm/llm_model_base.h"
#include "models/model_registry.h"
#include "models/vlm/qwen3_vl_base.h"
#include "processors/input_processor.h"
#include "processors/qwen2_vl_image_processor.h"
#include "qwen3_vl.h"
namespace xllm {
class Qwen3_5ModelImpl final
: public LlmModelImplBase<layer::Qwen3_5DecoderLayer> {
public:
Qwen3_5ModelImpl(const ModelContext& context)
: LlmModelImplBase<layer::Qwen3_5DecoderLayer>("qwen3_5",
context.get_model_args()) {
auto model_args = context.get_model_args();
auto options = context.get_tensor_options();
auto parallel_args = context.get_parallel_args();
dp_size_ = parallel_args.dp_size();
if (!mrope_section_.empty()) {
int64_t rotary_dim = static_cast<int64_t>(
model_args.head_dim() * model_args.partial_rotary_factor());
cos_sin_ = layer::rotary::get_concat_rotary_embedding(
rotary_dim,
model_args.max_position_embeddings(),
model_args.rope_theta(),
options);
}
layers_.reserve(model_args.n_layers());
rms_norm_ = register_module(
"norm",
layer::Qwen3NextRMSNorm(
model_args.hidden_size(), model_args.rms_norm_eps(), options));
embed_tokens_ =
register_module("embed_tokens", layer::WordEmbedding(context));
for (int32_t i = 0; i < model_args.n_layers(); i++) {
auto layer = layer::Qwen3_5DecoderLayer(context, i);
layers_.push_back(layer);
}
}
void load_state_dict(const StateDict& state_dict) override {
embed_tokens_->load_state_dict(
state_dict.get_dict_with_prefix("embed_tokens."));
// call each layer's load_state_dict function
for (size_t i = 0; i < layers_.size(); i++) {
layers_[i]->load_state_dict(
state_dict.get_dict_with_prefix("layers." + std::to_string(i) + "."));
}
rms_norm_->load_state_dict(state_dict.get_dict_with_prefix("norm."));
}
std::pair<torch::Tensor, torch::Tensor> apply_mrope(
const torch::Tensor positions) override {
auto target_cos_sin = cos_sin_.index({positions});
auto target_cos_sin_chunks = target_cos_sin.chunk(/*chunks=*/2, /*dim=*/-1);
auto cos_pos = target_cos_sin_chunks[0].contiguous();
auto sin_pos = target_cos_sin_chunks[1].contiguous();
auto apply = [this](torch::Tensor x) {
auto freqs_t = x[0].clone();
int64_t mrop_length = static_cast<int64_t>(freqs_t.size(-1) / 2);
for (int32_t dim_idx = 1; dim_idx <= 2; ++dim_idx) {
int64_t offset = dim_idx;
int64_t section_len = mrope_section_[dim_idx];
int64_t length = section_len * 3;
auto idx_first_half = torch::arange(offset, length, 3, torch::kLong);
auto idx_second_half = torch::arange(
offset + mrop_length, length + mrop_length, 3, torch::kLong);
auto idx_tensor =
torch::cat({idx_first_half, idx_second_half}, 0).to(x.device());
auto src = x[dim_idx].index_select(-1, idx_tensor);
freqs_t.index_copy_(-1, idx_tensor, src);
}
return freqs_t;
};
cos_pos = apply(cos_pos.reshape({positions.size(0), -1, cos_pos.size(-1)}));
sin_pos = apply(sin_pos.reshape({positions.size(0), -1, sin_pos.size(-1)}));
return std::make_pair(cos_pos, sin_pos);
}
virtual ModelOutput forward(torch::Tensor tokens,
torch::Tensor positions,
std::vector<KVCache>& kv_caches,
const ModelInputParams& input_params) {
ModelInputParams& input_params_new =
const_cast<ModelInputParams&>(input_params);
std::vector<torch::Tensor> deep_stacks;
if (dp_size_ > 1) {
if (tokens.numel() == 0) {
tokens = torch::tensor({1}).to(torch::kInt32).to(tokens.device());
positions = torch::tensor({1}).to(torch::kInt32).to(positions.device());
}
auto& dp_token_nums = input_params_new.dp_global_token_nums;
std::replace(dp_token_nums.begin(), dp_token_nums.end(), 0, 1);
}
auto inputs_embeds = input_params.input_embedding;
torch::Tensor h;
if (inputs_embeds.defined()) {
h = inputs_embeds;
} else {
h = embed_tokens_(tokens);
}
if (!input_params_new.attn_metadata) {
input_params_new.attn_metadata =
std::make_shared<layer::AttentionMetadata>(
get_attention_metadata(input_params_new, h));
}
auto& attn_metadata = *(input_params_new.attn_metadata);
bool only_prefill =
(attn_metadata.is_prefill || attn_metadata.is_chunked_prefill);
if (positions.dim() == 2 && only_prefill && !mrope_section_.empty()) {
std::tie(attn_metadata.mrope_cos, attn_metadata.mrope_sin) =
apply_mrope(positions);
}
std::optional<torch::Tensor> residual;
for (size_t i = 0; i < layers_.size(); i++) {
auto& layer = layers_[i];
h = layer(h,
residual,
positions,
attn_metadata,
kv_caches[i],
input_params_new);
}
if (residual.has_value()) {
h = h + residual.value();
}
auto hidden_states = std::get<0>(rms_norm_(h));
return ModelOutput(hidden_states);
}
private:
int32_t dp_size_ = 1;
layer::Qwen3NextRMSNorm rms_norm_{nullptr};
layer::AttentionMetadata get_attention_metadata(
const ModelInputParams& params,
const torch::Tensor& h) {
auto attn_metadata = layer::AttentionMetadataBuilder::build(params, false);
// TODO: support linear attention
return attn_metadata;
}
};
TORCH_MODULE(Qwen3_5Model);
class Qwen3_5ForCausalLMImpl : public LlmForCausalLMImplBase<Qwen3_5Model> {
public:
Qwen3_5ForCausalLMImpl(const ModelContext& context)
: LlmForCausalLMImplBase<Qwen3_5Model>(context) {}
torch::Tensor pooler(const torch::Tensor& hidden_states,
const torch::Tensor& seleted_idxes) {
auto h = hidden_states;
if (seleted_idxes.defined()) {
h = h.index_select(/*dim=*/0, seleted_idxes);
}
namespace F = torch::nn::functional;
return F::normalize(h, F::NormalizeFuncOptions().p(2).dim(1));
}
};
TORCH_MODULE(Qwen3_5ForCausalLM);
using Qwen3_5ForConditionalGenerationImpl =
Qwen3VLForConditionalGenerationBase<Qwen3_VisionTransformer,
Qwen3_5ForCausalLM>;
TORCH_MODULE(Qwen3_5ForConditionalGeneration);
#define LOAD_QWEN3_5_COMMON_ARGS() \
LOAD_ARG_OR(model_type, "model_type", "qwen3_5"); \
LOAD_ARG_OR(dtype, "text_config.dtype", "bfloat16"); \
LOAD_ARG_OR(vocab_size, "text_config.vocab_size", 248320); \
LOAD_ARG_OR(hidden_size, "text_config.hidden_size", 5120); \
LOAD_ARG_OR(hidden_act, "text_config.hidden_act", "silu"); \
LOAD_ARG_OR(intermediate_size, "text_config.intermediate_size", 17408); \
LOAD_ARG_OR(n_layers, "text_config.num_hidden_layers", 64); \
LOAD_ARG_OR(n_heads, "text_config.num_attention_heads", 24); \
LOAD_ARG(n_kv_heads, "text_config.num_key_value_heads"); \
LOAD_ARG_OR( \
max_position_embeddings, "text_config.max_position_embeddings", 262144); \
LOAD_ARG_OR(rms_norm_eps, "text_config.rms_norm_eps", 1e-6); \
LOAD_ARG_OR(eos_token_id, "text_config.eos_token_id", 248044); \
LOAD_ARG_OR( \
rope_theta, "text_config.rope_parameters.rope_theta", 10000000.0f); \
LOAD_ARG_OR(head_dim, "text_config.head_dim", 256); \
LOAD_ARG_OR(tie_word_embeddings, "tie_word_embeddings", false); \
LOAD_ARG(layer_types, "text_config.layer_types"); \
LOAD_ARG_OR( \
linear_conv_kernel_dim, "text_config.linear_conv_kernel_dim", 4); \
LOAD_ARG_OR(linear_key_head_dim, "text_config.linear_key_head_dim", 128); \
LOAD_ARG_OR( \
linear_value_head_dim, "text_config.linear_value_head_dim", 128); \
LOAD_ARG_OR(linear_num_key_heads, "text_config.linear_num_key_heads", 16); \
LOAD_ARG_OR( \
linear_num_value_heads, "text_config.linear_num_value_heads", 48); \
LOAD_ARG_OR( \
full_attention_interval, "text_config.full_attention_interval", 4); \
LOAD_ARG_OR(attn_output_gate, "text_config.attn_output_gate", false); \
LOAD_ARG_OR( \
num_nextn_predict_layers, "text_config.mtp_num_hidden_layers", 0); \
LOAD_ARG_OR(attention_bias, "text_config.attention_bias", false); \
LOAD_ARG_OR(attention_dropout, "text_config.attention_dropout", 0.0f); \
LOAD_ARG_OR(initializer_range, "text_config.initializer_range", 0.02f); \
LOAD_ARG_OR( \
mlp_only_layers, "text_config.mlp_only_layers", std::vector<int32_t>()); \
LOAD_ARG(rope_scaling_mrope_section, \
"text_config.rope_parameters.mrope_section"); \
LOAD_ARG_OR(rope_scaling_rope_type, \
"text_config.rope_parameters.rope_type", \
"default"); \
LOAD_ARG_OR(partial_rotary_factor, \
"text_config.rope_parameters.partial_rotary_factor", \
0.25f)
#define LOAD_QWEN3_5_VISION_ARGS() \
LOAD_ARG_OR(image_token_id, "image_token_id", 248056); \
LOAD_ARG_OR(video_token_id, "video_token_id", 248057); \
LOAD_ARG_OR(vision_start_token_id, "vision_start_token_id", 248053); \
LOAD_ARG_OR(vision_end_token_id, "vision_end_token_id", 248054); \
LOAD_ARG(mm_deepstack_visual_indexes, \
"vision_config.deepstack_visual_indexes"); \
LOAD_ARG_OR(mm_num_hidden_layers, "vision_config.depth", 27); \
LOAD_ARG_OR(mm_hidden_act, "vision_config.hidden_act", "gelu_pytorch_tanh"); \
LOAD_ARG_OR(mm_hidden_size, "vision_config.hidden_size", 1152); \
LOAD_ARG_OR(mm_num_channels, "vision_config.in_channels", 3); \
LOAD_ARG_OR(mm_initializer_range, "vision_config.initializer_range", 0.02f); \
LOAD_ARG_OR(mm_intermediate_size, "vision_config.intermediate_size", 4304); \
LOAD_ARG_OR(mm_num_attention_heads, "vision_config.num_heads", 16); \
LOAD_ARG_OR(mm_num_position_embeddings, \
"vision_config.num_position_embeddings", \
2304); \
LOAD_ARG_OR(mm_projection_dim, "vision_config.out_hidden_size", 5120); \
LOAD_ARG_OR(mm_patch_size, "vision_config.patch_size", 16); \
LOAD_ARG_OR(mm_spatial_merge_size, "vision_config.spatial_merge_size", 2); \
LOAD_ARG_OR(mm_temporal_patch_size, "vision_config.temporal_patch_size", 2); \
LOAD_ARG_OR_FUNC(mm_head_dim, "head_dim", [&] { \
return args->mm_hidden_size() / args->mm_num_attention_heads(); \
}); \
LOAD_ARG_OR( \
rope_scaling_rope_type, "vision_config.rope_scaling.type", "mrope")
REGISTER_INPUT_PROCESSOR(qwen3_5, Qwen2_5_VLInputProcessor);
REGISTER_CAUSAL_VLM_MODEL(qwen3_5, Qwen3_5ForConditionalGeneration);
REGISTER_IMAGE_PROCESSOR(qwen3_5, Qwen2VLImageProcessor);
REGISTER_MODEL_ARGS(qwen3_5, [&] {
LOAD_QWEN3_5_COMMON_ARGS();
LOAD_QWEN3_5_VISION_ARGS();
SET_ARG(stop_token_ids, std::unordered_set<int32_t>({args->eos_token_id()}));
});
REGISTER_INPUT_PROCESSOR(qwen3_5_moe, Qwen2_5_VLInputProcessor);
REGISTER_CAUSAL_VLM_MODEL(qwen3_5_moe, Qwen3_5ForConditionalGeneration);
REGISTER_IMAGE_PROCESSOR(qwen3_5_moe, Qwen2VLImageProcessor);
REGISTER_MODEL_ARGS(qwen3_5_moe, [&] {
LOAD_QWEN3_5_COMMON_ARGS();
LOAD_QWEN3_5_VISION_ARGS();
LOAD_ARG_OR(decoder_sparse_step, "text_config.decoder_sparse_step", 1);
LOAD_ARG_OR(moe_intermediate_size, "text_config.moe_intermediate_size", 512);
LOAD_ARG_OR(num_experts, "text_config.num_experts", 512);
LOAD_ARG_OR(num_experts_per_tok, "text_config.num_experts_per_tok", 10);
LOAD_ARG_OR(shared_expert_intermediate_size,
"text_config.shared_expert_intermediate_size",
512);
LOAD_ARG_OR(norm_topk_prob, "text_config.norm_topk_prob", true);
LOAD_ARG_OR(
n_routed_experts, "text_config.n_routed_experts", args->num_experts());
SET_ARG(n_shared_experts,
args->shared_expert_intermediate_size() > 0 ? 1 : 0);
SET_ARG(scoring_func, "softmax");
SET_ARG(topk_method, "");
SET_ARG(n_group, -1);
SET_ARG(topk_group, 0);
SET_ARG(routed_scaling_factor, 1.0f);
SET_ARG(stop_token_ids, std::unordered_set<int32_t>({args->eos_token_id()}));
});
} // namespace xllm