fix(critical): fold max_completion_tokens + max_num_seqs=2 + max_model_len=80000 + xllm_latest layer import
Sub 655 root causes (confirmed from log analysis): 1. protocol.py: max_completion_tokens never folded into max_tokens → 162/881 replay requests rejected 400 (extra_forbidden) 2. max_num_seqs=1 → t2_n_2 test fails (needs n=2) 3. max_model_len=131072 → OOM crash at 62% replay, opencompass all 0 Fixes: - protocol.py: model_validator fold_max_completion_tokens - yaml: max_num_seqs=2, max_model_len=80000, PYTORCH_CUDA_ALLOC_CONF - topk_softmax stays =0 (corex CUB BlockReduce incompatible on BI-V100) xllm_latest import to ex_engine/: - npu_torch layers: GDN(1164L), Qwen3.5 GDN, attention, fused_moe - cuda/moe kernels: topk_softmax_kernels.cuh, moe_combine, moe_compute_index - npu kernels: causal_conv1d, recurrent_gated_delta_rule - model headers: qwen3_5.h, qwen3_next.h
This commit is contained in:
@@ -8,14 +8,14 @@ command:
|
||||
- --served-model-name
|
||||
- llm
|
||||
- --max-model-len
|
||||
- '131072'
|
||||
- '80000'
|
||||
- --gpu-memory-utilization
|
||||
- '0.90'
|
||||
- --trust-remote-code
|
||||
- -tp
|
||||
- '4'
|
||||
- --max-num-seqs
|
||||
- '1'
|
||||
- '2'
|
||||
- --disable-log-requests
|
||||
- --disable-frontend-multiprocessing
|
||||
- --max-num-batched-tokens
|
||||
@@ -47,3 +47,5 @@ env:
|
||||
value: hybrid64
|
||||
- name: BI100_MOE_COREX_TOPK_SOFTMAX
|
||||
value: '0'
|
||||
- name: PYTORCH_CUDA_ALLOC_CONF
|
||||
value: expandable_segments:True
|
||||
|
||||
28
ex_engine/csrc/ilu_CMakeLists.txt
Normal file
28
ex_engine/csrc/ilu_CMakeLists.txt
Normal file
@@ -0,0 +1,28 @@
|
||||
include(cc_library)
|
||||
set(CMAKE_CUDA_ARCHITECTURES ivcore11)
|
||||
file(GLOB_RECURSE ILU_HEADER_FILES
|
||||
"${CMAKE_CURRENT_LIST_DIR}/*.h"
|
||||
)
|
||||
|
||||
file(GLOB_RECURSE ILU_SOURCE_FILES
|
||||
"${CMAKE_CURRENT_LIST_DIR}/*.cpp"
|
||||
"${CMAKE_CURRENT_LIST_DIR}/*.cu"
|
||||
)
|
||||
|
||||
find_package(Python3 REQUIRED COMPONENTS Interpreter Development)
|
||||
|
||||
cc_library(
|
||||
NAME
|
||||
ilu_kernels
|
||||
HDRS
|
||||
${ILU_HEADER_FILES}
|
||||
SRCS
|
||||
${ILU_SOURCE_FILES}
|
||||
DEPS
|
||||
torch
|
||||
:util
|
||||
ixformer_kernels
|
||||
ixformer
|
||||
${Python3_LIBRARIES}
|
||||
cuinfer
|
||||
)
|
||||
14
ex_engine/csrc/ilu_layers_CMakeLists.txt
Executable file
14
ex_engine/csrc/ilu_layers_CMakeLists.txt
Executable file
@@ -0,0 +1,14 @@
|
||||
include(cc_library)
|
||||
|
||||
cc_library(
|
||||
NAME
|
||||
ilu_layers
|
||||
HDRS
|
||||
attention.h
|
||||
fused_moe.h
|
||||
SRCS
|
||||
attention.cpp
|
||||
fused_moe.cpp
|
||||
DEPS
|
||||
:common_layers
|
||||
)
|
||||
124
ex_engine/xllm_kernels/cuda/moe/fused_moe.cpp
Normal file
124
ex_engine/xllm_kernels/cuda/moe/fused_moe.cpp
Normal file
@@ -0,0 +1,124 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "kernels/cuda/cuda_ops_api.h"
|
||||
#include "kernels/cuda/utils.h"
|
||||
#include "platform/device.h"
|
||||
#include "platform/platform.h"
|
||||
|
||||
namespace xllm::kernel::cuda {
|
||||
|
||||
torch::Tensor cutlass_fused_moe(
|
||||
const torch::Tensor& input, // [num_tokens, hidden]
|
||||
const torch::Tensor& token_selected_experts, // [num_tokens, top_k]
|
||||
const torch::Tensor& token_final_scales, // [num_tokens, top_k]
|
||||
const torch::Tensor&
|
||||
fc1_expert_weights, // [num_experts, inter_dim, hidden]
|
||||
const torch::Tensor&
|
||||
fc2_expert_weights, // [num_experts, hidden, inter_dim]
|
||||
torch::ScalarType output_dtype,
|
||||
const std::vector<torch::Tensor>& quant_scales,
|
||||
int32_t tp_size,
|
||||
int32_t tp_rank,
|
||||
int32_t ep_size,
|
||||
int32_t ep_rank,
|
||||
int32_t cluster_size,
|
||||
int32_t cluster_rank,
|
||||
const std::optional<torch::Tensor>& fc1_expert_biases,
|
||||
const std::optional<torch::Tensor>& fc2_expert_biases,
|
||||
const std::optional<torch::Tensor>& input_sf,
|
||||
const std::optional<torch::Tensor>& swiglu_alpha,
|
||||
const std::optional<torch::Tensor>& swiglu_beta,
|
||||
const std::optional<torch::Tensor>& swiglu_limit,
|
||||
const std::optional<torch::Tensor>& output,
|
||||
bool enable_alltoall,
|
||||
bool use_deepseek_fp8_block_scale,
|
||||
bool use_w4_group_scaling,
|
||||
bool use_mxfp8_act_scaling,
|
||||
bool min_latency_mode,
|
||||
bool use_packed_weights,
|
||||
int32_t tune_max_num_tokens,
|
||||
ActivationType activation_type) {
|
||||
int64_t num_rows = input.size(0);
|
||||
int64_t hidden_size = fc2_expert_weights.size(1);
|
||||
|
||||
if (min_latency_mode) {
|
||||
num_rows *= fc2_expert_weights.size(0);
|
||||
}
|
||||
|
||||
std::vector<int64_t> output_shape = {num_rows, hidden_size};
|
||||
torch::Tensor result_output;
|
||||
if (output.has_value() && output.value().defined()) {
|
||||
result_output = output.value();
|
||||
} else {
|
||||
torch::TensorOptions options = input.options().dtype(output_dtype);
|
||||
result_output = torch::empty(output_shape, options);
|
||||
}
|
||||
|
||||
std::string fused_moe_uri = "fused_moe";
|
||||
if (Platform::is_support_sm90a()) {
|
||||
fused_moe_uri += "_90";
|
||||
} else if (Platform::is_support_sm100a() || Platform::is_support_sm100f()) {
|
||||
fused_moe_uri += "_100";
|
||||
} else if (Platform::is_support_sm120a()) {
|
||||
fused_moe_uri += "_120";
|
||||
} else {
|
||||
LOG(FATAL) << "FusedMoE is only supported on sm90, sm100, sm120.";
|
||||
}
|
||||
|
||||
bind_tvmffi_stream_to_current_torch_stream(input.device());
|
||||
|
||||
ffi::Module fused_moe_runner =
|
||||
get_function(fused_moe_uri, "init")(
|
||||
to_dl_data_type(input.scalar_type()),
|
||||
to_dl_data_type(fc1_expert_weights.scalar_type()),
|
||||
to_dl_data_type(output_dtype),
|
||||
use_deepseek_fp8_block_scale,
|
||||
use_w4_group_scaling,
|
||||
use_mxfp8_act_scaling,
|
||||
use_packed_weights)
|
||||
.cast<ffi::Module>();
|
||||
|
||||
fused_moe_runner->GetFunction("run_moe").value()(
|
||||
to_ffi_tensor(result_output),
|
||||
to_ffi_tensor(input),
|
||||
to_ffi_tensor(token_selected_experts),
|
||||
to_ffi_optional_tensor(token_final_scales),
|
||||
to_ffi_tensor(fc1_expert_weights),
|
||||
to_ffi_optional_tensor(fc1_expert_biases),
|
||||
to_ffi_tensor(fc2_expert_weights),
|
||||
to_ffi_optional_tensor(fc2_expert_biases),
|
||||
to_ffi_optional_array_tensors(quant_scales),
|
||||
to_ffi_optional_tensor(input_sf),
|
||||
to_ffi_optional_tensor(swiglu_alpha),
|
||||
to_ffi_optional_tensor(swiglu_beta),
|
||||
to_ffi_optional_tensor(swiglu_limit),
|
||||
tp_size,
|
||||
tp_rank,
|
||||
ep_size,
|
||||
ep_rank,
|
||||
cluster_size,
|
||||
cluster_rank,
|
||||
enable_alltoall,
|
||||
min_latency_mode,
|
||||
/*profile_ids=*/ffi::Optional<ffi::Array<int64_t>>(), // TODO: support
|
||||
// auto tuning
|
||||
// profile ids
|
||||
support_pdl(),
|
||||
activation_type);
|
||||
|
||||
return result_output;
|
||||
}
|
||||
} // namespace xllm::kernel::cuda
|
||||
105
ex_engine/xllm_kernels/cuda/moe/moe_combine.cu
Executable file
105
ex_engine/xllm_kernels/cuda/moe/moe_combine.cu
Executable file
@@ -0,0 +1,105 @@
|
||||
/* Copyright 2025-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.
|
||||
==============================================================================*/
|
||||
|
||||
// Fused MoE combine kernel — reorder + weighted sum in one pass.
|
||||
// Replaces: torch::zeros + index_copy_ + view + multiply + sum
|
||||
//
|
||||
// Algorithm per token (each block handles one token):
|
||||
// 1. For each of its topk experts, read gemm2 at flat_idx directly
|
||||
// (gemm2 is flat-index-ordered after scatter via index_copy_ with dst_src)
|
||||
// 2. Multiply by router weight
|
||||
// 3. Accumulate into output[token]
|
||||
//
|
||||
// Grid: num_tokens (N) blocks
|
||||
// Block: HIDDEN_DIM / HIDDEN_TILE threads
|
||||
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include "device_utils.cuh"
|
||||
#include "kernels/cuda/cuda_ops_api.h"
|
||||
|
||||
namespace xllm::kernel::cuda {
|
||||
|
||||
constexpr int32_t kCombineBlockSize = 256;
|
||||
|
||||
template <typename scalar_t>
|
||||
__global__ void XLLM_KERNEL_ATTR(kCombineBlockSize) moe_combine_kernel(
|
||||
const scalar_t* __restrict__ gemm2, // [N*topk, H] flat-index-ordered
|
||||
const float* __restrict__ reduce_weight, // [N, topk]
|
||||
scalar_t* __restrict__ output, // [N, H]
|
||||
int64_t N,
|
||||
int32_t topk,
|
||||
int64_t H) {
|
||||
int64_t token_id = blockIdx.x; // 0 .. N-1
|
||||
if (token_id >= N) return;
|
||||
|
||||
int32_t tid = threadIdx.x;
|
||||
int32_t stride = kCombineBlockSize;
|
||||
|
||||
// Accumulate over topk experts for this token
|
||||
for (int64_t h = tid; h < H; h += stride) {
|
||||
float acc = 0.0f;
|
||||
for (int32_t k = 0; k < topk; ++k) {
|
||||
int64_t flat_idx = token_id * topk + k;
|
||||
float w = reduce_weight[flat_idx];
|
||||
acc += w * static_cast<float>(gemm2[flat_idx * H + h]);
|
||||
}
|
||||
output[token_id * H + h] = static_cast<scalar_t>(acc);
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Host-side orchestrator ----
|
||||
torch::Tensor moe_combine_result(
|
||||
const torch::Tensor& gemm2, // [N*topk, H] flat-index-ordered
|
||||
const torch::Tensor& reduce_weight, // [N, topk] float or same as gemm2
|
||||
int64_t N,
|
||||
int32_t topk) {
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
int64_t H = gemm2.size(1);
|
||||
auto dtype = gemm2.scalar_type();
|
||||
|
||||
auto output = torch::empty({N, H}, gemm2.options());
|
||||
auto rw = reduce_weight.to(gemm2.device(), torch::kFloat32).contiguous();
|
||||
|
||||
if (dtype == torch::kFloat16) {
|
||||
moe_combine_kernel<c10::Half>
|
||||
<<<N, kCombineBlockSize, 0, stream>>>(gemm2.data_ptr<c10::Half>(),
|
||||
rw.data_ptr<float>(),
|
||||
output.data_ptr<c10::Half>(),
|
||||
N,
|
||||
topk,
|
||||
H);
|
||||
} else if (dtype == torch::kBFloat16) {
|
||||
moe_combine_kernel<c10::BFloat16>
|
||||
<<<N, kCombineBlockSize, 0, stream>>>(gemm2.data_ptr<c10::BFloat16>(),
|
||||
rw.data_ptr<float>(),
|
||||
output.data_ptr<c10::BFloat16>(),
|
||||
N,
|
||||
topk,
|
||||
H);
|
||||
} else {
|
||||
moe_combine_kernel<float>
|
||||
<<<N, kCombineBlockSize, 0, stream>>>(gemm2.data_ptr<float>(),
|
||||
rw.data_ptr<float>(),
|
||||
output.data_ptr<float>(),
|
||||
N,
|
||||
topk,
|
||||
H);
|
||||
}
|
||||
|
||||
return output;
|
||||
}
|
||||
|
||||
} // namespace xllm::kernel::cuda
|
||||
155
ex_engine/xllm_kernels/cuda/moe/moe_compute_index.cu
Normal file
155
ex_engine/xllm_kernels/cuda/moe/moe_compute_index.cu
Normal file
@@ -0,0 +1,155 @@
|
||||
/* Copyright 2025-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.
|
||||
==============================================================================*/
|
||||
|
||||
// Fused MoE token index computation — 3 kernels replacing:
|
||||
// torch::bincount + 2 × torch::argsort + torch::cumsum + CPU sync
|
||||
//
|
||||
// Phase 1 histogram: atomicAdd per-expert token counts
|
||||
// Phase 2 prefix_sum: 1 block, exclusive scan → expert_offsets
|
||||
// Phase 3 place_indices: atomicAdd on offsets, write dst_src + src_dst
|
||||
//
|
||||
// expert_sizes = per-expert token count [num_experts] (preserved)
|
||||
// expert_offsets = exclusive prefix sum of counts (scratch, reused)
|
||||
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include <cub/block/block_scan.cuh>
|
||||
|
||||
#include "kernels/cuda/cuda_ops_api.h"
|
||||
|
||||
namespace xllm::kernel::cuda {
|
||||
|
||||
constexpr int32_t kMoeIndexBlock = 256;
|
||||
|
||||
// ---- Phase 1: histogram ----
|
||||
__global__ void
|
||||
#ifdef USE_DCU
|
||||
__launch_bounds__(kMoeIndexBlock, 1)
|
||||
#endif
|
||||
moe_histogram_kernel(const int32_t* __restrict__ expert_id,
|
||||
int32_t* __restrict__ expert_sizes,
|
||||
int64_t num_elements,
|
||||
int32_t num_experts) {
|
||||
int64_t tid = int64_t(blockIdx.x) * kMoeIndexBlock + threadIdx.x;
|
||||
if (tid < num_elements) {
|
||||
int32_t eid = expert_id[tid];
|
||||
if (eid >= 0 && eid < num_experts) {
|
||||
atomicAdd(&expert_sizes[eid], 1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Phase 2: exclusive prefix sum (1 block) ----
|
||||
// input: expert_sizes (per-expert counts)
|
||||
// output: expert_offsets (exclusive scan of counts)
|
||||
// total_out (total number of tokens, scalar)
|
||||
__global__ void
|
||||
#ifdef USE_DCU
|
||||
__launch_bounds__(kMoeIndexBlock, 1)
|
||||
#endif
|
||||
moe_prefix_sum_kernel(const int32_t* __restrict__ expert_sizes,
|
||||
int32_t* __restrict__ expert_offsets,
|
||||
int32_t num_experts,
|
||||
int64_t* __restrict__ total_out) {
|
||||
using BlockScan = cub::BlockScan<int32_t, kMoeIndexBlock>;
|
||||
__shared__ typename BlockScan::TempStorage s_scan;
|
||||
|
||||
int32_t val = (threadIdx.x < num_experts) ? expert_sizes[threadIdx.x] : 0;
|
||||
int32_t offset;
|
||||
BlockScan(s_scan).ExclusiveSum(val, offset);
|
||||
__syncthreads();
|
||||
|
||||
// total = all elements sum = last thread's exclusive output + its input
|
||||
int32_t total = offset + val;
|
||||
|
||||
if (threadIdx.x < num_experts) {
|
||||
expert_offsets[threadIdx.x] = offset;
|
||||
}
|
||||
if (threadIdx.x == 0 && total_out != nullptr) {
|
||||
*total_out = total;
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Phase 3: place indices ----
|
||||
// atomicAdd on expert_offsets to assign a unique position within
|
||||
// [start(e), start(e)+count(e)), then write both direction mappings.
|
||||
__global__ void
|
||||
#ifdef USE_DCU
|
||||
__launch_bounds__(kMoeIndexBlock, 1)
|
||||
#endif
|
||||
moe_place_indices_kernel(const int32_t* __restrict__ expert_id,
|
||||
int32_t* __restrict__ expert_offsets,
|
||||
int32_t* __restrict__ dst_src,
|
||||
int32_t* __restrict__ src_dst,
|
||||
int64_t num_elements,
|
||||
int32_t num_experts) {
|
||||
int64_t flat_idx = int64_t(blockIdx.x) * kMoeIndexBlock + threadIdx.x;
|
||||
if (flat_idx >= num_elements) return;
|
||||
|
||||
int32_t eid = expert_id[flat_idx];
|
||||
if (eid < 0 || eid >= num_experts) return;
|
||||
|
||||
int32_t pos = atomicAdd(&expert_offsets[eid], 1);
|
||||
dst_src[pos] = static_cast<int32_t>(flat_idx);
|
||||
src_dst[flat_idx] = pos;
|
||||
}
|
||||
|
||||
// ---- Host-side orchestrator ----
|
||||
// Returns {src_dst, dst_src, expert_sizes}
|
||||
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> moe_compute_index(
|
||||
const torch::Tensor& expert_id,
|
||||
int64_t num_experts) {
|
||||
auto device = expert_id.device();
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
int64_t N = expert_id.numel();
|
||||
int32_t E = static_cast<int32_t>(num_experts);
|
||||
CHECK_LE(E, kMoeIndexBlock) << "num_experts cannot exceed " << kMoeIndexBlock;
|
||||
auto expert_id_i32 = expert_id.to(torch::kInt32).contiguous();
|
||||
auto opt_i32 = expert_id_i32.options();
|
||||
|
||||
auto expert_sizes = torch::zeros({num_experts}, opt_i32);
|
||||
auto expert_offsets = torch::empty({num_experts}, opt_i32);
|
||||
auto dst_src = torch::empty({N}, opt_i32);
|
||||
auto src_dst = torch::empty({N}, opt_i32);
|
||||
|
||||
int64_t grid = (N + kMoeIndexBlock - 1) / kMoeIndexBlock;
|
||||
|
||||
// Phase 1: histogram
|
||||
moe_histogram_kernel<<<grid, kMoeIndexBlock, 0, stream>>>(
|
||||
expert_id_i32.data_ptr<int32_t>(),
|
||||
expert_sizes.data_ptr<int32_t>(),
|
||||
N,
|
||||
E);
|
||||
|
||||
// Phase 2: prefix sum (1 block)
|
||||
moe_prefix_sum_kernel<<<1, kMoeIndexBlock, 0, stream>>>(
|
||||
expert_sizes.data_ptr<int32_t>(),
|
||||
expert_offsets.data_ptr<int32_t>(),
|
||||
E,
|
||||
nullptr);
|
||||
|
||||
// Phase 3: place indices
|
||||
moe_place_indices_kernel<<<grid, kMoeIndexBlock, 0, stream>>>(
|
||||
expert_id_i32.data_ptr<int32_t>(),
|
||||
expert_offsets.data_ptr<int32_t>(),
|
||||
dst_src.data_ptr<int32_t>(),
|
||||
src_dst.data_ptr<int32_t>(),
|
||||
N,
|
||||
E);
|
||||
|
||||
return std::make_tuple(src_dst, dst_src, expert_sizes);
|
||||
}
|
||||
|
||||
} // namespace xllm::kernel::cuda
|
||||
59
ex_engine/xllm_kernels/cuda/moe/moe_fused_topk.cu
Normal file
59
ex_engine/xllm_kernels/cuda/moe/moe_fused_topk.cu
Normal file
@@ -0,0 +1,59 @@
|
||||
/* Copyright 2025-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.
|
||||
==============================================================================*/
|
||||
#if defined(USE_DCU)
|
||||
#include "kernels/dcu/dcu_ops_api.h"
|
||||
#else
|
||||
#include "kernels/cuda/cuda_ops_api.h"
|
||||
#endif
|
||||
#include "moe_topk_sigmoid_kernels.cuh"
|
||||
#include "moe_topk_softmax_kernels.cuh"
|
||||
|
||||
namespace xllm::kernel::cuda {
|
||||
|
||||
std::tuple<torch::Tensor, torch::Tensor> moe_fused_topk(
|
||||
torch::Tensor& gating_output,
|
||||
int64_t topk,
|
||||
bool renormalize,
|
||||
const std::optional<torch::Tensor>& correction_bias,
|
||||
const std::string& scoring_func) {
|
||||
int64_t num_tokens = gating_output.size(0);
|
||||
|
||||
torch::Tensor topk_weights = torch::empty(
|
||||
{num_tokens, topk},
|
||||
torch::dtype(torch::kFloat32).device(gating_output.device()));
|
||||
torch::Tensor topk_ids =
|
||||
torch::empty({num_tokens, topk},
|
||||
torch::dtype(torch::kInt32).device(gating_output.device()));
|
||||
|
||||
if (scoring_func == "softmax") {
|
||||
std::optional<torch::Tensor> none_correction_bias = std::nullopt;
|
||||
topk_softmax(topk_weights,
|
||||
topk_ids,
|
||||
gating_output,
|
||||
renormalize,
|
||||
/*moe_softcapping=*/0.0,
|
||||
none_correction_bias);
|
||||
} else if (scoring_func == "sigmoid") {
|
||||
topk_sigmoid(
|
||||
topk_weights, topk_ids, gating_output, renormalize, correction_bias);
|
||||
} else {
|
||||
LOG(FATAL) << "Unsupported scoring function for moe topk: " << scoring_func
|
||||
<< "only softmax and sigmoid are supported";
|
||||
}
|
||||
|
||||
return std::make_tuple(topk_weights, topk_ids);
|
||||
}
|
||||
|
||||
} // namespace xllm::kernel::cuda
|
||||
345
ex_engine/xllm_kernels/cuda/moe/moe_topk.cuh
Normal file
345
ex_engine/xllm_kernels/cuda/moe/moe_topk.cuh
Normal file
@@ -0,0 +1,345 @@
|
||||
|
||||
/*
|
||||
* Copyright (c) 2025, NVIDIA CORPORATION. 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
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* 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.
|
||||
*/
|
||||
|
||||
// refers to
|
||||
// https://github.com/NVIDIA/TensorRT-LLM/blob/main/cpp/tensorrt_llm/kernels/moeTopKFuncs.cuh
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cooperative_groups.h>
|
||||
#if !defined(USE_DCU)
|
||||
#include <cooperative_groups/reduce.h>
|
||||
#endif
|
||||
|
||||
#if defined(USE_MACA)
|
||||
#include <cuda_bf16.h>
|
||||
#endif
|
||||
|
||||
#if !defined(USE_DCU)
|
||||
#include <cub/cub.cuh>
|
||||
#else
|
||||
#include <hipcub/hipcub.hpp>
|
||||
#endif
|
||||
|
||||
#include "core/kernels/cuda/arch_condition.h"
|
||||
|
||||
#if defined(USE_DCU)
|
||||
#include <hip/hip_bfloat16.h>
|
||||
#include <hip/hip_fp16.h>
|
||||
#endif
|
||||
|
||||
#include "core/kernels/cuda/device_utils.cuh"
|
||||
|
||||
namespace xllm::kernel::cuda {
|
||||
namespace reduce_topk {
|
||||
namespace cg = cooperative_groups;
|
||||
static constexpr int kWarpSize = 32;
|
||||
#if !defined(USE_DCU)
|
||||
static constexpr bool kTllmGenHasFastRedux = arch::is_major_v<10>;
|
||||
#else
|
||||
static constexpr bool kTllmGenHasFastRedux = false;
|
||||
#endif
|
||||
|
||||
template <typename T_>
|
||||
struct TopKRedType {
|
||||
using T = T_;
|
||||
static_assert(
|
||||
std::is_same_v<T, float> || std::is_same_v<T, half> ||
|
||||
std::is_same_v<T, BFloat16Type> || std::is_same_v<T, int>,
|
||||
"Top K reduction only implemented for int, float, float16 and bfloat16");
|
||||
|
||||
using TypeCmp = std::conditional_t<sizeof(T) == 4, uint64_t, uint32_t>;
|
||||
using IdxT = std::conditional_t<sizeof(T) == 4, int32_t, int16_t>;
|
||||
#if defined(USE_DCU)
|
||||
using UnsignedBits = std::conditional_t<sizeof(T) == 4, uint32_t, uint16_t>;
|
||||
#endif
|
||||
|
||||
static constexpr int kMoveBits = (sizeof(T) == 4) ? 32 : 16;
|
||||
static constexpr int kMaxIdx = 65535;
|
||||
TypeCmp compValIdx;
|
||||
|
||||
static __host__ __device__ inline TypeCmp makeCmpVal(T val, int32_t idx = 0) {
|
||||
#if !defined(USE_DCU)
|
||||
auto valueBits = cub::Traits<T>::TwiddleIn(
|
||||
reinterpret_cast<typename cub::Traits<T>::UnsignedBits&>(val));
|
||||
#else
|
||||
UnsignedBits valueBits = reinterpret_cast<UnsignedBits&>(val);
|
||||
constexpr UnsignedBits kSignMask =
|
||||
static_cast<UnsignedBits>(UnsignedBits{1} << (sizeof(T) * 8 - 1));
|
||||
if constexpr (std::is_same_v<T, int>) {
|
||||
valueBits = static_cast<UnsignedBits>(valueBits ^ kSignMask);
|
||||
} else {
|
||||
valueBits = (valueBits & kSignMask)
|
||||
? static_cast<UnsignedBits>(~valueBits)
|
||||
: static_cast<UnsignedBits>(valueBits ^ kSignMask);
|
||||
}
|
||||
#endif
|
||||
TypeCmp compactTmp = valueBits;
|
||||
compactTmp = (compactTmp << kMoveBits) | (0xFFFF & (kMaxIdx - idx));
|
||||
// Use 65535 minus idx to give higher priority to elements with smaller
|
||||
// indices.
|
||||
return compactTmp;
|
||||
}
|
||||
|
||||
static __host__ __device__ void unpack(T& value,
|
||||
int32_t& index,
|
||||
TypeCmp cmp) {
|
||||
// Since "65535-idx" is always smaller than 65536 and positive, we can
|
||||
// directly use it as the lower 16 bits
|
||||
index = kMaxIdx - static_cast<int32_t>((cmp & 0xFFFF));
|
||||
|
||||
auto compactTmp = cmp >> kMoveBits;
|
||||
#if !defined(USE_DCU)
|
||||
auto valueBits = cub::Traits<T>::TwiddleOut(
|
||||
reinterpret_cast<typename cub::Traits<T>::UnsignedBits&>(compactTmp));
|
||||
#else
|
||||
UnsignedBits valueBits = static_cast<UnsignedBits>(compactTmp);
|
||||
constexpr UnsignedBits kSignMask =
|
||||
static_cast<UnsignedBits>(UnsignedBits{1} << (sizeof(T) * 8 - 1));
|
||||
if constexpr (std::is_same_v<T, int>) {
|
||||
valueBits = static_cast<UnsignedBits>(valueBits ^ kSignMask);
|
||||
} else {
|
||||
valueBits = (valueBits & kSignMask)
|
||||
? static_cast<UnsignedBits>(valueBits ^ kSignMask)
|
||||
: static_cast<UnsignedBits>(~valueBits);
|
||||
}
|
||||
#endif
|
||||
value = reinterpret_cast<T&>(valueBits);
|
||||
}
|
||||
|
||||
__host__ __device__ TopKRedType() = default;
|
||||
|
||||
__host__ __device__ TopKRedType(T val, int32_t idx)
|
||||
: compValIdx(makeCmpVal(val, idx)) {}
|
||||
|
||||
__host__ __device__ operator TypeCmp() const noexcept { return compValIdx; }
|
||||
|
||||
__device__ inline TypeCmp reduce(
|
||||
cg::thread_block_tile<kWarpSize> const& warp) {
|
||||
#if defined(USE_DCU)
|
||||
TypeCmp result = compValIdx;
|
||||
#pragma unroll
|
||||
for (int offset = kWarpSize / 2; offset > 0; offset >>= 1) {
|
||||
TypeCmp other = warp.shfl_down(result, offset);
|
||||
result = other > result ? other : result;
|
||||
}
|
||||
return warp.shfl(result, 0);
|
||||
#else
|
||||
if constexpr (!kTllmGenHasFastRedux || sizeof(TypeCmp) == 8) {
|
||||
return cg::reduce(warp, compValIdx, cg::greater<TypeCmp>{});
|
||||
} else {
|
||||
TypeCmp result;
|
||||
asm("redux.sync.max.u32 %0, %1, 0xffffffff;\n"
|
||||
: "=r"(result)
|
||||
: "r"(compValIdx));
|
||||
return result;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <int K_, bool Enable_>
|
||||
struct TopKIdx {
|
||||
// by default, empty
|
||||
};
|
||||
|
||||
template <int K_>
|
||||
struct TopKIdx<K_, true> {
|
||||
static constexpr int K = K_;
|
||||
int32_t val[K];
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#define TOPK_SWAP(I, J) \
|
||||
{ \
|
||||
auto pairMin = min(topK[I].compValIdx, topK[J].compValIdx); \
|
||||
auto pairMax = max(topK[I].compValIdx, topK[J].compValIdx); \
|
||||
topK[I].compValIdx = pairMax; \
|
||||
topK[J].compValIdx = pairMin; \
|
||||
}
|
||||
|
||||
template <int N, typename RedType>
|
||||
struct Sort;
|
||||
|
||||
template <typename RedType>
|
||||
struct Sort<1, RedType> {
|
||||
static __device__ void run(RedType* topK) {}
|
||||
};
|
||||
|
||||
template <typename RedType>
|
||||
struct Sort<2, RedType> {
|
||||
static __device__ void run(RedType* topK) { TOPK_SWAP(0, 1); }
|
||||
};
|
||||
|
||||
template <typename RedType>
|
||||
struct Sort<3, RedType> {
|
||||
static __device__ void run(RedType* topK) {
|
||||
TOPK_SWAP(0, 1);
|
||||
TOPK_SWAP(1, 2);
|
||||
TOPK_SWAP(0, 1);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename RedType>
|
||||
struct Sort<4, RedType> {
|
||||
static __device__ void run(RedType* topK) {
|
||||
TOPK_SWAP(0, 2);
|
||||
TOPK_SWAP(1, 3);
|
||||
TOPK_SWAP(0, 1);
|
||||
TOPK_SWAP(2, 3);
|
||||
TOPK_SWAP(1, 2);
|
||||
}
|
||||
};
|
||||
|
||||
template <int K, typename Type>
|
||||
__forceinline__ __device__ void reduceTopK(
|
||||
cg::thread_block_tile<kWarpSize> const& warp,
|
||||
Type (&out)[K],
|
||||
int32_t (&outIdx)[K],
|
||||
Type value,
|
||||
int32_t idx,
|
||||
Type const minValue,
|
||||
int actualK = K) {
|
||||
static_assert(K > 0, "Top K must have K > 0");
|
||||
static_assert(K < kWarpSize, "Top K must have K < kWarpSize");
|
||||
using RedType = TopKRedType<Type>;
|
||||
RedType topK{value, idx};
|
||||
typename RedType::TypeCmp packedMax{};
|
||||
#pragma unroll
|
||||
for (int kk = 0; kk < actualK; ++kk) //@todo: check if actualK is correct
|
||||
{
|
||||
topK =
|
||||
kk > 0 && packedMax == topK.compValIdx ? RedType{minValue, idx} : topK;
|
||||
// get the next largest value
|
||||
packedMax = topK.reduce(warp);
|
||||
RedType::unpack(out[kk], outIdx[kk], packedMax);
|
||||
}
|
||||
};
|
||||
|
||||
template <int K, typename Type, int N, bool IsSorted = false>
|
||||
__device__ void reduceTopKFunc(cg::thread_block_tile<kWarpSize> const& warp,
|
||||
Type (&out)[K],
|
||||
int32_t (&outIdx)[K],
|
||||
Type (&value)[N],
|
||||
int32_t (&idx)[N],
|
||||
Type minValue,
|
||||
int actualK = K) {
|
||||
static_assert(K > 0, "Top K must have K > 0");
|
||||
static_assert(K < kWarpSize, "Top K must have K < kWarpSize");
|
||||
static_assert(N > 0, "Top K must have N > 0");
|
||||
static_assert(N < 5,
|
||||
"Only support candidates number less than or equal to 128");
|
||||
using RedType = TopKRedType<Type>;
|
||||
RedType topK[N];
|
||||
#pragma unroll
|
||||
for (int nn = 0; nn < N; ++nn) {
|
||||
topK[nn] = RedType{value[nn], idx[nn]};
|
||||
}
|
||||
|
||||
if constexpr (!IsSorted) {
|
||||
Sort<N, RedType>::run(topK);
|
||||
}
|
||||
typename RedType::TypeCmp packedMax{};
|
||||
#pragma unroll
|
||||
for (int kk = 0; kk < actualK; ++kk) {
|
||||
bool update = kk > 0 && packedMax == topK[0].compValIdx;
|
||||
#pragma unroll
|
||||
for (int nn = 0; nn < N; ++nn) {
|
||||
topK[nn] = update && nn == N - 1 ? RedType{minValue, idx[nn]}
|
||||
: update ? topK[nn + 1]
|
||||
: topK[nn];
|
||||
}
|
||||
// get the next largest value
|
||||
packedMax = topK[0].reduce(warp);
|
||||
RedType::unpack(out[kk], outIdx[kk], packedMax);
|
||||
}
|
||||
};
|
||||
|
||||
template <int K, typename Type, int N>
|
||||
__forceinline__ __device__ void reduceTopK(
|
||||
cg::thread_block_tile<kWarpSize> const& warp,
|
||||
Type (&out)[K],
|
||||
int32_t (&outIdx)[K],
|
||||
Type (&value)[N],
|
||||
int32_t (&idx)[N],
|
||||
Type const minValue,
|
||||
int actualK = K) {
|
||||
static_assert(K > 0, "Top K must have K > 0");
|
||||
static_assert(K < kWarpSize, "Top K must have K < kWarpSize");
|
||||
static_assert(N > 0, "Top K must have N > 0");
|
||||
static_assert(
|
||||
N <= 16,
|
||||
"Only support candidates number less than or equal to 16*32=512");
|
||||
static_assert(N <= 4 || N % 4 == 0,
|
||||
"Only support candidates number is a multiple of 4*32=128 or "
|
||||
"less than or equal to 4");
|
||||
using RedType = TopKRedType<Type>;
|
||||
|
||||
if constexpr (N <= 4) {
|
||||
reduceTopKFunc<K, Type, N>(
|
||||
warp, out, outIdx, value, idx, minValue, actualK);
|
||||
} else {
|
||||
constexpr int kNumLoops = N / 4;
|
||||
constexpr int kNumResults = (kNumLoops * K - 1) / kWarpSize + 1;
|
||||
|
||||
Type topKBufferValue[kNumResults];
|
||||
int32_t topKBufferIdx[kNumResults];
|
||||
int32_t laneIdx = threadIdx.x % kWarpSize;
|
||||
|
||||
// Sentinel index must be in [0, kMaxIdx] to survive makeCmpVal pack/unpack
|
||||
// (kMaxIdx - idx is stored in 16 bits; -1 would become 0 and unpack to
|
||||
// 65535). Use kMaxIdx so sentinel slots have smallest compValIdx for
|
||||
// minValue and lose to any real candidate.
|
||||
for (int ii = 0; ii < kNumResults; ++ii) {
|
||||
topKBufferValue[ii] = minValue;
|
||||
topKBufferIdx[ii] = RedType::kMaxIdx;
|
||||
}
|
||||
for (int loop = 0; loop < kNumLoops; ++loop) {
|
||||
int start = loop * 4;
|
||||
Type topKValue[K];
|
||||
int32_t topKIdx[K];
|
||||
Type inValue[4];
|
||||
int32_t inIdx[4];
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
inValue[i] = value[start + i];
|
||||
inIdx[i] = idx[start + i];
|
||||
}
|
||||
reduceTopKFunc<K, Type, 4>(
|
||||
warp, topKValue, topKIdx, inValue, inIdx, minValue, actualK);
|
||||
int inOffset = laneIdx % K;
|
||||
if (laneIdx >= loop * K && laneIdx < (loop + 1) * K) {
|
||||
topKBufferValue[0] = topKValue[inOffset];
|
||||
topKBufferIdx[0] = topKIdx[inOffset];
|
||||
}
|
||||
if (loop == kNumLoops - 1 && (laneIdx < (kNumLoops * K - kWarpSize))) {
|
||||
topKBufferValue[1] = topKValue[inOffset];
|
||||
topKBufferIdx[1] = topKIdx[inOffset];
|
||||
}
|
||||
}
|
||||
|
||||
reduceTopKFunc<K, Type, kNumResults>(
|
||||
warp, out, outIdx, topKBufferValue, topKBufferIdx, minValue, actualK);
|
||||
}
|
||||
};
|
||||
|
||||
#undef TOPK_SWAP
|
||||
|
||||
} // namespace reduce_topk
|
||||
} // namespace xllm::kernel::cuda
|
||||
609
ex_engine/xllm_kernels/cuda/moe/moe_topk_sigmoid_kernels.cuh
Normal file
609
ex_engine/xllm_kernels/cuda/moe/moe_topk_sigmoid_kernels.cuh
Normal file
@@ -0,0 +1,609 @@
|
||||
// Adapt from
|
||||
// https://github.com/vllm-project/vllm/blob/v0.7.3/csrc/moe/topk_softmax_kernels.cu
|
||||
// which is originally adapted from
|
||||
// https://github.com/NVIDIA/TensorRT-LLM/blob/v0.7.1/cpp/tensorrt_llm/kernels/mixtureOfExperts/moe_kernels.cu
|
||||
/* Copyright 2025 SGLang Team. 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
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
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.
|
||||
==============================================================================*/
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <torch/all.h>
|
||||
|
||||
#include <cub/util_type.cuh>
|
||||
|
||||
#if !defined(USE_DCU) && !defined(USE_MACA)
|
||||
#include <cuda/functional>
|
||||
#endif
|
||||
|
||||
#include "kernels/cuda/device_utils.cuh"
|
||||
|
||||
namespace {
|
||||
|
||||
using namespace xllm::kernel::cuda;
|
||||
|
||||
#if defined(USE_DCU)
|
||||
static constexpr unsigned long long kSigmoidFullMask = 0xffffffffffffffffULL;
|
||||
#else
|
||||
static constexpr unsigned int kSigmoidFullMask = 0xffffffffU;
|
||||
#endif
|
||||
|
||||
// ====================== Sigmoid things ===============================
|
||||
// We have our own implementation of sigmoid here so we can support transposing
|
||||
// the output in the sigmoid kernel when we extend this module to support
|
||||
// expert-choice routing.
|
||||
template <typename T, int TPB>
|
||||
__launch_bounds__(TPB) __global__
|
||||
void moe_sigmoid(const T* input,
|
||||
const bool* finished,
|
||||
float* output,
|
||||
const int num_cols,
|
||||
const float* correction_bias) {
|
||||
const int thread_row_offset = blockIdx.x * num_cols;
|
||||
|
||||
// Don't touch finished rows.
|
||||
if ((finished != nullptr) && finished[blockIdx.x]) {
|
||||
return;
|
||||
}
|
||||
|
||||
// First pass: Apply transformation, find max, and write transformed values to
|
||||
// output
|
||||
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||
const int idx = thread_row_offset + ii;
|
||||
float val = convert_to_float<T>(input[idx]);
|
||||
|
||||
val = 1.0f / (1.0f + expf(-val));
|
||||
|
||||
// Apply correction bias if provided
|
||||
if (correction_bias != nullptr) {
|
||||
val = val + correction_bias[ii];
|
||||
}
|
||||
|
||||
output[idx] = val; // Store transformed value
|
||||
}
|
||||
}
|
||||
|
||||
template <int TPB>
|
||||
__launch_bounds__(TPB) __global__
|
||||
void moe_topK(const float* inputs_after_sigmoid,
|
||||
const bool* finished,
|
||||
float* output,
|
||||
int* indices,
|
||||
const int num_experts,
|
||||
const int k,
|
||||
const int start_expert,
|
||||
const int end_expert,
|
||||
const bool renormalize,
|
||||
const float* correction_bias) {
|
||||
using cub_kvp = cub::KeyValuePair<int, float>;
|
||||
using BlockReduce = cub::BlockReduce<cub_kvp, TPB>;
|
||||
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||
|
||||
cub_kvp thread_kvp;
|
||||
cub::ArgMax arg_max;
|
||||
|
||||
const int block_row = blockIdx.x;
|
||||
|
||||
const bool row_is_active = finished ? !finished[block_row] : true;
|
||||
const int thread_read_offset = blockIdx.x * num_experts;
|
||||
float row_sum_for_renormalize = 0;
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
thread_kvp.key = 0;
|
||||
thread_kvp.value = -1.f; // This is OK because inputs are probabilities
|
||||
|
||||
cub_kvp inp_kvp;
|
||||
for (int expert = threadIdx.x; expert < num_experts; expert += TPB) {
|
||||
const int idx = thread_read_offset + expert;
|
||||
inp_kvp.key = expert;
|
||||
inp_kvp.value = inputs_after_sigmoid[idx];
|
||||
|
||||
for (int prior_k = 0; prior_k < k_idx; ++prior_k) {
|
||||
const int prior_winning_expert = indices[k * block_row + prior_k];
|
||||
|
||||
if (prior_winning_expert == expert) {
|
||||
inp_kvp = thread_kvp;
|
||||
}
|
||||
}
|
||||
|
||||
thread_kvp = arg_max(inp_kvp, thread_kvp);
|
||||
}
|
||||
|
||||
const cub_kvp result_kvp =
|
||||
BlockReduce(tmpStorage).Reduce(thread_kvp, arg_max);
|
||||
if (threadIdx.x == 0) {
|
||||
// Ignore experts the node isn't responsible for with expert parallelism
|
||||
const int expert = result_kvp.key;
|
||||
const bool node_uses_expert =
|
||||
expert >= start_expert && expert < end_expert;
|
||||
const bool should_process_row = row_is_active && node_uses_expert;
|
||||
|
||||
const int idx = k * block_row + k_idx;
|
||||
float val = result_kvp.value;
|
||||
if (correction_bias != nullptr) {
|
||||
val -= correction_bias[expert];
|
||||
}
|
||||
output[idx] = val;
|
||||
indices[idx] = should_process_row ? (expert - start_expert) : num_experts;
|
||||
assert(indices[idx] >= 0);
|
||||
row_sum_for_renormalize += val;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
if (renormalize && threadIdx.x == 0) {
|
||||
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
const int idx = k * block_row + k_idx;
|
||||
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ====================== TopK sigmoid things ===============================
|
||||
|
||||
/*
|
||||
A Top-K gating sigmoid written to exploit when the number of experts in the
|
||||
MoE layers are a small power of 2. This allows us to cleanly share the rows
|
||||
among the threads in a single warp and eliminate communication between warps
|
||||
(so no need to use shared mem).
|
||||
|
||||
It fuses the sigmoid, max and argmax into a single kernel.
|
||||
|
||||
Limitations:
|
||||
1) This implementation is intended for when the number of experts is a small
|
||||
power of 2. 2) This implementation assumes k is small, but will work for any
|
||||
k.
|
||||
*/
|
||||
|
||||
template <typename T,
|
||||
int VPT,
|
||||
int NUM_EXPERTS,
|
||||
int WARPS_PER_CTA,
|
||||
int BYTES_PER_LDG>
|
||||
__launch_bounds__(WARPS_PER_CTA* WARP_SIZE) __global__
|
||||
void topk_gating_sigmoid(const T* input,
|
||||
const bool* finished,
|
||||
float* output,
|
||||
const int num_rows,
|
||||
int* indices,
|
||||
const int k,
|
||||
const int start_expert,
|
||||
const int end_expert,
|
||||
const bool renormalize,
|
||||
const float* correction_bias) {
|
||||
// We begin by enforcing compile time assertions and setting up compile time
|
||||
// constants.
|
||||
static_assert(VPT == (VPT & -VPT), "VPT must be power of 2");
|
||||
static_assert(NUM_EXPERTS == (NUM_EXPERTS & -NUM_EXPERTS),
|
||||
"NUM_EXPERTS must be power of 2");
|
||||
static_assert(BYTES_PER_LDG == (BYTES_PER_LDG & -BYTES_PER_LDG),
|
||||
"BYTES_PER_LDG must be power of 2");
|
||||
static_assert(BYTES_PER_LDG <= 16, "BYTES_PER_LDG must be leq 16");
|
||||
|
||||
// Number of bytes each thread pulls in per load
|
||||
static constexpr int kEltsPerLdg = BYTES_PER_LDG / sizeof(T);
|
||||
static constexpr int kEltsPerRow = NUM_EXPERTS;
|
||||
static constexpr int kThreadsPerRow = kEltsPerRow / VPT;
|
||||
static constexpr int kLdgPerThread = VPT / kEltsPerLdg;
|
||||
|
||||
// Restrictions based on previous section.
|
||||
static_assert(
|
||||
VPT % kEltsPerLdg == 0,
|
||||
"The elements per thread must be a multiple of the elements per ldg");
|
||||
static_assert(WARP_SIZE % kThreadsPerRow == 0,
|
||||
"The threads per row must cleanly divide the threads per warp");
|
||||
static_assert(kThreadsPerRow == (kThreadsPerRow & -kThreadsPerRow),
|
||||
"THREADS_PER_ROW must be power of 2");
|
||||
static_assert(kThreadsPerRow <= WARP_SIZE,
|
||||
"THREADS_PER_ROW can be at most warp size");
|
||||
|
||||
// We have NUM_EXPERTS elements per row. We specialize for small #experts
|
||||
static constexpr int kEltsPerWarp = WARP_SIZE * VPT;
|
||||
static constexpr int kRowsPerWarp = kEltsPerWarp / kEltsPerRow;
|
||||
static constexpr int kRowsPerCta = WARPS_PER_CTA * kRowsPerWarp;
|
||||
|
||||
// Restrictions for previous section.
|
||||
static_assert(kEltsPerWarp % kEltsPerRow == 0,
|
||||
"The elts per row must cleanly divide the total elt per warp");
|
||||
|
||||
// ===================== From this point, we finally start computing run-time
|
||||
// variables. ========================
|
||||
|
||||
// Compute CTA and warp rows. We pack multiple rows into a single warp, and a
|
||||
// block contains WARPS_PER_CTA warps. This, each block processes a chunk of
|
||||
// rows. We start by computing the start row for each block.
|
||||
const int cta_base_row = blockIdx.x * kRowsPerCta;
|
||||
|
||||
// Now, using the base row per thread block, we compute the base row per warp.
|
||||
const int warp_base_row = cta_base_row + threadIdx.y * kRowsPerWarp;
|
||||
|
||||
// The threads in a warp are split into sub-groups that will work on a row.
|
||||
// We compute row offset for each thread sub-group
|
||||
const int thread_row_in_warp = threadIdx.x / kThreadsPerRow;
|
||||
const int thread_row = warp_base_row + thread_row_in_warp;
|
||||
|
||||
// Threads with indices out of bounds should early exit here.
|
||||
if (thread_row >= num_rows) {
|
||||
return;
|
||||
}
|
||||
const bool row_is_active = finished ? !finished[thread_row] : true;
|
||||
|
||||
// We finally start setting up the read pointers for each thread. First, each
|
||||
// thread jumps to the start of the row it will read.
|
||||
const T* thread_row_ptr = input + thread_row * kEltsPerRow;
|
||||
|
||||
// Now, we compute the group each thread belong to in order to determine the
|
||||
// first column to start loads.
|
||||
const int thread_group_idx = threadIdx.x % kThreadsPerRow;
|
||||
const int first_elt_read_by_thread = thread_group_idx * kEltsPerLdg;
|
||||
const T* thread_read_ptr = thread_row_ptr + first_elt_read_by_thread;
|
||||
|
||||
// Determine the pointer type to use to read in the data depending on the
|
||||
// BYTES_PER_LDG template param. In theory, this can support all powers of 2
|
||||
// up to 16. NOTE(woosuk): The original implementation uses CUTLASS aligned
|
||||
// array here. We defined our own aligned array and use it here to avoid the
|
||||
// dependency on CUTLASS.
|
||||
using AccessType = AlignedArray<T, kEltsPerLdg>;
|
||||
|
||||
// Finally, we pull in the data from global mem
|
||||
T row_chunk_temp[VPT];
|
||||
AccessType* row_chunk_vec_ptr =
|
||||
reinterpret_cast<AccessType*>(&row_chunk_temp);
|
||||
const AccessType* vec_thread_read_ptr =
|
||||
reinterpret_cast<const AccessType*>(thread_read_ptr);
|
||||
#pragma unroll
|
||||
// Note(Byron): interleaved loads to achieve better memory coalescing
|
||||
// | thread[0] | thread[1] | thread[2] | thread[3] | thread[0] | thread[1] |
|
||||
// thread[2] | thread[3] | ...
|
||||
for (int ii = 0; ii < kLdgPerThread; ++ii) {
|
||||
row_chunk_vec_ptr[ii] = vec_thread_read_ptr[ii * kThreadsPerRow];
|
||||
}
|
||||
|
||||
float row_chunk[VPT];
|
||||
#pragma unroll
|
||||
// Note(Byron): upcast logits to float32
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
float val = convert_to_float<T>(row_chunk_temp[ii]);
|
||||
val = 1.0f / (1.0f + expf(-val));
|
||||
// Apply correction bias if provided
|
||||
if (correction_bias != nullptr) {
|
||||
/*
|
||||
LDG is interleaved
|
||||
|thread0 LDG| |thread1 LDG| |thread0 LDG| |thread1 LDG|
|
||||
|--------- group0 --------| |----------group1 --------|
|
||||
^ local2
|
||||
*/
|
||||
const int group_id = ii / kEltsPerLdg;
|
||||
const int local_id = ii % kEltsPerLdg;
|
||||
const int expert_idx = first_elt_read_by_thread +
|
||||
group_id * kThreadsPerRow * kEltsPerLdg + local_id;
|
||||
val = val + correction_bias[expert_idx];
|
||||
}
|
||||
|
||||
row_chunk[ii] = val;
|
||||
}
|
||||
|
||||
// Now, row_chunk contains the sigmoid of the row chunk. Now, I want to find
|
||||
// the topk elements in each row, along with the max index.
|
||||
int start_col = first_elt_read_by_thread;
|
||||
static constexpr int kColsPerGroupLdg = kEltsPerLdg * kThreadsPerRow;
|
||||
|
||||
float row_sum_for_renormalize = 0;
|
||||
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
// First, each thread does the local argmax
|
||||
float max_val = row_chunk[0];
|
||||
int expert = start_col;
|
||||
#pragma unroll
|
||||
for (int ldg = 0, col = start_col; ldg < kLdgPerThread;
|
||||
++ldg, col += kColsPerGroupLdg) {
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < kEltsPerLdg; ++ii) {
|
||||
float val = row_chunk[ldg * kEltsPerLdg + ii];
|
||||
|
||||
// No check on the experts here since columns with the smallest index
|
||||
// are processed first and only updated if > (not >=)
|
||||
if (val > max_val) {
|
||||
max_val = val;
|
||||
expert = col + ii;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Now, we perform the argmax reduce. We use the butterfly pattern so threads
|
||||
// reach consensus about the max. This will be useful for K > 1 so that the
|
||||
// threads can agree on "who" had the max value. That thread can then blank out
|
||||
// their max with -inf and the warp can run more iterations...
|
||||
#pragma unroll
|
||||
for (int mask = kThreadsPerRow / 2; mask > 0; mask /= 2) {
|
||||
float other_max = XLLM_SHFL_XOR_SYNC_WIDTH(
|
||||
kSigmoidFullMask, max_val, mask, kThreadsPerRow);
|
||||
int other_expert = XLLM_SHFL_XOR_SYNC_WIDTH(
|
||||
kSigmoidFullMask, expert, mask, kThreadsPerRow);
|
||||
|
||||
// We want lower indices to "win" in every thread so we break ties this
|
||||
// way
|
||||
if (other_max > max_val ||
|
||||
(other_max == max_val && other_expert < expert)) {
|
||||
max_val = other_max;
|
||||
expert = other_expert;
|
||||
}
|
||||
}
|
||||
|
||||
// Write the max for this k iteration to global memory.
|
||||
if (thread_group_idx == 0) {
|
||||
// Add a guard to ignore experts not included by this node
|
||||
const bool node_uses_expert =
|
||||
expert >= start_expert && expert < end_expert;
|
||||
const bool should_process_row = row_is_active && node_uses_expert;
|
||||
|
||||
// The lead thread from each sub-group will write out the final results to
|
||||
// global memory. (This will be a single) thread per row of the
|
||||
// input/output matrices.
|
||||
const int idx = k * thread_row + k_idx;
|
||||
if (correction_bias != nullptr) {
|
||||
max_val -= correction_bias[expert];
|
||||
}
|
||||
output[idx] = max_val;
|
||||
indices[idx] = should_process_row ? (expert - start_expert) : NUM_EXPERTS;
|
||||
row_sum_for_renormalize += max_val;
|
||||
}
|
||||
|
||||
// Finally, we clear the value in the thread with the current max if there
|
||||
// is another iteration to run.
|
||||
if (k_idx + 1 < k) {
|
||||
const int ldg_group_for_expert = expert / kColsPerGroupLdg;
|
||||
const int thread_to_clear_in_group =
|
||||
(expert / kEltsPerLdg) % kThreadsPerRow;
|
||||
|
||||
// Only the thread in the group which produced the max will reset the
|
||||
// "winning" value to -inf.
|
||||
if (thread_group_idx == thread_to_clear_in_group) {
|
||||
const int offset_for_expert = expert % kEltsPerLdg;
|
||||
// Safe to set to any negative value since row_chunk values must be
|
||||
// between 0 and 1.
|
||||
row_chunk[ldg_group_for_expert * kEltsPerLdg + offset_for_expert] =
|
||||
-10000.f;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fuse renormalization of topk_weights into this kernel
|
||||
if (renormalize && thread_group_idx == 0) {
|
||||
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||
#pragma unroll
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
const int idx = k * thread_row + k_idx;
|
||||
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, int EXPERTS, int WARPS_PER_TB>
|
||||
void topk_gating_sigmoid_launcher_helper(const T* input,
|
||||
const bool* finished,
|
||||
float* output,
|
||||
int* indices,
|
||||
const int num_rows,
|
||||
const int k,
|
||||
const int start_expert,
|
||||
const int end_expert,
|
||||
const bool renormalize,
|
||||
const float* correction_bias,
|
||||
cudaStream_t stream) {
|
||||
static constexpr std::size_t kMaxBytesPerLdg = 16;
|
||||
|
||||
static constexpr int kBytesPerLdg = MIN(kMaxBytesPerLdg, sizeof(T) * EXPERTS);
|
||||
using Constants = TopkConstants<T, EXPERTS, kBytesPerLdg>;
|
||||
static constexpr int kVpt = Constants::VPT;
|
||||
static constexpr int kRowsPerWarp = Constants::ROWS_PER_WARP;
|
||||
const int num_warps = (num_rows + kRowsPerWarp - 1) / kRowsPerWarp;
|
||||
const int num_blocks = (num_warps + WARPS_PER_TB - 1) / WARPS_PER_TB;
|
||||
|
||||
dim3 block_dim(WARP_SIZE, WARPS_PER_TB);
|
||||
topk_gating_sigmoid<T, kVpt, EXPERTS, WARPS_PER_TB, kBytesPerLdg>
|
||||
<<<num_blocks, block_dim, 0, stream>>>(input,
|
||||
finished,
|
||||
output,
|
||||
num_rows,
|
||||
indices,
|
||||
k,
|
||||
start_expert,
|
||||
end_expert,
|
||||
renormalize,
|
||||
correction_bias);
|
||||
}
|
||||
|
||||
#define LAUNCH_SIGMOID(TYPE, NUM_EXPERTS, WARPS_PER_TB) \
|
||||
topk_gating_sigmoid_launcher_helper<TYPE, NUM_EXPERTS, WARPS_PER_TB>( \
|
||||
gating_output, \
|
||||
nullptr, \
|
||||
topk_weights, \
|
||||
topk_indices, \
|
||||
num_tokens, \
|
||||
topk, \
|
||||
0, \
|
||||
num_experts, \
|
||||
renormalize, \
|
||||
correction_bias, \
|
||||
stream);
|
||||
|
||||
template <typename T>
|
||||
void topk_gating_sigmoid_kernel_launcher(const T* gating_output,
|
||||
float* topk_weights,
|
||||
int* topk_indices,
|
||||
float* sigmoid_workspace,
|
||||
const int num_tokens,
|
||||
const int num_experts,
|
||||
const int topk,
|
||||
const bool renormalize,
|
||||
const float* correction_bias,
|
||||
cudaStream_t stream) {
|
||||
static constexpr int kWarpsPerTb = 4;
|
||||
switch (num_experts) {
|
||||
case 1:
|
||||
LAUNCH_SIGMOID(T, 1, kWarpsPerTb);
|
||||
break;
|
||||
case 2:
|
||||
LAUNCH_SIGMOID(T, 2, kWarpsPerTb);
|
||||
break;
|
||||
case 4:
|
||||
LAUNCH_SIGMOID(T, 4, kWarpsPerTb);
|
||||
break;
|
||||
case 8:
|
||||
LAUNCH_SIGMOID(T, 8, kWarpsPerTb);
|
||||
break;
|
||||
case 16:
|
||||
LAUNCH_SIGMOID(T, 16, kWarpsPerTb);
|
||||
break;
|
||||
case 32:
|
||||
LAUNCH_SIGMOID(T, 32, kWarpsPerTb);
|
||||
break;
|
||||
case 64:
|
||||
LAUNCH_SIGMOID(T, 64, kWarpsPerTb);
|
||||
break;
|
||||
case 128:
|
||||
LAUNCH_SIGMOID(T, 128, kWarpsPerTb);
|
||||
break;
|
||||
case 256:
|
||||
LAUNCH_SIGMOID(T, 256, kWarpsPerTb);
|
||||
break;
|
||||
default: {
|
||||
TORCH_CHECK(sigmoid_workspace != nullptr,
|
||||
"sigmoid_workspace must be provided for num_experts that are "
|
||||
"not a power of 2.");
|
||||
static constexpr int kTpb = 256;
|
||||
moe_sigmoid<T, kTpb><<<num_tokens, kTpb, 0, stream>>>(gating_output,
|
||||
nullptr,
|
||||
sigmoid_workspace,
|
||||
num_experts,
|
||||
correction_bias);
|
||||
moe_topK<kTpb><<<num_tokens, kTpb, 0, stream>>>(sigmoid_workspace,
|
||||
nullptr,
|
||||
topk_weights,
|
||||
topk_indices,
|
||||
num_experts,
|
||||
topk,
|
||||
0,
|
||||
num_experts,
|
||||
renormalize,
|
||||
correction_bias);
|
||||
}
|
||||
}
|
||||
}
|
||||
} // namespace
|
||||
|
||||
namespace xllm::kernel::cuda {
|
||||
void topk_sigmoid(torch::Tensor& topk_weights, // [num_tokens, topk]
|
||||
torch::Tensor& topk_indices, // [num_tokens, topk]
|
||||
torch::Tensor& gating_output, // [num_tokens, num_experts]
|
||||
const bool renormalize,
|
||||
const std::optional<torch::Tensor>& correction_bias) {
|
||||
// Check data type
|
||||
CHECK(gating_output.scalar_type() == at::ScalarType::Float ||
|
||||
gating_output.scalar_type() == at::ScalarType::Half ||
|
||||
gating_output.scalar_type() == at::ScalarType::BFloat16)
|
||||
<< "gating_output must be float32, float16, or bfloat16";
|
||||
|
||||
// Check dimensions
|
||||
CHECK(gating_output.dim() == 2)
|
||||
<< "gating_output must be 2D tensor [num_tokens, num_experts]";
|
||||
CHECK(topk_weights.dim() == 2)
|
||||
<< "topk_weights must be 2D tensor [num_tokens, topk]";
|
||||
CHECK(topk_indices.dim() == 2)
|
||||
<< "topk_indices must be 2D tensor [num_tokens, topk]";
|
||||
|
||||
// Check shapes
|
||||
CHECK(gating_output.size(0) == topk_weights.size(0))
|
||||
<< "First dimension of topk_weights must match num_tokens in "
|
||||
"gating_output";
|
||||
CHECK(gating_output.size(0) == topk_indices.size(0))
|
||||
<< "First dimension of topk_indices must match num_tokens in "
|
||||
"gating_output";
|
||||
CHECK(topk_weights.size(-1) == topk_indices.size(-1))
|
||||
<< "Second dimension of topk_indices must match topk in topk_weights";
|
||||
CHECK(topk_weights.size(-1) <= gating_output.size(-1))
|
||||
<< "topk must be less than or equal to num_experts";
|
||||
|
||||
const int num_experts = static_cast<int>(gating_output.size(-1));
|
||||
const int num_tokens = static_cast<int>(gating_output.size(0));
|
||||
const int topk = static_cast<int>(topk_weights.size(-1));
|
||||
|
||||
const bool is_pow_2 =
|
||||
(num_experts != 0) && ((num_experts & (num_experts - 1)) == 0);
|
||||
const bool needs_workspace = !is_pow_2 || num_experts > 256;
|
||||
const int64_t workspace_size = needs_workspace ? num_tokens * num_experts : 0;
|
||||
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(gating_output));
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
torch::Tensor sigmoid_workspace = torch::empty(
|
||||
{workspace_size}, gating_output.options().dtype(at::ScalarType::Float));
|
||||
|
||||
const at::ScalarType dtype = gating_output.scalar_type();
|
||||
|
||||
// Validate correction_bias if provided - must always be float32
|
||||
const float* bias_ptr = nullptr;
|
||||
if (correction_bias.has_value()) {
|
||||
const torch::Tensor& bias_tensor = correction_bias.value();
|
||||
CHECK(bias_tensor.dim() == 1)
|
||||
<< "correction_bias must be 1D tensor [num_experts]";
|
||||
CHECK(bias_tensor.size(0) == num_experts)
|
||||
<< "correction_bias size must match num_experts";
|
||||
CHECK(bias_tensor.scalar_type() == at::ScalarType::Float)
|
||||
<< "correction_bias must be float32, got " << bias_tensor.scalar_type();
|
||||
bias_ptr = bias_tensor.data_ptr<float>();
|
||||
}
|
||||
|
||||
if (dtype == at::ScalarType::Float) {
|
||||
topk_gating_sigmoid_kernel_launcher<float>(
|
||||
gating_output.data_ptr<float>(),
|
||||
topk_weights.data_ptr<float>(),
|
||||
topk_indices.data_ptr<int>(),
|
||||
sigmoid_workspace.data_ptr<float>(),
|
||||
num_tokens,
|
||||
num_experts,
|
||||
topk,
|
||||
renormalize,
|
||||
bias_ptr,
|
||||
stream);
|
||||
} else if (dtype == at::ScalarType::Half) {
|
||||
topk_gating_sigmoid_kernel_launcher<__half>(
|
||||
reinterpret_cast<const __half*>(gating_output.data_ptr<at::Half>()),
|
||||
topk_weights.data_ptr<float>(),
|
||||
topk_indices.data_ptr<int>(),
|
||||
sigmoid_workspace.data_ptr<float>(),
|
||||
num_tokens,
|
||||
num_experts,
|
||||
topk,
|
||||
renormalize,
|
||||
bias_ptr,
|
||||
stream);
|
||||
} else if (dtype == at::ScalarType::BFloat16) {
|
||||
topk_gating_sigmoid_kernel_launcher<BFloat16Type>(
|
||||
reinterpret_cast<const BFloat16Type*>(
|
||||
gating_output.data_ptr<at::BFloat16>()),
|
||||
topk_weights.data_ptr<float>(),
|
||||
topk_indices.data_ptr<int>(),
|
||||
sigmoid_workspace.data_ptr<float>(),
|
||||
num_tokens,
|
||||
num_experts,
|
||||
topk,
|
||||
renormalize,
|
||||
bias_ptr,
|
||||
stream);
|
||||
} else {
|
||||
LOG(FATAL) << "Unsupported gating_output dtype: " << dtype;
|
||||
}
|
||||
}
|
||||
} // namespace xllm::kernel::cuda
|
||||
867
ex_engine/xllm_kernels/cuda/moe/moe_topk_softmax_kernels.cuh
Normal file
867
ex_engine/xllm_kernels/cuda/moe/moe_topk_softmax_kernels.cuh
Normal file
@@ -0,0 +1,867 @@
|
||||
// Adapt from
|
||||
// https://github.com/vllm-project/vllm/blob/v0.7.3/csrc/moe/topk_softmax_kernels.cu
|
||||
// which is originally adapted from
|
||||
// https://github.com/NVIDIA/TensorRT-LLM/blob/v0.7.1/cpp/tensorrt_llm/kernels/mixtureOfExperts/moe_kernels.cu
|
||||
/* Copyright 2025 SGLang Team. 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
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
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.
|
||||
==============================================================================*/
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <torch/all.h>
|
||||
|
||||
#include <cub/util_type.cuh>
|
||||
|
||||
#if !defined(USE_DCU) && !defined(USE_MACA)
|
||||
#include <cuda/functional>
|
||||
#endif
|
||||
|
||||
#include "kernels/cuda/device_utils.cuh"
|
||||
|
||||
using cub_kvp = cub::KeyValuePair<int, float>;
|
||||
|
||||
namespace {
|
||||
|
||||
using namespace xllm::kernel::cuda;
|
||||
|
||||
#if defined(USE_DCU)
|
||||
static constexpr unsigned long long kSoftmaxFullMask = 0xffffffffffffffffULL;
|
||||
#else
|
||||
static constexpr unsigned int kSoftmaxFullMask = 0xffffffffU;
|
||||
#endif
|
||||
|
||||
// ====================== Softmax things ===============================
|
||||
// We have our own implementation of softmax here so we can support transposing
|
||||
// the output in the softmax kernel when we extend this module to support
|
||||
// expert-choice routing.
|
||||
template <typename T, int TPB>
|
||||
__launch_bounds__(TPB) __global__
|
||||
void moe_softmax(const T* input,
|
||||
const bool* finished,
|
||||
float* output,
|
||||
const int num_cols,
|
||||
const float moe_softcapping,
|
||||
const float* correction_bias) {
|
||||
using BlockReduce = cub::BlockReduce<float, TPB>;
|
||||
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||
|
||||
__shared__ float normalizing_factor;
|
||||
__shared__ float float_max;
|
||||
|
||||
const int thread_row_offset = blockIdx.x * num_cols;
|
||||
|
||||
float threadData(-FLT_MAX);
|
||||
|
||||
// Don't touch finished rows.
|
||||
if ((finished != nullptr) && finished[blockIdx.x]) {
|
||||
return;
|
||||
}
|
||||
|
||||
// First pass: Apply transformation, find max, and write transformed values to
|
||||
// output
|
||||
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||
const int idx = thread_row_offset + ii;
|
||||
float val = convert_to_float<T>(input[idx]);
|
||||
|
||||
// Apply tanh softcapping if enabled
|
||||
if (moe_softcapping != 0.0f) {
|
||||
val = tanhf(val / moe_softcapping) * moe_softcapping;
|
||||
}
|
||||
|
||||
// Apply correction bias if provided
|
||||
if (correction_bias != nullptr) {
|
||||
val = val + correction_bias[ii];
|
||||
}
|
||||
|
||||
output[idx] = val; // Store transformed value
|
||||
threadData = max(val, threadData);
|
||||
}
|
||||
|
||||
const float maxElem =
|
||||
BlockReduce(tmpStorage).Reduce(threadData, MaxReduceOp());
|
||||
|
||||
if (threadIdx.x == 0) {
|
||||
float_max = maxElem;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Second pass: Compute sum using transformed values from output
|
||||
threadData = 0;
|
||||
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||
const int idx = thread_row_offset + ii;
|
||||
threadData += exp((output[idx] - float_max));
|
||||
}
|
||||
|
||||
const auto Z = BlockReduce(tmpStorage).Sum(threadData);
|
||||
|
||||
if (threadIdx.x == 0) {
|
||||
normalizing_factor = 1.f / Z;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Third pass: Compute final softmax using transformed values from output
|
||||
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||
const int idx = thread_row_offset + ii;
|
||||
const float softmax_val =
|
||||
exp((output[idx] - float_max)) * normalizing_factor;
|
||||
output[idx] = softmax_val;
|
||||
}
|
||||
}
|
||||
|
||||
namespace moe {
|
||||
class TopKPair {
|
||||
public:
|
||||
static constexpr int kPair = 2;
|
||||
static constexpr int kMaxIndex = 0;
|
||||
cub_kvp max;
|
||||
cub_kvp secondMax;
|
||||
|
||||
__device__ TopKPair() {}
|
||||
__device__ TopKPair(cub_kvp max, cub_kvp secondMax)
|
||||
: max(max), secondMax(secondMax) {}
|
||||
};
|
||||
|
||||
class TopKPairArgMax {
|
||||
public:
|
||||
__device__ TopKPairArgMax() {}
|
||||
__device__ __forceinline__ TopKPair
|
||||
operator()(const TopKPair& candidate1, const TopKPair& candidate2) const {
|
||||
cub_kvp globalMax, globalSecondMax;
|
||||
|
||||
// Determine the global maximum
|
||||
if (candidate1.max.value > candidate2.max.value) {
|
||||
globalMax = candidate1.max;
|
||||
} else {
|
||||
globalMax = candidate2.max;
|
||||
}
|
||||
|
||||
// Determine the global second maximum
|
||||
if (globalMax.key == candidate1.max.key) {
|
||||
// If candidate1 contributed the max, compare its secondMax with
|
||||
// candidate2's max
|
||||
globalSecondMax = (candidate1.secondMax.value > candidate2.max.value)
|
||||
? candidate1.secondMax
|
||||
: candidate2.max;
|
||||
} else {
|
||||
// If candidate2 contributed the max, compare its secondMax with
|
||||
// candidate1's max
|
||||
globalSecondMax = (candidate2.secondMax.value > candidate1.max.value)
|
||||
? candidate2.secondMax
|
||||
: candidate1.max;
|
||||
}
|
||||
return TopKPair(globalMax, globalSecondMax);
|
||||
}
|
||||
};
|
||||
} // namespace moe
|
||||
|
||||
template <int TPB>
|
||||
__launch_bounds__(TPB) __global__
|
||||
void moe_topk_fast(float* inputs_after_softmax,
|
||||
const bool* finished,
|
||||
float* output,
|
||||
int* indices,
|
||||
const int num_experts,
|
||||
const int k,
|
||||
const int start_expert,
|
||||
const int end_expert,
|
||||
const bool renormalize) {
|
||||
using namespace moe;
|
||||
using BlockReduce = cub::BlockReduce<TopKPair, TPB>;
|
||||
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||
TopKPair thread_pair;
|
||||
|
||||
const int block_row = blockIdx.x;
|
||||
|
||||
const bool row_is_active = finished ? !finished[block_row] : true;
|
||||
const int thread_read_offset = blockIdx.x * num_experts;
|
||||
float row_sum_for_renormalize = 0;
|
||||
// Each loop finds the top 2 elements,
|
||||
// thus requiring only ceil(k / 2) loops (calculated as (k + 1) / 2).
|
||||
for (int k_idx = 0; k_idx < (k + TopKPair::kPair - 1) / TopKPair::kPair;
|
||||
++k_idx) {
|
||||
// Initializing the top 2 elements by the minimum value.
|
||||
thread_pair.max.key = 0;
|
||||
thread_pair.max.value = -1.f;
|
||||
thread_pair.secondMax.key = 0;
|
||||
thread_pair.secondMax.value = -1.f;
|
||||
|
||||
cub_kvp inp_kvp;
|
||||
for (int expert = threadIdx.x; expert < num_experts; expert += TPB) {
|
||||
const int idx = thread_read_offset + expert;
|
||||
inp_kvp.key = expert;
|
||||
inp_kvp.value = inputs_after_softmax[idx];
|
||||
// updating the thread_pair according to inp_kvp's value
|
||||
if (inp_kvp.value > thread_pair.max.value) {
|
||||
thread_pair.secondMax = thread_pair.max;
|
||||
thread_pair.max = inp_kvp;
|
||||
} else if (inp_kvp.value > thread_pair.secondMax.value) {
|
||||
thread_pair.secondMax = inp_kvp;
|
||||
}
|
||||
}
|
||||
|
||||
TopKPairArgMax reducer;
|
||||
const TopKPair result_pair =
|
||||
BlockReduce(tmpStorage).Reduce(thread_pair, reducer);
|
||||
if (threadIdx.x == 0) {
|
||||
#pragma unroll
|
||||
// updating 2 elements to the result.
|
||||
for (int i = 0; i < TopKPair::kPair; i++) {
|
||||
if (k_idx * 2 + i >= k) {
|
||||
break;
|
||||
}
|
||||
cub_kvp result = (i == TopKPair::kMaxIndex) ? result_pair.max
|
||||
: result_pair.secondMax;
|
||||
int expert = result.key;
|
||||
bool node_uses_expert = expert >= start_expert && expert < end_expert;
|
||||
bool should_process_row = row_is_active && node_uses_expert;
|
||||
// The inputs_after_softmax is modified in-place to avoid unnecessary
|
||||
// loops for finding the top k-1 value. 1.f represents the minimum
|
||||
// value.
|
||||
inputs_after_softmax[thread_read_offset + expert] = -1.f;
|
||||
int idx = k * block_row + k_idx * 2 + i;
|
||||
output[idx] = result.value;
|
||||
indices[idx] =
|
||||
should_process_row ? (expert - start_expert) : num_experts;
|
||||
assert(indices[idx] >= 0);
|
||||
row_sum_for_renormalize += result.value;
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
if (renormalize && threadIdx.x == 0) {
|
||||
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
const int idx = k * block_row + k_idx;
|
||||
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <int TPB>
|
||||
__launch_bounds__(TPB) __global__ void moe_topK(float* inputs_after_softmax,
|
||||
const bool* finished,
|
||||
float* output,
|
||||
int* indices,
|
||||
const int num_experts,
|
||||
const int k,
|
||||
const int start_expert,
|
||||
const int end_expert,
|
||||
const bool renormalize) {
|
||||
using cub_kvp = cub::KeyValuePair<int, float>;
|
||||
using BlockReduce = cub::BlockReduce<cub_kvp, TPB>;
|
||||
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||
|
||||
cub_kvp thread_kvp;
|
||||
cub::ArgMax arg_max;
|
||||
|
||||
const int block_row = blockIdx.x;
|
||||
|
||||
const bool row_is_active = finished ? !finished[block_row] : true;
|
||||
const int thread_read_offset = blockIdx.x * num_experts;
|
||||
float row_sum_for_renormalize = 0;
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
thread_kvp.key = 0;
|
||||
thread_kvp.value = -1.f; // This is OK because inputs are probabilities
|
||||
|
||||
cub_kvp inp_kvp;
|
||||
for (int expert = threadIdx.x; expert < num_experts; expert += TPB) {
|
||||
const int idx = thread_read_offset + expert;
|
||||
inp_kvp.key = expert;
|
||||
inp_kvp.value = inputs_after_softmax[idx];
|
||||
thread_kvp = arg_max(inp_kvp, thread_kvp);
|
||||
}
|
||||
|
||||
const cub_kvp result_kvp =
|
||||
BlockReduce(tmpStorage).Reduce(thread_kvp, arg_max);
|
||||
if (threadIdx.x == 0) {
|
||||
// Ignore experts the node isn't responsible for with expert parallelism
|
||||
const int expert = result_kvp.key;
|
||||
const bool node_uses_expert =
|
||||
expert >= start_expert && expert < end_expert;
|
||||
const bool should_process_row = row_is_active && node_uses_expert;
|
||||
|
||||
const int idx = k * block_row + k_idx;
|
||||
output[idx] = result_kvp.value;
|
||||
indices[idx] = should_process_row ? (expert - start_expert) : num_experts;
|
||||
assert(indices[idx] >= 0);
|
||||
row_sum_for_renormalize += result_kvp.value;
|
||||
// The inputs_after_softmax is modified in-place to avoid unnecessary
|
||||
// loops for finding the top k-1 value. 1.f represents the minimum value.
|
||||
inputs_after_softmax[thread_read_offset + expert] = -1.f;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
if (renormalize && threadIdx.x == 0) {
|
||||
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
const int idx = k * block_row + k_idx;
|
||||
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ====================== TopK softmax things ===============================
|
||||
|
||||
/*
|
||||
A Top-K gating softmax written to exploit when the number of experts in the
|
||||
MoE layers are a small power of 2. This allows us to cleanly share the rows
|
||||
among the threads in a single warp and eliminate communication between warps
|
||||
(so no need to use shared mem).
|
||||
|
||||
It fuses the softmax, max and argmax into a single kernel.
|
||||
|
||||
Limitations:
|
||||
1) This implementation is intended for when the number of experts is a small
|
||||
power of 2. 2) This implementation assumes k is small, but will work for any
|
||||
k.
|
||||
*/
|
||||
|
||||
template <typename T,
|
||||
int VPT,
|
||||
int NUM_EXPERTS,
|
||||
int WARPS_PER_CTA,
|
||||
int BYTES_PER_LDG>
|
||||
__launch_bounds__(WARPS_PER_CTA* WARP_SIZE) __global__
|
||||
void topk_gating_softmax(const T* input,
|
||||
const bool* finished,
|
||||
float* output,
|
||||
const int num_rows,
|
||||
int* indices,
|
||||
const int k,
|
||||
const int start_expert,
|
||||
const int end_expert,
|
||||
const bool renormalize,
|
||||
const float moe_softcapping,
|
||||
const float* correction_bias) {
|
||||
// We begin by enforcing compile time assertions and setting up compile time
|
||||
// constants.
|
||||
static_assert(VPT == (VPT & -VPT), "VPT must be power of 2");
|
||||
static_assert(NUM_EXPERTS == (NUM_EXPERTS & -NUM_EXPERTS),
|
||||
"NUM_EXPERTS must be power of 2");
|
||||
static_assert(BYTES_PER_LDG == (BYTES_PER_LDG & -BYTES_PER_LDG),
|
||||
"BYTES_PER_LDG must be power of 2");
|
||||
static_assert(BYTES_PER_LDG <= 16, "BYTES_PER_LDG must be leq 16");
|
||||
|
||||
// Number of bytes each thread pulls in per load
|
||||
static constexpr int kEltsPerLdg = BYTES_PER_LDG / sizeof(T);
|
||||
static constexpr int kEltsPerRow = NUM_EXPERTS;
|
||||
static constexpr int kThreadsPerRow = kEltsPerRow / VPT;
|
||||
static constexpr int kLdgPerThread = VPT / kEltsPerLdg;
|
||||
|
||||
// Restrictions based on previous section.
|
||||
static_assert(
|
||||
VPT % kEltsPerLdg == 0,
|
||||
"The elements per thread must be a multiple of the elements per ldg");
|
||||
static_assert(WARP_SIZE % kThreadsPerRow == 0,
|
||||
"The threads per row must cleanly divide the threads per warp");
|
||||
static_assert(kThreadsPerRow == (kThreadsPerRow & -kThreadsPerRow),
|
||||
"THREADS_PER_ROW must be power of 2");
|
||||
static_assert(kThreadsPerRow <= WARP_SIZE,
|
||||
"THREADS_PER_ROW can be at most warp size");
|
||||
|
||||
// We have NUM_EXPERTS elements per row. We specialize for small #experts
|
||||
static constexpr int kEltsPerWarp = WARP_SIZE * VPT;
|
||||
static constexpr int kRowsPerWarp = kEltsPerWarp / kEltsPerRow;
|
||||
static constexpr int kRowsPerCta = WARPS_PER_CTA * kRowsPerWarp;
|
||||
|
||||
// Restrictions for previous section.
|
||||
static_assert(kEltsPerWarp % kEltsPerRow == 0,
|
||||
"The elts per row must cleanly divide the total elt per warp");
|
||||
|
||||
// ===================== From this point, we finally start computing run-time
|
||||
// variables. ========================
|
||||
|
||||
// Compute CTA and warp rows. We pack multiple rows into a single warp, and a
|
||||
// block contains WARPS_PER_CTA warps. This, each block processes a chunk of
|
||||
// rows. We start by computing the start row for each block.
|
||||
const int cta_base_row = blockIdx.x * kRowsPerCta;
|
||||
|
||||
// Now, using the base row per thread block, we compute the base row per warp.
|
||||
const int warp_base_row = cta_base_row + threadIdx.y * kRowsPerWarp;
|
||||
|
||||
// The threads in a warp are split into sub-groups that will work on a row.
|
||||
// We compute row offset for each thread sub-group
|
||||
const int thread_row_in_warp = threadIdx.x / kThreadsPerRow;
|
||||
const int thread_row = warp_base_row + thread_row_in_warp;
|
||||
|
||||
// Threads with indices out of bounds should early exit here.
|
||||
if (thread_row >= num_rows) {
|
||||
return;
|
||||
}
|
||||
const bool row_is_active = finished ? !finished[thread_row] : true;
|
||||
|
||||
// We finally start setting up the read pointers for each thread. First, each
|
||||
// thread jumps to the start of the row it will read.
|
||||
const T* thread_row_ptr = input + thread_row * kEltsPerRow;
|
||||
|
||||
// Now, we compute the group each thread belong to in order to determine the
|
||||
// first column to start loads.
|
||||
const int thread_group_idx = threadIdx.x % kThreadsPerRow;
|
||||
const int first_elt_read_by_thread = thread_group_idx * kEltsPerLdg;
|
||||
const T* thread_read_ptr = thread_row_ptr + first_elt_read_by_thread;
|
||||
|
||||
// Determine the pointer type to use to read in the data depending on the
|
||||
// BYTES_PER_LDG template param. In theory, this can support all powers of 2
|
||||
// up to 16. NOTE(woosuk): The original implementation uses CUTLASS aligned
|
||||
// array here. We defined our own aligned array and use it here to avoid the
|
||||
// dependency on CUTLASS.
|
||||
using AccessType = AlignedArray<T, kEltsPerLdg>;
|
||||
|
||||
// Finally, we pull in the data from global mem
|
||||
T row_chunk_temp[VPT];
|
||||
AccessType* row_chunk_vec_ptr =
|
||||
reinterpret_cast<AccessType*>(&row_chunk_temp);
|
||||
const AccessType* vec_thread_read_ptr =
|
||||
reinterpret_cast<const AccessType*>(thread_read_ptr);
|
||||
#pragma unroll
|
||||
// Note(Byron): interleaved loads to achieve better memory coalescing
|
||||
// | thread[0] | thread[1] | thread[2] | thread[3] | thread[0] | thread[1] |
|
||||
// thread[2] | thread[3] | ...
|
||||
for (int ii = 0; ii < kLdgPerThread; ++ii) {
|
||||
row_chunk_vec_ptr[ii] = vec_thread_read_ptr[ii * kThreadsPerRow];
|
||||
}
|
||||
|
||||
float row_chunk[VPT];
|
||||
#pragma unroll
|
||||
// Note(Byron): upcast logits to float32
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
row_chunk[ii] = convert_to_float<T>(row_chunk_temp[ii]);
|
||||
}
|
||||
|
||||
// Apply tanh softcapping and correction bias
|
||||
if (moe_softcapping != 0.0f || correction_bias != nullptr) {
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
float val = row_chunk[ii];
|
||||
|
||||
// Apply tanh softcapping if enabled
|
||||
if (moe_softcapping != 0.0f) {
|
||||
val = tanhf(val / moe_softcapping) * moe_softcapping;
|
||||
}
|
||||
|
||||
// Apply correction bias if provided
|
||||
if (correction_bias != nullptr) {
|
||||
/*
|
||||
LDG is interleaved
|
||||
|thread0 LDG| |thread1 LDG| |thread0 LDG| |thread1 LDG|
|
||||
|--------- group0 --------| |----------group1 --------|
|
||||
^ local2
|
||||
*/
|
||||
const int group_id = ii / kEltsPerLdg;
|
||||
const int local_id = ii % kEltsPerLdg;
|
||||
const int expert_idx = first_elt_read_by_thread +
|
||||
group_id * kThreadsPerRow * kEltsPerLdg +
|
||||
local_id;
|
||||
val = val + correction_bias[expert_idx];
|
||||
}
|
||||
|
||||
row_chunk[ii] = val;
|
||||
}
|
||||
}
|
||||
|
||||
// First, we perform a max reduce within the thread. We can do the max in fp16
|
||||
// safely (I think) and just convert to float afterwards for the exp + sum
|
||||
// reduction.
|
||||
float thread_max = row_chunk[0];
|
||||
#pragma unroll
|
||||
for (int ii = 1; ii < VPT; ++ii) {
|
||||
thread_max = max(thread_max, row_chunk[ii]);
|
||||
}
|
||||
|
||||
/*********************************/
|
||||
/********* Softmax Begin *********/
|
||||
/*********************************/
|
||||
|
||||
// Now, we find the max within the thread group and distribute among the
|
||||
// threads. We use a butterfly reduce. lane id: 0-31 within a warp
|
||||
#pragma unroll
|
||||
for (int mask = kThreadsPerRow / 2; mask > 0; mask /= 2) {
|
||||
// butterfly reduce with (lane id ^ mask)
|
||||
thread_max = max(thread_max,
|
||||
XLLM_SHFL_XOR_SYNC_WIDTH(
|
||||
kSoftmaxFullMask, thread_max, mask, kThreadsPerRow));
|
||||
}
|
||||
|
||||
// From this point, thread max in all the threads have the max within the row.
|
||||
// Now, we subtract the max from each element in the thread and take the exp.
|
||||
// We also compute the thread local sum.
|
||||
float row_sum = 0;
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
row_chunk[ii] = expf(row_chunk[ii] - thread_max);
|
||||
row_sum += row_chunk[ii];
|
||||
}
|
||||
|
||||
// Now, we perform the sum reduce within each thread group. Similar to the max
|
||||
// reduce, we use a bufferfly pattern.
|
||||
#pragma unroll
|
||||
for (int mask = kThreadsPerRow / 2; mask > 0; mask /= 2) {
|
||||
row_sum += XLLM_SHFL_XOR_SYNC_WIDTH(
|
||||
kSoftmaxFullMask, row_sum, mask, kThreadsPerRow);
|
||||
}
|
||||
|
||||
// From this point, all threads have the max and the sum for their rows in the
|
||||
// thread_max and thread_sum variables respectively. Finally, we can scale the
|
||||
// rows for the softmax. Technically, for top-k gating we don't need to
|
||||
// compute the entire softmax row. We can likely look at the maxes and only
|
||||
// compute for the top-k values in the row. However, this kernel will likely
|
||||
// not be a bottle neck and it seems better to closer match torch and find the
|
||||
// argmax after computing the softmax.
|
||||
const float reciprocal_row_sum = 1.f / row_sum;
|
||||
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
row_chunk[ii] = row_chunk[ii] * reciprocal_row_sum;
|
||||
}
|
||||
/*******************************/
|
||||
/********* Softmax End *********/
|
||||
/*******************************/
|
||||
|
||||
// Now, softmax_res contains the softmax of the row chunk. Now, I want to find
|
||||
// the topk elements in each row, along with the max index.
|
||||
int start_col = first_elt_read_by_thread;
|
||||
static constexpr int kColsPerGroupLdg = kEltsPerLdg * kThreadsPerRow;
|
||||
|
||||
float row_sum_for_renormalize = 0;
|
||||
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
// First, each thread does the local argmax
|
||||
float max_val = row_chunk[0];
|
||||
int expert = start_col;
|
||||
#pragma unroll
|
||||
for (int ldg = 0, col = start_col; ldg < kLdgPerThread;
|
||||
++ldg, col += kColsPerGroupLdg) {
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < kEltsPerLdg; ++ii) {
|
||||
float val = row_chunk[ldg * kEltsPerLdg + ii];
|
||||
|
||||
// No check on the experts here since columns with the smallest index
|
||||
// are processed first and only updated if > (not >=)
|
||||
if (val > max_val) {
|
||||
max_val = val;
|
||||
expert = col + ii;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Now, we perform the argmax reduce. We use the butterfly pattern so threads
|
||||
// reach consensus about the max. This will be useful for K > 1 so that the
|
||||
// threads can agree on "who" had the max value. That thread can then blank out
|
||||
// their max with -inf and the warp can run more iterations...
|
||||
#pragma unroll
|
||||
for (int mask = kThreadsPerRow / 2; mask > 0; mask /= 2) {
|
||||
float other_max = XLLM_SHFL_XOR_SYNC_WIDTH(
|
||||
kSoftmaxFullMask, max_val, mask, kThreadsPerRow);
|
||||
int other_expert = XLLM_SHFL_XOR_SYNC_WIDTH(
|
||||
kSoftmaxFullMask, expert, mask, kThreadsPerRow);
|
||||
|
||||
// We want lower indices to "win" in every thread so we break ties this
|
||||
// way
|
||||
if (other_max > max_val ||
|
||||
(other_max == max_val && other_expert < expert)) {
|
||||
max_val = other_max;
|
||||
expert = other_expert;
|
||||
}
|
||||
}
|
||||
|
||||
// Write the max for this k iteration to global memory.
|
||||
if (thread_group_idx == 0) {
|
||||
// Add a guard to ignore experts not included by this node
|
||||
const bool node_uses_expert =
|
||||
expert >= start_expert && expert < end_expert;
|
||||
const bool should_process_row = row_is_active && node_uses_expert;
|
||||
|
||||
// The lead thread from each sub-group will write out the final results to
|
||||
// global memory. (This will be a single) thread per row of the
|
||||
// input/output matrices.
|
||||
const int idx = k * thread_row + k_idx;
|
||||
output[idx] = max_val;
|
||||
indices[idx] = should_process_row ? (expert - start_expert) : NUM_EXPERTS;
|
||||
row_sum_for_renormalize += max_val;
|
||||
}
|
||||
|
||||
// Finally, we clear the value in the thread with the current max if there
|
||||
// is another iteration to run.
|
||||
if (k_idx + 1 < k) {
|
||||
const int ldg_group_for_expert = expert / kColsPerGroupLdg;
|
||||
const int thread_to_clear_in_group =
|
||||
(expert / kEltsPerLdg) % kThreadsPerRow;
|
||||
|
||||
// Only the thread in the group which produced the max will reset the
|
||||
// "winning" value to -inf.
|
||||
if (thread_group_idx == thread_to_clear_in_group) {
|
||||
const int offset_for_expert = expert % kEltsPerLdg;
|
||||
// Safe to set to any negative value since row_chunk values must be
|
||||
// between 0 and 1.
|
||||
row_chunk[ldg_group_for_expert * kEltsPerLdg + offset_for_expert] =
|
||||
-10000.f;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fuse renormalization of topk_weights into this kernel
|
||||
if (renormalize && thread_group_idx == 0) {
|
||||
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||
#pragma unroll
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
const int idx = k * thread_row + k_idx;
|
||||
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, int EXPERTS, int WARPS_PER_TB>
|
||||
void topk_gating_softmax_launcher_helper(const T* input,
|
||||
const bool* finished,
|
||||
float* output,
|
||||
int* indices,
|
||||
const int num_rows,
|
||||
const int k,
|
||||
const int start_expert,
|
||||
const int end_expert,
|
||||
const bool renormalize,
|
||||
const float moe_softcapping,
|
||||
const float* correction_bias,
|
||||
cudaStream_t stream) {
|
||||
static constexpr std::size_t kMaxBytesPerLdg = 16;
|
||||
|
||||
static constexpr int kBytesPerLdg = MIN(kMaxBytesPerLdg, sizeof(T) * EXPERTS);
|
||||
using Constants = TopkConstants<T, EXPERTS, kBytesPerLdg>;
|
||||
static constexpr int kVpt = Constants::VPT;
|
||||
static constexpr int kRowsPerWarp = Constants::ROWS_PER_WARP;
|
||||
const int num_warps = (num_rows + kRowsPerWarp - 1) / kRowsPerWarp;
|
||||
const int num_blocks = (num_warps + WARPS_PER_TB - 1) / WARPS_PER_TB;
|
||||
|
||||
dim3 block_dim(WARP_SIZE, WARPS_PER_TB);
|
||||
topk_gating_softmax<T, kVpt, EXPERTS, WARPS_PER_TB, kBytesPerLdg>
|
||||
<<<num_blocks, block_dim, 0, stream>>>(input,
|
||||
finished,
|
||||
output,
|
||||
num_rows,
|
||||
indices,
|
||||
k,
|
||||
start_expert,
|
||||
end_expert,
|
||||
renormalize,
|
||||
moe_softcapping,
|
||||
correction_bias);
|
||||
}
|
||||
|
||||
#define LAUNCH_SOFTMAX(TYPE, NUM_EXPERTS, WARPS_PER_TB) \
|
||||
topk_gating_softmax_launcher_helper<TYPE, NUM_EXPERTS, WARPS_PER_TB>( \
|
||||
gating_output, \
|
||||
nullptr, \
|
||||
topk_weights, \
|
||||
topk_indices, \
|
||||
num_tokens, \
|
||||
topk, \
|
||||
0, \
|
||||
num_experts, \
|
||||
renormalize, \
|
||||
moe_softcapping, \
|
||||
correction_bias, \
|
||||
stream);
|
||||
|
||||
template <typename T>
|
||||
void topk_gating_softmax_kernel_launcher(const T* gating_output,
|
||||
float* topk_weights,
|
||||
int* topk_indices,
|
||||
float* softmax_workspace,
|
||||
const int num_tokens,
|
||||
const int num_experts,
|
||||
const int topk,
|
||||
const bool renormalize,
|
||||
const float moe_softcapping,
|
||||
const float* correction_bias,
|
||||
cudaStream_t stream) {
|
||||
static constexpr int kWarpsPerTb = 4;
|
||||
switch (num_experts) {
|
||||
case 1:
|
||||
LAUNCH_SOFTMAX(T, 1, kWarpsPerTb);
|
||||
break;
|
||||
case 2:
|
||||
LAUNCH_SOFTMAX(T, 2, kWarpsPerTb);
|
||||
break;
|
||||
case 4:
|
||||
LAUNCH_SOFTMAX(T, 4, kWarpsPerTb);
|
||||
break;
|
||||
case 8:
|
||||
LAUNCH_SOFTMAX(T, 8, kWarpsPerTb);
|
||||
break;
|
||||
case 16:
|
||||
LAUNCH_SOFTMAX(T, 16, kWarpsPerTb);
|
||||
break;
|
||||
case 32:
|
||||
LAUNCH_SOFTMAX(T, 32, kWarpsPerTb);
|
||||
break;
|
||||
case 64:
|
||||
LAUNCH_SOFTMAX(T, 64, kWarpsPerTb);
|
||||
break;
|
||||
case 128:
|
||||
LAUNCH_SOFTMAX(T, 128, kWarpsPerTb);
|
||||
break;
|
||||
case 256:
|
||||
LAUNCH_SOFTMAX(T, 256, kWarpsPerTb);
|
||||
break;
|
||||
default: {
|
||||
CHECK(softmax_workspace != nullptr)
|
||||
<< "softmax_workspace must be provided for num_experts that are "
|
||||
"not a power of 2.";
|
||||
static constexpr int kTpb = 256;
|
||||
moe_softmax<T, kTpb><<<num_tokens, kTpb, 0, stream>>>(gating_output,
|
||||
nullptr,
|
||||
softmax_workspace,
|
||||
num_experts,
|
||||
moe_softcapping,
|
||||
correction_bias);
|
||||
if (topk == 1) {
|
||||
// Note: As an optimization for better performance,
|
||||
// the softmax_workspace is overwritten in-place by both moeTopK and
|
||||
// moe_topk_fast.
|
||||
moe_topK<kTpb><<<num_tokens, kTpb, 0, stream>>>(softmax_workspace,
|
||||
nullptr,
|
||||
topk_weights,
|
||||
topk_indices,
|
||||
num_experts,
|
||||
topk,
|
||||
0,
|
||||
num_experts,
|
||||
renormalize);
|
||||
} else {
|
||||
moe_topk_fast<kTpb><<<num_tokens, kTpb, 0, stream>>>(softmax_workspace,
|
||||
nullptr,
|
||||
topk_weights,
|
||||
topk_indices,
|
||||
num_experts,
|
||||
topk,
|
||||
0,
|
||||
num_experts,
|
||||
renormalize);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} // namespace
|
||||
|
||||
namespace xllm::kernel::cuda {
|
||||
void topk_softmax(torch::Tensor& topk_weights, // [num_tokens, topk]
|
||||
torch::Tensor& topk_indices, // [num_tokens, topk]
|
||||
torch::Tensor& gating_output, // [num_tokens, num_experts]
|
||||
const bool renormalize,
|
||||
const double moe_softcapping,
|
||||
const std::optional<torch::Tensor>& correction_bias) {
|
||||
// Check data type
|
||||
CHECK(gating_output.scalar_type() == at::ScalarType::Float ||
|
||||
gating_output.scalar_type() == at::ScalarType::Half ||
|
||||
gating_output.scalar_type() == at::ScalarType::BFloat16)
|
||||
<< "gating_output must be float32, float16, or bfloat16";
|
||||
|
||||
// Check dimensions
|
||||
CHECK(gating_output.dim() == 2)
|
||||
<< "gating_output must be 2D tensor [num_tokens, num_experts]";
|
||||
CHECK(topk_weights.dim() == 2)
|
||||
<< "topk_weights must be 2D tensor [num_tokens, topk]";
|
||||
CHECK(topk_indices.dim() == 2)
|
||||
<< "topk_indices must be 2D tensor [num_tokens, topk]";
|
||||
|
||||
// Check shapes
|
||||
CHECK(gating_output.size(0) == topk_weights.size(0))
|
||||
<< "First dimension of topk_weights must match num_tokens in "
|
||||
"gating_output"
|
||||
<< "First dimension of topk_indices must match num_tokens in "
|
||||
"gating_output";
|
||||
|
||||
CHECK(topk_weights.size(-1) == topk_indices.size(-1))
|
||||
<< "Second dimension of topk_indices must match topk in topk_weights"
|
||||
<< "topk must be less than or equal to num_experts";
|
||||
|
||||
const int num_experts = static_cast<int>(gating_output.size(-1));
|
||||
const int num_tokens = static_cast<int>(gating_output.size(0));
|
||||
const int topk = static_cast<int>(topk_weights.size(-1));
|
||||
|
||||
const bool is_pow_2 =
|
||||
(num_experts != 0) && ((num_experts & (num_experts - 1)) == 0);
|
||||
const bool needs_workspace = !is_pow_2 || num_experts > 256;
|
||||
const int64_t workspace_size = needs_workspace ? num_tokens * num_experts : 0;
|
||||
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(gating_output));
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
torch::Tensor softmax_workspace = torch::empty(
|
||||
{workspace_size}, gating_output.options().dtype(at::ScalarType::Float));
|
||||
|
||||
const at::ScalarType dtype = gating_output.scalar_type();
|
||||
|
||||
// Validate correction_bias if provided - must always be float32
|
||||
const float* bias_ptr = nullptr;
|
||||
if (correction_bias.has_value()) {
|
||||
const torch::Tensor& bias_tensor = correction_bias.value();
|
||||
CHECK(bias_tensor.dim() == 1)
|
||||
<< "correction_bias must be 1D tensor [num_experts]";
|
||||
CHECK(bias_tensor.size(0) == num_experts)
|
||||
<< "correction_bias size must match num_experts";
|
||||
CHECK(bias_tensor.scalar_type() == at::ScalarType::Float)
|
||||
<< "correction_bias must be float32, got " << bias_tensor.scalar_type();
|
||||
bias_ptr = bias_tensor.data_ptr<float>();
|
||||
}
|
||||
|
||||
// Cast moe_softcapping from double to float for CUDA kernels
|
||||
const float moe_softcapping_f = static_cast<float>(moe_softcapping);
|
||||
|
||||
if (dtype == at::ScalarType::Float) {
|
||||
topk_gating_softmax_kernel_launcher<float>(
|
||||
gating_output.data_ptr<float>(),
|
||||
topk_weights.data_ptr<float>(),
|
||||
topk_indices.data_ptr<int>(),
|
||||
softmax_workspace.data_ptr<float>(),
|
||||
num_tokens,
|
||||
num_experts,
|
||||
topk,
|
||||
renormalize,
|
||||
moe_softcapping_f,
|
||||
bias_ptr,
|
||||
stream);
|
||||
} else if (dtype == at::ScalarType::Half) {
|
||||
topk_gating_softmax_kernel_launcher<__half>(
|
||||
reinterpret_cast<const __half*>(gating_output.data_ptr<at::Half>()),
|
||||
topk_weights.data_ptr<float>(),
|
||||
topk_indices.data_ptr<int>(),
|
||||
softmax_workspace.data_ptr<float>(),
|
||||
num_tokens,
|
||||
num_experts,
|
||||
topk,
|
||||
renormalize,
|
||||
moe_softcapping_f,
|
||||
bias_ptr,
|
||||
stream);
|
||||
} else if (dtype == at::ScalarType::BFloat16) {
|
||||
topk_gating_softmax_kernel_launcher<BFloat16Type>(
|
||||
reinterpret_cast<const BFloat16Type*>(
|
||||
gating_output.data_ptr<at::BFloat16>()),
|
||||
topk_weights.data_ptr<float>(),
|
||||
topk_indices.data_ptr<int>(),
|
||||
softmax_workspace.data_ptr<float>(),
|
||||
num_tokens,
|
||||
num_experts,
|
||||
topk,
|
||||
renormalize,
|
||||
moe_softcapping_f,
|
||||
bias_ptr,
|
||||
stream);
|
||||
} else {
|
||||
LOG(FATAL) << "Unsupported gating_output dtype: " << dtype;
|
||||
}
|
||||
}
|
||||
} // namespace xllm::kernel::cuda
|
||||
59
ex_engine/xllm_kernels/npu/npu_causal_conv1d.cpp
Normal file
59
ex_engine/xllm_kernels/npu/npu_causal_conv1d.cpp
Normal file
@@ -0,0 +1,59 @@
|
||||
/* 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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "core/kernels/npu/aclnn/pytorch_npu_helper.hpp"
|
||||
#include "core/kernels/npu/utils.h"
|
||||
#include "core/kernels/npu/xllm_ops/xllm_ops_api.h"
|
||||
|
||||
namespace xllm::kernel::npu {
|
||||
|
||||
torch::Tensor causal_conv1d(const torch::Tensor& x,
|
||||
const torch::Tensor& weight,
|
||||
const torch::Tensor& conv_state,
|
||||
const std::optional<torch::Tensor>& bias_opt,
|
||||
const torch::IntArrayRef query_start_loc_opt,
|
||||
const torch::IntArrayRef cache_indices_opt,
|
||||
const torch::IntArrayRef initial_state_mode_opt,
|
||||
const torch::IntArrayRef num_accepted_tokens_opt,
|
||||
int64_t activation_mode,
|
||||
int64_t pad_slot_id,
|
||||
int64_t run_mode) {
|
||||
check_tensor(x, "x", "causal_conv1d");
|
||||
check_tensor(weight, "weight", "causal_conv1d");
|
||||
check_tensor(conv_state, "conv_state", "causal_conv1d");
|
||||
|
||||
c10::optional<torch::Tensor> bias_tensor = c10::nullopt;
|
||||
if (bias_opt.has_value() && bias_opt.value().defined()) {
|
||||
bias_tensor = bias_opt.value();
|
||||
}
|
||||
|
||||
torch::Tensor output = torch::empty(x.sizes(), x.options());
|
||||
EXEC_NPU_CMD(aclnnCausalConv1d,
|
||||
x,
|
||||
weight,
|
||||
bias_tensor,
|
||||
conv_state,
|
||||
query_start_loc_opt,
|
||||
cache_indices_opt,
|
||||
initial_state_mode_opt,
|
||||
num_accepted_tokens_opt,
|
||||
activation_mode,
|
||||
pad_slot_id,
|
||||
run_mode,
|
||||
output);
|
||||
return output;
|
||||
}
|
||||
|
||||
} // namespace xllm::kernel::npu
|
||||
@@ -0,0 +1,83 @@
|
||||
/* 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.
|
||||
==============================================================================*/
|
||||
|
||||
#include <glog/logging.h>
|
||||
|
||||
#include "core/kernels/npu/aclnn/pytorch_npu_helper.hpp"
|
||||
#include "core/kernels/npu/npu_ops_api.h"
|
||||
#include "core/kernels/npu/utils.h"
|
||||
|
||||
namespace {
|
||||
|
||||
c10::optional<torch::Tensor> to_c10_optional_tensor(
|
||||
const std::optional<torch::Tensor>& tensor_opt) {
|
||||
if (tensor_opt.has_value() && tensor_opt.value().defined()) {
|
||||
return tensor_opt.value();
|
||||
}
|
||||
return c10::nullopt;
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
namespace xllm::kernel::npu {
|
||||
|
||||
torch::Tensor npu_recurrent_gated_delta_rule(
|
||||
const torch::Tensor& query,
|
||||
const torch::Tensor& key,
|
||||
const torch::Tensor& value,
|
||||
torch::Tensor& state,
|
||||
const std::optional<torch::Tensor>& beta,
|
||||
const std::optional<double> scale,
|
||||
const std::optional<torch::Tensor>& actual_seq_lengths,
|
||||
const std::optional<torch::Tensor>& ssm_state_indices,
|
||||
const std::optional<torch::Tensor>& num_accepted_tokens,
|
||||
const std::optional<torch::Tensor>& g,
|
||||
const std::optional<torch::Tensor>& gk) {
|
||||
check_tensor(query, "query", "recurrent_gated_delta_rule");
|
||||
check_tensor(key, "key", "recurrent_gated_delta_rule");
|
||||
check_tensor(value, "value", "recurrent_gated_delta_rule");
|
||||
check_tensor(state, "state", "recurrent_gated_delta_rule");
|
||||
CHECK(scale.has_value())
|
||||
<< "recurrent_gated_delta_rule requires a valid scale value";
|
||||
|
||||
c10::optional<torch::Tensor> beta_tensor = to_c10_optional_tensor(beta);
|
||||
c10::optional<torch::Tensor> actual_seq_lengths_tensor =
|
||||
to_c10_optional_tensor(actual_seq_lengths);
|
||||
c10::optional<torch::Tensor> ssm_state_indices_tensor =
|
||||
to_c10_optional_tensor(ssm_state_indices);
|
||||
c10::optional<torch::Tensor> num_accepted_tokens_tensor =
|
||||
to_c10_optional_tensor(num_accepted_tokens);
|
||||
c10::optional<torch::Tensor> g_tensor = to_c10_optional_tensor(g);
|
||||
c10::optional<torch::Tensor> gk_tensor = to_c10_optional_tensor(gk);
|
||||
float scale_value = static_cast<float>(scale.value());
|
||||
torch::Tensor output = torch::empty_like(value);
|
||||
|
||||
EXEC_NPU_CMD(aclnnRecurrentGatedDeltaRule,
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
beta_tensor,
|
||||
state,
|
||||
actual_seq_lengths_tensor,
|
||||
ssm_state_indices_tensor,
|
||||
g_tensor,
|
||||
gk_tensor,
|
||||
num_accepted_tokens_tensor,
|
||||
scale_value,
|
||||
output);
|
||||
return output;
|
||||
}
|
||||
|
||||
} // namespace xllm::kernel::npu
|
||||
236
ex_engine/xllm_layers/mlu/qwen3_5_attention.cpp
Normal file
236
ex_engine/xllm_layers/mlu/qwen3_5_attention.cpp
Normal file
@@ -0,0 +1,236 @@
|
||||
/* 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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "qwen3_5_attention.h"
|
||||
|
||||
#include <glog/logging.h>
|
||||
|
||||
#include <tuple>
|
||||
|
||||
#include "kernels/ops_api.h"
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
Qwen3_5AttentionImpl::Qwen3_5AttentionImpl(const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options,
|
||||
int32_t layer_id) {
|
||||
const int64_t tp_size = parallel_args.tp_group_->world_size();
|
||||
const int64_t total_num_heads = args.n_heads();
|
||||
const int64_t total_num_kv_heads = args.n_kv_heads().value_or(args.n_heads());
|
||||
layer_id_ = layer_id;
|
||||
rank_ = parallel_args.tp_group_->rank();
|
||||
CHECK(total_num_heads % tp_size == 0);
|
||||
num_heads_ = total_num_heads / tp_size;
|
||||
|
||||
if (total_num_kv_heads >= tp_size) {
|
||||
CHECK(total_num_kv_heads % tp_size == 0);
|
||||
num_kv_heads_ = total_num_kv_heads / tp_size;
|
||||
num_kv_head_replicas_ = 1;
|
||||
} else {
|
||||
CHECK(tp_size % total_num_kv_heads == 0);
|
||||
num_kv_heads_ = 1;
|
||||
num_kv_head_replicas_ = tp_size / total_num_kv_heads;
|
||||
}
|
||||
|
||||
head_dim_ = args.head_dim();
|
||||
q_size_ = num_heads_ * head_dim_;
|
||||
kv_size_ = num_kv_heads_ * head_dim_;
|
||||
scaling_ = 1.0f / std::sqrt(static_cast<float>(head_dim_));
|
||||
attn_output_gate_ = args.attn_output_gate();
|
||||
mrope_cu_seq_lens_ = torch::zeros(2, torch::kInt32).to(options.device());
|
||||
// 1. QKV linear
|
||||
qkv_proj_ = register_module(
|
||||
"qkv_proj",
|
||||
QKVParallelLinear(args.hidden_size(),
|
||||
attn_output_gate_ ? num_heads_ * 2 : num_heads_,
|
||||
num_kv_heads_,
|
||||
args.head_dim(),
|
||||
num_kv_head_replicas_,
|
||||
/*bias=*/args.attention_bias(),
|
||||
/*gather_output=*/false,
|
||||
parallel_args,
|
||||
options));
|
||||
|
||||
// 2. O proj
|
||||
o_proj_ = register_module("o_proj",
|
||||
RowParallelLinear(total_num_heads * head_dim_,
|
||||
args.hidden_size(),
|
||||
/*bias=*/false,
|
||||
/*input_is_parallelized=*/true,
|
||||
/*if_reduce_results=*/true,
|
||||
quant_args,
|
||||
parallel_args.tp_group_,
|
||||
options));
|
||||
|
||||
// 3. Q norm
|
||||
q_norm_ = register_module(
|
||||
"q_norm", Qwen3NextRMSNorm(head_dim_, args.rms_norm_eps(), options));
|
||||
|
||||
// 4. K norm
|
||||
k_norm_ = register_module(
|
||||
"k_norm", Qwen3NextRMSNorm(head_dim_, args.rms_norm_eps(), options));
|
||||
|
||||
// 5. Attention
|
||||
attn_ = register_module("attn",
|
||||
Attention(num_heads_,
|
||||
head_dim_,
|
||||
scaling_,
|
||||
num_kv_heads_,
|
||||
args.sliding_window()));
|
||||
|
||||
// 6. Rotary embedding
|
||||
const int32_t rotary_dim =
|
||||
static_cast<int32_t>(head_dim_ * args.partial_rotary_factor());
|
||||
rotary_emb_ =
|
||||
register_module("rope",
|
||||
MRotaryEmbedding(rotary_dim,
|
||||
args.max_position_embeddings(),
|
||||
args.rope_theta(),
|
||||
/*interleaved=*/false,
|
||||
args.rope_scaling_mrope_section(),
|
||||
options));
|
||||
}
|
||||
|
||||
void Qwen3_5AttentionImpl::rotary_emb_forward(
|
||||
torch::Tensor& q,
|
||||
torch::Tensor& k,
|
||||
const torch::Tensor& positions,
|
||||
const AttentionMetadata& attn_metadata) {
|
||||
auto q_shape = q.sizes();
|
||||
auto k_shape = k.sizes();
|
||||
auto num_tokens = positions.size(-1);
|
||||
mrope_cu_seq_lens_[1] = num_tokens;
|
||||
|
||||
xllm::kernel::RotaryParams rotary_params;
|
||||
bool only_prefill =
|
||||
(attn_metadata.is_prefill || attn_metadata.is_chunked_prefill);
|
||||
if (only_prefill) {
|
||||
rotary_params.sin = attn_metadata.mrope_sin;
|
||||
rotary_params.cos = attn_metadata.mrope_cos;
|
||||
rotary_params.position_ids = std::nullopt;
|
||||
rotary_params.cu_query_lens = mrope_cu_seq_lens_;
|
||||
rotary_params.interleaved = false;
|
||||
rotary_params.discrete = false;
|
||||
rotary_params.max_query_len = num_tokens;
|
||||
|
||||
rotary_params.q = q.view({num_tokens, -1, head_dim_});
|
||||
xllm::kernel::apply_rotary(rotary_params);
|
||||
q = rotary_params.q.reshape(q_shape);
|
||||
|
||||
rotary_params.q = k.view({num_tokens, -1, head_dim_});
|
||||
xllm::kernel::apply_rotary(rotary_params);
|
||||
k = rotary_params.q.reshape(k_shape);
|
||||
} else {
|
||||
if (positions.dim() == 2) {
|
||||
rotary_params.position_ids = positions[0];
|
||||
} else {
|
||||
rotary_params.position_ids = positions;
|
||||
}
|
||||
rotary_params.sin = rotary_emb_->get_sin_cache();
|
||||
rotary_params.cos = rotary_emb_->get_cos_cache();
|
||||
|
||||
rotary_params.interleaved = false;
|
||||
rotary_params.discrete = true;
|
||||
rotary_params.max_query_len = num_tokens;
|
||||
rotary_params.q = q.view({1, num_tokens, -1, head_dim_});
|
||||
xllm::kernel::apply_rotary(rotary_params);
|
||||
q = rotary_params.q.reshape(q_shape);
|
||||
|
||||
rotary_params.q = k.view({1, num_tokens, -1, head_dim_});
|
||||
xllm::kernel::apply_rotary(rotary_params);
|
||||
k = rotary_params.q.reshape(k_shape);
|
||||
}
|
||||
}
|
||||
|
||||
torch::Tensor Qwen3_5AttentionImpl::forward(
|
||||
const torch::Tensor& positions,
|
||||
const torch::Tensor& hidden_states,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache) {
|
||||
// 1. qkv projection
|
||||
auto qkv = qkv_proj_->forward(hidden_states);
|
||||
torch::Tensor q, k, v;
|
||||
torch::Tensor gate;
|
||||
|
||||
if (attn_output_gate_) {
|
||||
// Split qkv for attn_output_gate case: [q_size*2, kv_size, kv_size]
|
||||
auto q_gate = qkv.slice(/*dim=*/-1, 0, q_size_ * 2);
|
||||
k = qkv.slice(/*dim=*/-1, q_size_ * 2, q_size_ * 2 + kv_size_);
|
||||
v = qkv.slice(
|
||||
/*dim=*/-1, q_size_ * 2 + kv_size_, q_size_ * 2 + kv_size_ * 2);
|
||||
v = v.contiguous();
|
||||
|
||||
std::vector<int64_t> orig_shape;
|
||||
for (int64_t i = 0; i < q_gate.dim() - 1; i++) {
|
||||
orig_shape.push_back(q_gate.size(i));
|
||||
}
|
||||
std::vector<int64_t> new_shape = orig_shape;
|
||||
new_shape.push_back(num_heads_);
|
||||
new_shape.push_back(-1);
|
||||
torch::Tensor q_gate_reshaped = q_gate.reshape(new_shape);
|
||||
auto chunks = torch::chunk(q_gate_reshaped, 2, /*dim=*/-1);
|
||||
q = chunks[0];
|
||||
gate = chunks[1];
|
||||
|
||||
std::vector<int64_t> q_new_shape = orig_shape;
|
||||
q_new_shape.push_back(-1);
|
||||
q = q.reshape(q_new_shape);
|
||||
|
||||
std::vector<int64_t> gate_new_shape = orig_shape;
|
||||
gate_new_shape.push_back(-1);
|
||||
gate = gate.reshape(gate_new_shape);
|
||||
} else {
|
||||
// Normal case: [q_size, kv_size, kv_size]
|
||||
q = qkv.slice(/*dim=*/-1, 0, q_size_);
|
||||
k = qkv.slice(/*dim=*/-1, q_size_, q_size_ + kv_size_);
|
||||
v = qkv.slice(/*dim=*/-1, q_size_ + kv_size_, q_size_ + 2 * kv_size_);
|
||||
}
|
||||
|
||||
const int64_t T = q.size(0);
|
||||
|
||||
auto q_reshaped = q.reshape({T, num_heads_, head_dim_});
|
||||
auto q_normed = std::get<0>(q_norm_->forward(q_reshaped));
|
||||
auto k_reshaped = k.reshape({T, num_kv_heads_, head_dim_});
|
||||
auto k_normed = std::get<0>(k_norm_->forward(k_reshaped));
|
||||
|
||||
q = q_normed.view({T, q_size_});
|
||||
k = k_normed.view({T, kv_size_});
|
||||
rotary_emb_forward(q, k, positions, attn_metadata);
|
||||
auto out = std::get<0>(attn_->forward(attn_metadata, q, k, v, kv_cache));
|
||||
|
||||
if (attn_output_gate_) {
|
||||
gate = torch::sigmoid(gate);
|
||||
out = out * gate;
|
||||
}
|
||||
|
||||
out = o_proj_->forward(out);
|
||||
return out;
|
||||
}
|
||||
|
||||
void Qwen3_5AttentionImpl::load_state_dict(const StateDict& state_dict) {
|
||||
qkv_proj_->load_state_dict(state_dict, {"q_proj.", "k_proj.", "v_proj."});
|
||||
o_proj_->load_state_dict(state_dict.get_dict_with_prefix("o_proj."));
|
||||
if (auto w = state_dict.get_tensor("q_norm.weight"); w.defined()) {
|
||||
q_norm_->load_state_dict(StateDict({{"weight", w}}));
|
||||
}
|
||||
if (auto w = state_dict.get_tensor("k_norm.weight"); w.defined()) {
|
||||
k_norm_->load_state_dict(StateDict({{"weight", w}}));
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
79
ex_engine/xllm_layers/mlu/qwen3_5_attention.h
Normal file
79
ex_engine/xllm_layers/mlu/qwen3_5_attention.h
Normal file
@@ -0,0 +1,79 @@
|
||||
/* 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 "attention.h"
|
||||
#include "framework/kv_cache/kv_cache.h"
|
||||
#include "framework/model/model_args.h"
|
||||
#include "framework/parallel_state/parallel_args.h"
|
||||
#include "framework/quant_args.h"
|
||||
#include "framework/state_dict/state_dict.h"
|
||||
#include "layers/common/linear.h"
|
||||
#include "layers/common/partial_rotary_embedding.h"
|
||||
#include "layers/common/qwen3_next_rms_norm.h"
|
||||
#include "layers/common/rotary_embedding.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3_5AttentionImpl : public torch::nn::Module {
|
||||
public:
|
||||
Qwen3_5AttentionImpl() = default;
|
||||
Qwen3_5AttentionImpl(const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options,
|
||||
int32_t layer_id);
|
||||
|
||||
torch::Tensor forward(const torch::Tensor& positions,
|
||||
const torch::Tensor& hidden_states,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache);
|
||||
|
||||
void load_state_dict(const StateDict& state_dict);
|
||||
void rotary_emb_forward(torch::Tensor& q,
|
||||
torch::Tensor& k,
|
||||
const torch::Tensor& positions,
|
||||
const AttentionMetadata& attn_metadata);
|
||||
|
||||
private:
|
||||
int64_t num_heads_;
|
||||
int64_t num_kv_heads_;
|
||||
int64_t num_kv_head_replicas_;
|
||||
int64_t head_dim_;
|
||||
int64_t q_size_;
|
||||
int64_t kv_size_;
|
||||
float scaling_;
|
||||
bool attn_output_gate_;
|
||||
int32_t layer_id_;
|
||||
int32_t rank_;
|
||||
|
||||
QKVParallelLinear qkv_proj_{nullptr};
|
||||
RowParallelLinear o_proj_{nullptr};
|
||||
|
||||
Qwen3NextRMSNorm q_norm_{nullptr};
|
||||
Qwen3NextRMSNorm k_norm_{nullptr};
|
||||
|
||||
Attention attn_{nullptr};
|
||||
MRotaryEmbedding rotary_emb_{nullptr};
|
||||
torch::Tensor mrope_cu_seq_lens_;
|
||||
};
|
||||
TORCH_MODULE(Qwen3_5Attention);
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
193
ex_engine/xllm_layers/mlu/qwen3_5_decoder_layer.cpp
Normal file
193
ex_engine/xllm_layers/mlu/qwen3_5_decoder_layer.cpp
Normal file
@@ -0,0 +1,193 @@
|
||||
/* 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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "qwen3_5_decoder_layer.h"
|
||||
|
||||
#include <glog/logging.h>
|
||||
|
||||
#include "common/global_flags.h"
|
||||
#include "layers/common/dp_utils.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
namespace {
|
||||
bool use_moe_all2all(bool enable_deep_ep,
|
||||
const ModelInputParams& input_params) {
|
||||
return enable_deep_ep && all_dp_ranks_are_decode(input_params);
|
||||
}
|
||||
|
||||
bool is_moe_layer(const ModelArgs& model_args, int32_t layer_id) {
|
||||
const auto& mlp_only_layers = model_args.mlp_only_layers();
|
||||
return std::count(mlp_only_layers.begin(), mlp_only_layers.end(), layer_id) ==
|
||||
0 &&
|
||||
model_args.n_routed_experts() > 0 &&
|
||||
(layer_id + 1) % model_args.decoder_sparse_step() == 0;
|
||||
}
|
||||
} // namespace
|
||||
|
||||
Qwen3_5DecoderLayerImpl::Qwen3_5DecoderLayerImpl(const ModelContext& context,
|
||||
int32_t layer_id)
|
||||
: parallel_args_(context.get_parallel_args()) {
|
||||
const auto& model_args = context.get_model_args();
|
||||
const auto& quant_args = context.get_quant_args();
|
||||
const auto& options = context.get_tensor_options();
|
||||
|
||||
const bool use_moe = is_moe_layer(model_args, layer_id);
|
||||
|
||||
enable_deep_ep_ = use_moe && FLAGS_expert_parallel_degree == 2;
|
||||
if (enable_deep_ep_) {
|
||||
CHECK_EQ(parallel_args_.dp_size(), parallel_args_.world_size())
|
||||
<< "Qwen3.5 MoE only support deep ep all2all when dp_size == "
|
||||
"world_size";
|
||||
CHECK_EQ(parallel_args_.dp_size(), parallel_args_.ep_size())
|
||||
<< "Qwen3.5 MoE only support deep ep all2all when dp_size == ep_size";
|
||||
}
|
||||
|
||||
auto layer_types = model_args.layer_types();
|
||||
if (layer_types.empty()) {
|
||||
int32_t interval = model_args.full_attention_interval();
|
||||
for (int32_t i = 0; i < model_args.n_layers(); i++) {
|
||||
layer_types.push_back((i + 1) % interval == 0 ? "full_attention"
|
||||
: "linear_attention");
|
||||
}
|
||||
}
|
||||
|
||||
if (layer_id >= 0 && layer_id < static_cast<int32_t>(layer_types.size())) {
|
||||
layer_type_ = layer_types[layer_id];
|
||||
} else {
|
||||
layer_type_ = "full_attention";
|
||||
}
|
||||
|
||||
if (layer_type_ == "linear_attention") {
|
||||
// TODO: support linear attention
|
||||
} else {
|
||||
full_attention_ = register_module(
|
||||
"self_attn",
|
||||
Qwen3_5Attention(
|
||||
model_args, quant_args, parallel_args_, options, layer_id));
|
||||
}
|
||||
|
||||
input_norm_ = register_module(
|
||||
"input_layernorm",
|
||||
Qwen3NextRMSNorm(
|
||||
model_args.hidden_size(), model_args.rms_norm_eps(), options));
|
||||
|
||||
post_norm_ = register_module(
|
||||
"post_attention_layernorm",
|
||||
Qwen3NextRMSNorm(
|
||||
model_args.hidden_size(), model_args.rms_norm_eps(), options));
|
||||
|
||||
if (use_moe) {
|
||||
moe_mlp_ = register_module("mlp",
|
||||
Qwen3_5FusedMoE(model_args,
|
||||
FusedMoEArgs{.is_gated = true},
|
||||
quant_args,
|
||||
parallel_args_,
|
||||
options));
|
||||
} else {
|
||||
mlp_ = register_module("mlp",
|
||||
DenseMLP(model_args.hidden_size(),
|
||||
model_args.intermediate_size(),
|
||||
true,
|
||||
false,
|
||||
model_args.hidden_act(),
|
||||
/*enable_result_reduction=*/true,
|
||||
quant_args,
|
||||
parallel_args_.tp_group_,
|
||||
options));
|
||||
}
|
||||
}
|
||||
|
||||
void Qwen3_5DecoderLayerImpl::load_state_dict(const StateDict& state_dict) {
|
||||
if (layer_type_ == "linear_attention") {
|
||||
// TODO: support linear attention
|
||||
} else {
|
||||
full_attention_->load_state_dict(
|
||||
state_dict.get_dict_with_prefix("self_attn."));
|
||||
}
|
||||
input_norm_->load_state_dict(
|
||||
state_dict.get_dict_with_prefix("input_layernorm."));
|
||||
post_norm_->load_state_dict(
|
||||
state_dict.get_dict_with_prefix("post_attention_layernorm."));
|
||||
if (moe_mlp_) {
|
||||
moe_mlp_->load_state_dict(state_dict.get_dict_with_prefix("mlp."));
|
||||
} else {
|
||||
mlp_->load_state_dict(state_dict.get_dict_with_prefix("mlp."));
|
||||
}
|
||||
}
|
||||
|
||||
torch::Tensor Qwen3_5DecoderLayerImpl::run_moe(
|
||||
torch::Tensor x,
|
||||
const ModelInputParams& input_params) {
|
||||
const bool enable_moe_all2all =
|
||||
use_moe_all2all(enable_deep_ep_, input_params);
|
||||
if (need_dp_moe_gather(parallel_args_, enable_moe_all2all)) {
|
||||
x = gather_dp_tokens(x, input_params, parallel_args_);
|
||||
x = moe_mlp_->forward_experts(x, enable_moe_all2all);
|
||||
return get_dp_local_slice(x, input_params, parallel_args_);
|
||||
}
|
||||
return moe_mlp_->forward_experts(x, enable_moe_all2all);
|
||||
}
|
||||
|
||||
std::tuple<torch::Tensor, std::optional<torch::Tensor>>
|
||||
Qwen3_5DecoderLayerImpl::apply_norm(Qwen3NextRMSNorm& norm,
|
||||
torch::Tensor& input,
|
||||
std::optional<torch::Tensor>& residual) {
|
||||
if (!residual.has_value()) {
|
||||
auto new_residual = input;
|
||||
auto output = std::get<0>(norm->forward(input));
|
||||
return {output, new_residual};
|
||||
}
|
||||
auto orig_dtype = input.dtype();
|
||||
input = input + residual.value();
|
||||
auto new_residual = input;
|
||||
input = input.to(orig_dtype);
|
||||
auto output = std::get<0>(norm->forward(input));
|
||||
return {output, new_residual};
|
||||
}
|
||||
|
||||
torch::Tensor Qwen3_5DecoderLayerImpl::forward(
|
||||
torch::Tensor& x,
|
||||
std::optional<torch::Tensor>& residual,
|
||||
torch::Tensor& positions,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache,
|
||||
const ModelInputParams& input_params) {
|
||||
// Pre-attention norm
|
||||
std::tie(x, residual) = apply_norm(input_norm_, x, residual);
|
||||
|
||||
// Attention
|
||||
if (full_attention_) {
|
||||
x = full_attention_->forward(positions, x, attn_metadata, kv_cache);
|
||||
} else {
|
||||
// TODO: support linear attention
|
||||
}
|
||||
|
||||
auto orig_dtype = x.dtype();
|
||||
// Post-attention norm
|
||||
std::tie(x, residual) = apply_norm(post_norm_, x, residual);
|
||||
|
||||
// MLP/MoE
|
||||
if (moe_mlp_) {
|
||||
x = run_moe(x, input_params);
|
||||
} else {
|
||||
x = mlp_->forward(x);
|
||||
}
|
||||
x = x.to(orig_dtype);
|
||||
return x;
|
||||
}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
73
ex_engine/xllm_layers/mlu/qwen3_5_decoder_layer.h
Normal file
73
ex_engine/xllm_layers/mlu/qwen3_5_decoder_layer.h
Normal file
@@ -0,0 +1,73 @@
|
||||
/* 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 <optional>
|
||||
#include <string>
|
||||
|
||||
#include "framework/kv_cache/kv_cache.h"
|
||||
#include "framework/model/model_args.h"
|
||||
#include "framework/model/model_input_params.h"
|
||||
#include "framework/model_context.h"
|
||||
#include "framework/parallel_state/parallel_args.h"
|
||||
#include "framework/state_dict/state_dict.h"
|
||||
#include "layers/common/dense_mlp.h"
|
||||
#include "layers/common/qwen3_next_rms_norm.h"
|
||||
#include "layers/mlu/qwen3_5_attention.h"
|
||||
#include "layers/mlu/qwen3_5_fused_moe.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3_5DecoderLayerImpl final : public torch::nn::Module {
|
||||
public:
|
||||
Qwen3_5DecoderLayerImpl(const ModelContext& context, int32_t layer_id);
|
||||
|
||||
void load_state_dict(const StateDict& state_dict);
|
||||
|
||||
torch::Tensor forward(torch::Tensor& x,
|
||||
std::optional<torch::Tensor>& residual,
|
||||
torch::Tensor& positions,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache,
|
||||
const ModelInputParams& input_params);
|
||||
|
||||
private:
|
||||
std::tuple<torch::Tensor, std::optional<torch::Tensor>> apply_norm(
|
||||
Qwen3NextRMSNorm& norm,
|
||||
torch::Tensor& input,
|
||||
std::optional<torch::Tensor>& residual);
|
||||
|
||||
torch::Tensor run_moe(torch::Tensor x, const ModelInputParams& input_params);
|
||||
|
||||
std::string layer_type_;
|
||||
Qwen3_5Attention full_attention_{nullptr};
|
||||
// TODO: support linear attention
|
||||
// Qwen3_5GatedDeltaNet linear_attention_{nullptr};
|
||||
DenseMLP mlp_{nullptr};
|
||||
Qwen3_5FusedMoE moe_mlp_{nullptr};
|
||||
Qwen3NextRMSNorm input_norm_{nullptr};
|
||||
Qwen3NextRMSNorm post_norm_{nullptr};
|
||||
ParallelArgs parallel_args_;
|
||||
bool enable_deep_ep_ = false;
|
||||
};
|
||||
|
||||
TORCH_MODULE(Qwen3_5DecoderLayer);
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
209
ex_engine/xllm_layers/mlu/qwen3_5_fused_moe.cpp
Normal file
209
ex_engine/xllm_layers/mlu/qwen3_5_fused_moe.cpp
Normal file
@@ -0,0 +1,209 @@
|
||||
/* 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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "qwen3_5_fused_moe.h"
|
||||
|
||||
#include <glog/logging.h>
|
||||
|
||||
#include "framework/parallel_state/parallel_state.h"
|
||||
#include "framework/state_dict/utils.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
namespace {
|
||||
torch::Tensor get_tensor_with_weight_suffix(const StateDict& state_dict,
|
||||
const std::string& tensor_name) {
|
||||
auto tensor = state_dict.get_tensor(tensor_name);
|
||||
if (!tensor.defined()) {
|
||||
tensor = state_dict.get_tensor(tensor_name + ".weight");
|
||||
}
|
||||
return tensor;
|
||||
}
|
||||
|
||||
torch::Tensor slice_expert_weights(const torch::Tensor& weight,
|
||||
int64_t start_expert_id,
|
||||
int64_t num_experts_per_rank) {
|
||||
return weight
|
||||
.slice(0, start_expert_id, start_expert_id + num_experts_per_rank)
|
||||
.contiguous();
|
||||
}
|
||||
|
||||
bool load_fused_gate_up_fallback(const StateDict& state_dict,
|
||||
int64_t rank,
|
||||
int64_t world_size,
|
||||
int64_t start_expert_id,
|
||||
int64_t num_experts_per_rank,
|
||||
torch::Tensor& w13) {
|
||||
auto fused_gate_up =
|
||||
get_tensor_with_weight_suffix(state_dict, "gate_up_proj");
|
||||
if (!fused_gate_up.defined()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (world_size > 1) {
|
||||
CHECK_EQ(fused_gate_up.size(1) % 2, 0)
|
||||
<< "gate_up_proj dim1 must be even, got " << fused_gate_up.size(1);
|
||||
const int64_t full_intermediate = fused_gate_up.size(1) / 2;
|
||||
CHECK_EQ(full_intermediate % world_size, 0)
|
||||
<< "gate_up_proj intermediate dim is not divisible by world_size";
|
||||
const int64_t inter_shard = full_intermediate / world_size;
|
||||
|
||||
auto gate_full = fused_gate_up.slice(1, 0, full_intermediate);
|
||||
auto up_full =
|
||||
fused_gate_up.slice(1, full_intermediate, full_intermediate * 2);
|
||||
auto gate_shard =
|
||||
gate_full.slice(1, rank * inter_shard, (rank + 1) * inter_shard);
|
||||
auto up_shard =
|
||||
up_full.slice(1, rank * inter_shard, (rank + 1) * inter_shard);
|
||||
fused_gate_up = torch::cat({gate_shard, up_shard}, 1);
|
||||
}
|
||||
|
||||
auto gate_up_slice = slice_expert_weights(
|
||||
fused_gate_up, start_expert_id, num_experts_per_rank);
|
||||
CHECK_EQ(w13.sizes(), gate_up_slice.sizes())
|
||||
<< "weight size mismatch for " << state_dict.prefix()
|
||||
<< "experts.gate_up_proj";
|
||||
w13.copy_(gate_up_slice);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool load_fused_down_fallback(const StateDict& state_dict,
|
||||
int64_t rank,
|
||||
int64_t world_size,
|
||||
int64_t start_expert_id,
|
||||
int64_t num_experts_per_rank,
|
||||
torch::Tensor& w2) {
|
||||
auto fused_down = get_tensor_with_weight_suffix(state_dict, "down_proj");
|
||||
if (!fused_down.defined()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (world_size > 1) {
|
||||
CHECK_EQ(fused_down.size(2) % world_size, 0)
|
||||
<< "down_proj dim2 is not divisible by world_size";
|
||||
const int64_t down_shard = fused_down.size(2) / world_size;
|
||||
fused_down =
|
||||
fused_down.slice(2, rank * down_shard, (rank + 1) * down_shard);
|
||||
}
|
||||
|
||||
auto down_slice =
|
||||
slice_expert_weights(fused_down, start_expert_id, num_experts_per_rank);
|
||||
CHECK_EQ(w2.sizes(), down_slice.sizes())
|
||||
<< "weight size mismatch for " << state_dict.prefix()
|
||||
<< "experts.down_proj";
|
||||
w2.copy_(down_slice);
|
||||
return true;
|
||||
}
|
||||
} // namespace
|
||||
|
||||
Qwen3_5FusedMoEImpl::Qwen3_5FusedMoEImpl(const ModelArgs& model_args,
|
||||
const FusedMoEArgs& moe_args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options)
|
||||
: FusedMoEImpl(model_args, moe_args, quant_args, parallel_args, options) {
|
||||
if (n_shared_experts_ > 0) {
|
||||
shared_expert_gate_ = register_module(
|
||||
"shared_expert_gate",
|
||||
torch::nn::Linear(
|
||||
torch::nn::LinearOptions(hidden_size_, 1).bias(false)));
|
||||
shared_expert_gate_->weight.set_data(
|
||||
shared_expert_gate_->weight.to(options));
|
||||
}
|
||||
}
|
||||
|
||||
void Qwen3_5FusedMoEImpl::load_experts(const StateDict& state_dict) {
|
||||
FusedMoEImpl::load_experts(state_dict);
|
||||
|
||||
if (!is_smoothquant_) {
|
||||
if (!w13_is_loaded_) {
|
||||
w13_is_loaded_ = load_fused_gate_up_fallback(state_dict,
|
||||
tp_pg_->rank(),
|
||||
tp_pg_->world_size(),
|
||||
start_expert_id_,
|
||||
num_experts_per_rank_,
|
||||
w13_);
|
||||
}
|
||||
|
||||
if (!w2_is_loaded_) {
|
||||
w2_is_loaded_ = load_fused_down_fallback(state_dict,
|
||||
tp_pg_->rank(),
|
||||
tp_pg_->world_size(),
|
||||
start_expert_id_,
|
||||
num_experts_per_rank_,
|
||||
w2_);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void Qwen3_5FusedMoEImpl::load_state_dict(const StateDict& state_dict) {
|
||||
if (state_dict.size() == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (n_shared_experts_ > 0) {
|
||||
shared_experts_->load_state_dict(
|
||||
state_dict.get_dict_with_prefix("shared_expert."));
|
||||
auto weight = state_dict.get_tensor("shared_expert_gate.weight");
|
||||
if (weight.defined()) {
|
||||
weight = weight.reshape({weight.size(0), -1});
|
||||
DCHECK_EQ(shared_expert_gate_->weight.sizes(), weight.sizes())
|
||||
<< "proj weight size mismatch for " << name();
|
||||
shared_expert_gate_->weight.data().copy_(weight);
|
||||
}
|
||||
}
|
||||
gate_->load_state_dict(state_dict.get_dict_with_prefix("gate."));
|
||||
load_experts(state_dict.get_dict_with_prefix("experts."));
|
||||
}
|
||||
|
||||
void Qwen3_5FusedMoEImpl::final_comm_allreduce(
|
||||
torch::Tensor& final_hidden_states,
|
||||
const torch::Tensor& hidden_states,
|
||||
torch::Tensor& shared_expert_output) {
|
||||
auto current_stream = device_.current_stream();
|
||||
routed_stream_->wait_stream(*current_stream);
|
||||
{
|
||||
torch::StreamGuard stream_guard = routed_stream_->set_stream_guard();
|
||||
if (tp_pg_->world_size() > 1) {
|
||||
final_hidden_states = parallel_state::reduce(final_hidden_states, tp_pg_);
|
||||
}
|
||||
if (parallel_args_.ep_size() > 1) {
|
||||
final_hidden_states = parallel_state::reduce(
|
||||
final_hidden_states, parallel_args_.moe_ep_group_);
|
||||
}
|
||||
}
|
||||
|
||||
if (n_shared_experts_ > 0) {
|
||||
shared_stream_->wait_stream(*current_stream);
|
||||
torch::StreamGuard stream_guard = shared_stream_->set_stream_guard();
|
||||
shared_expert_output = shared_experts_(hidden_states);
|
||||
if (shared_expert_gate_) {
|
||||
auto gate = torch::sigmoid(shared_expert_gate_->forward(hidden_states));
|
||||
shared_expert_output = gate * shared_expert_output;
|
||||
}
|
||||
shared_expert_output =
|
||||
shared_expert_output.reshape({-1, shared_expert_output.size(-1)});
|
||||
}
|
||||
|
||||
// join for parallelization
|
||||
current_stream->wait_stream(*routed_stream_);
|
||||
if (n_shared_experts_ > 0) {
|
||||
current_stream->wait_stream(*shared_stream_);
|
||||
final_hidden_states += shared_expert_output;
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
47
ex_engine/xllm_layers/mlu/qwen3_5_fused_moe.h
Normal file
47
ex_engine/xllm_layers/mlu/qwen3_5_fused_moe.h
Normal file
@@ -0,0 +1,47 @@
|
||||
/* 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 "layers/mlu/fused_moe.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3_5FusedMoEImpl final : public FusedMoEImpl {
|
||||
public:
|
||||
Qwen3_5FusedMoEImpl() = default;
|
||||
|
||||
Qwen3_5FusedMoEImpl(const ModelArgs& model_args,
|
||||
const FusedMoEArgs& moe_args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options);
|
||||
|
||||
void load_state_dict(const StateDict& state_dict) override;
|
||||
|
||||
protected:
|
||||
void final_comm_allreduce(torch::Tensor& final_hidden_states,
|
||||
const torch::Tensor& hidden_states,
|
||||
torch::Tensor& shared_expert_output) override;
|
||||
|
||||
private:
|
||||
void load_experts(const StateDict& state_dict);
|
||||
torch::nn::Linear shared_expert_gate_{nullptr};
|
||||
};
|
||||
|
||||
TORCH_MODULE(Qwen3_5FusedMoE);
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
28
ex_engine/xllm_layers/npu_torch/CMakeLists.txt
Executable file
28
ex_engine/xllm_layers/npu_torch/CMakeLists.txt
Executable file
@@ -0,0 +1,28 @@
|
||||
include(cc_library)
|
||||
|
||||
cc_library(
|
||||
NAME
|
||||
npu_torch_layers
|
||||
HDRS
|
||||
fused_moe.h
|
||||
attention.h
|
||||
qwen3_gated_delta_net_base.h
|
||||
qwen3_next_attention.h
|
||||
qwen3_next_gated_delta_net.h
|
||||
qwen3_5_gated_delta_net.h
|
||||
qwen3_next_hybrid_decoder_layer_base.h
|
||||
qwen3_next_decoder_layer_impl.h
|
||||
qwen3_5_decoder_layer_impl.h
|
||||
SRCS
|
||||
fused_moe.cpp
|
||||
attention.cpp
|
||||
qwen3_gated_delta_net_base.cpp
|
||||
qwen3_next_attention.cpp
|
||||
qwen3_next_gated_delta_net.cpp
|
||||
qwen3_next_hybrid_decoder_layer_base.cpp
|
||||
qwen3_5_gated_delta_net.cpp
|
||||
qwen3_next_decoder_layer_impl.cpp
|
||||
qwen3_5_decoder_layer_impl.cpp
|
||||
DEPS
|
||||
:common_layers
|
||||
)
|
||||
152
ex_engine/xllm_layers/npu_torch/attention.cpp
Normal file
152
ex_engine/xllm_layers/npu_torch/attention.cpp
Normal file
@@ -0,0 +1,152 @@
|
||||
/* 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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "attention.h"
|
||||
|
||||
#include "kernels/npu/npu_ops_api.h"
|
||||
#include "kernels/ops_api.h"
|
||||
|
||||
DECLARE_bool(enable_chunked_prefill);
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
AttentionImpl::AttentionImpl(int64_t num_heads,
|
||||
int64_t head_size,
|
||||
float scale,
|
||||
int64_t num_kv_heads,
|
||||
int64_t sliding_window)
|
||||
: num_heads_(num_heads),
|
||||
head_size_(head_size),
|
||||
num_kv_heads_(num_kv_heads),
|
||||
sliding_window_(sliding_window),
|
||||
scale_(scale) {
|
||||
if (sliding_window_ > -1) {
|
||||
sliding_window_ = sliding_window_ - 1;
|
||||
}
|
||||
}
|
||||
|
||||
std::tuple<torch::Tensor, std::optional<torch::Tensor>> AttentionImpl::forward(
|
||||
const AttentionMetadata& attn_metadata,
|
||||
torch::Tensor& query,
|
||||
torch::Tensor& key,
|
||||
torch::Tensor& value,
|
||||
KVCache& kv_cache) {
|
||||
std::optional<torch::Tensor> output_lse = std::nullopt;
|
||||
torch::Tensor output = torch::empty_like(query);
|
||||
|
||||
if (attn_metadata.is_dummy) {
|
||||
return std::make_tuple(output, output_lse);
|
||||
}
|
||||
|
||||
bool only_prefill =
|
||||
attn_metadata.is_prefill || attn_metadata.is_chunked_prefill;
|
||||
|
||||
torch::Tensor k_cache = kv_cache.get_k_cache();
|
||||
torch::Tensor v = value.view({-1, num_kv_heads_, head_size_});
|
||||
std::optional<torch::Tensor> v_cache = kv_cache.get_v_cache();
|
||||
|
||||
// Reshape and cache key/value
|
||||
xllm::kernel::ReshapePagedCacheParams reshape_paged_cache_params;
|
||||
reshape_paged_cache_params.key = key.view({-1, num_kv_heads_, head_size_});
|
||||
reshape_paged_cache_params.value = v;
|
||||
reshape_paged_cache_params.k_cache = k_cache;
|
||||
reshape_paged_cache_params.v_cache = v_cache;
|
||||
reshape_paged_cache_params.slot_mapping = attn_metadata.slot_mapping;
|
||||
xllm::kernel::reshape_paged_cache(reshape_paged_cache_params);
|
||||
|
||||
if (only_prefill) {
|
||||
prefill_forward(query, key, value, output, k_cache, v_cache, attn_metadata);
|
||||
} else {
|
||||
decoder_forward(query, output, k_cache, v_cache, attn_metadata);
|
||||
}
|
||||
|
||||
output = output.view({-1, num_heads_ * head_size_});
|
||||
return {output, output_lse};
|
||||
}
|
||||
|
||||
void AttentionImpl::prefill_forward(torch::Tensor& query,
|
||||
torch::Tensor& key,
|
||||
torch::Tensor& value,
|
||||
torch::Tensor& output,
|
||||
const torch::Tensor& k_cache,
|
||||
const std::optional<torch::Tensor>& v_cache,
|
||||
const AttentionMetadata& attn_metadata) {
|
||||
query = query.view({-1, num_heads_, head_size_});
|
||||
output = output.view({-1, num_heads_, head_size_});
|
||||
|
||||
if (attn_metadata.is_prefill) {
|
||||
key = key.view({-1, num_kv_heads_, head_size_});
|
||||
value = value.view({-1, num_kv_heads_, head_size_});
|
||||
|
||||
xllm::kernel::npu::batch_prefill(query,
|
||||
key,
|
||||
value,
|
||||
attn_metadata.attn_mask,
|
||||
attn_metadata.kv_seq_lens_host,
|
||||
scale_,
|
||||
output);
|
||||
} else if (attn_metadata.is_chunked_prefill) {
|
||||
xllm::kernel::npu::batch_prefill(query,
|
||||
k_cache,
|
||||
v_cache.value(),
|
||||
attn_metadata.attn_mask,
|
||||
attn_metadata.kv_seq_lens_host,
|
||||
scale_,
|
||||
output);
|
||||
}
|
||||
}
|
||||
|
||||
void AttentionImpl::decoder_forward(torch::Tensor& query,
|
||||
torch::Tensor& output,
|
||||
const torch::Tensor& k_cache,
|
||||
const std::optional<torch::Tensor>& v_cache,
|
||||
const AttentionMetadata& attn_metadata) {
|
||||
query = query.view({-1, 1, num_heads_, head_size_});
|
||||
output = output.view({-1, 1, num_heads_, head_size_});
|
||||
|
||||
torch::Tensor kv_seq_lens;
|
||||
if (attn_metadata.kv_seq_lens_host.defined()) {
|
||||
kv_seq_lens = attn_metadata.kv_seq_lens_host;
|
||||
} else {
|
||||
// Fallback if host tensor isn't prepared.
|
||||
kv_seq_lens = attn_metadata.kv_seq_lens;
|
||||
}
|
||||
|
||||
if (attn_metadata.paged_attention_tiling_data.defined()) {
|
||||
// Use CustomPagedAttention for ACL graph mode to avoid .to(kCPU) operations
|
||||
|
||||
xllm::kernel::npu::batch_decode_acl_graph(
|
||||
query,
|
||||
k_cache,
|
||||
v_cache.value_or(torch::Tensor()),
|
||||
scale_,
|
||||
attn_metadata.block_table,
|
||||
kv_seq_lens,
|
||||
attn_metadata.paged_attention_tiling_data,
|
||||
output);
|
||||
} else {
|
||||
// Standard PagedAttention path
|
||||
xllm::kernel::npu::batch_decode(query,
|
||||
k_cache,
|
||||
v_cache.value_or(torch::Tensor()),
|
||||
scale_,
|
||||
attn_metadata.block_table,
|
||||
kv_seq_lens,
|
||||
output);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
70
ex_engine/xllm_layers/npu_torch/attention.h
Normal file
70
ex_engine/xllm_layers/npu_torch/attention.h
Normal file
@@ -0,0 +1,70 @@
|
||||
/* 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 <torch/torch.h>
|
||||
|
||||
#include <tuple>
|
||||
|
||||
#include "framework/kv_cache/kv_cache.h"
|
||||
#include "framework/model/model_input_params.h"
|
||||
#include "layers/common/attention_metadata.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class AttentionImpl : public torch::nn::Module {
|
||||
public:
|
||||
AttentionImpl() = default;
|
||||
|
||||
AttentionImpl(int64_t num_heads,
|
||||
int64_t head_size,
|
||||
float scale,
|
||||
int64_t num_kv_heads,
|
||||
int64_t sliding_window);
|
||||
|
||||
std::tuple<torch::Tensor, std::optional<torch::Tensor>> forward(
|
||||
const AttentionMetadata& attn_metadata,
|
||||
torch::Tensor& query,
|
||||
torch::Tensor& key,
|
||||
torch::Tensor& value,
|
||||
KVCache& kv_cache);
|
||||
|
||||
void prefill_forward(torch::Tensor& query,
|
||||
torch::Tensor& key,
|
||||
torch::Tensor& value,
|
||||
torch::Tensor& output,
|
||||
const torch::Tensor& k_cache,
|
||||
const std::optional<torch::Tensor>& v_cache,
|
||||
const AttentionMetadata& attn_metadata);
|
||||
|
||||
void decoder_forward(torch::Tensor& query,
|
||||
torch::Tensor& output,
|
||||
const torch::Tensor& k_cache,
|
||||
const std::optional<torch::Tensor>& v_cache,
|
||||
const AttentionMetadata& attn_metadata);
|
||||
|
||||
private:
|
||||
int64_t num_heads_;
|
||||
int64_t head_size_;
|
||||
float scale_;
|
||||
int64_t num_kv_heads_;
|
||||
int64_t sliding_window_;
|
||||
};
|
||||
TORCH_MODULE(Attention);
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
513
ex_engine/xllm_layers/npu_torch/fused_moe.cpp
Normal file
513
ex_engine/xllm_layers/npu_torch/fused_moe.cpp
Normal file
@@ -0,0 +1,513 @@
|
||||
/* 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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "fused_moe.h"
|
||||
|
||||
#include <glog/logging.h>
|
||||
|
||||
#include <numeric>
|
||||
#include <vector>
|
||||
|
||||
#include "framework/parallel_state/parallel_state.h"
|
||||
#include "kernels/ops_api.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
namespace {
|
||||
// Generic local tensor helpers.
|
||||
torch::Tensor create_group_gemm_output(
|
||||
const torch::Tensor& a,
|
||||
const torch::Tensor& b,
|
||||
const torch::Tensor& group_list,
|
||||
torch::ScalarType dtype = torch::ScalarType::BFloat16) {
|
||||
torch::TensorOptions target_options = a.options().dtype(dtype);
|
||||
if (b.dim() != 2) {
|
||||
return torch::empty({a.size(0), b.size(1)}, target_options);
|
||||
}
|
||||
return torch::empty({group_list.size(0), a.size(0), b.size(0)},
|
||||
target_options);
|
||||
}
|
||||
|
||||
torch::Tensor get_tensor_with_weight_suffix(const StateDict& state_dict,
|
||||
const std::string& tensor_name) {
|
||||
auto tensor = state_dict.get_tensor(tensor_name);
|
||||
if (!tensor.defined()) {
|
||||
tensor = state_dict.get_tensor(tensor_name + ".weight");
|
||||
}
|
||||
return tensor;
|
||||
}
|
||||
|
||||
torch::Tensor slice_expert_weights(const torch::Tensor& weight,
|
||||
int64_t start_expert_id,
|
||||
int64_t num_experts_per_rank) {
|
||||
return weight
|
||||
.slice(0, start_expert_id, start_expert_id + num_experts_per_rank)
|
||||
.contiguous();
|
||||
}
|
||||
|
||||
// Qwen3.5-MoE fused checkpoint fallback helpers.
|
||||
bool load_fused_gate_up_fallback(const StateDict& state_dict,
|
||||
int64_t rank,
|
||||
int64_t world_size,
|
||||
int64_t start_expert_id,
|
||||
int64_t num_experts_per_rank,
|
||||
torch::Tensor& w13) {
|
||||
auto fused_gate_up =
|
||||
get_tensor_with_weight_suffix(state_dict, "gate_up_proj");
|
||||
if (!fused_gate_up.defined()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (world_size > 1) {
|
||||
CHECK_EQ(fused_gate_up.size(1) % 2, 0)
|
||||
<< "gate_up_proj dim1 must be even, got " << fused_gate_up.size(1);
|
||||
const int64_t full_intermediate = fused_gate_up.size(1) / 2;
|
||||
CHECK_EQ(full_intermediate % world_size, 0)
|
||||
<< "gate_up_proj intermediate dim is not divisible by world_size";
|
||||
const int64_t inter_shard = full_intermediate / world_size;
|
||||
|
||||
auto gate_full = fused_gate_up.slice(1, 0, full_intermediate);
|
||||
auto up_full =
|
||||
fused_gate_up.slice(1, full_intermediate, full_intermediate * 2);
|
||||
auto gate_shard =
|
||||
gate_full.slice(1, rank * inter_shard, (rank + 1) * inter_shard);
|
||||
auto up_shard =
|
||||
up_full.slice(1, rank * inter_shard, (rank + 1) * inter_shard);
|
||||
fused_gate_up = torch::cat({gate_shard, up_shard}, 1);
|
||||
}
|
||||
|
||||
auto gate_up_slice = slice_expert_weights(
|
||||
fused_gate_up, start_expert_id, num_experts_per_rank);
|
||||
CHECK_EQ(w13.sizes(), gate_up_slice.sizes())
|
||||
<< "weight size mismatch for " << state_dict.prefix()
|
||||
<< "experts.gate_up_proj";
|
||||
w13.copy_(gate_up_slice);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool load_fused_down_fallback(const StateDict& state_dict,
|
||||
int64_t rank,
|
||||
int64_t world_size,
|
||||
int64_t start_expert_id,
|
||||
int64_t num_experts_per_rank,
|
||||
torch::Tensor& w2) {
|
||||
auto fused_down = get_tensor_with_weight_suffix(state_dict, "down_proj");
|
||||
if (!fused_down.defined()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (world_size > 1) {
|
||||
CHECK_EQ(fused_down.size(2) % world_size, 0)
|
||||
<< "down_proj dim2 is not divisible by world_size";
|
||||
const int64_t down_shard = fused_down.size(2) / world_size;
|
||||
fused_down =
|
||||
fused_down.slice(2, rank * down_shard, (rank + 1) * down_shard);
|
||||
}
|
||||
|
||||
auto down_slice =
|
||||
slice_expert_weights(fused_down, start_expert_id, num_experts_per_rank);
|
||||
CHECK_EQ(w2.sizes(), down_slice.sizes())
|
||||
<< "weight size mismatch for " << state_dict.prefix()
|
||||
<< "experts.down_proj";
|
||||
w2.copy_(down_slice);
|
||||
return true;
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
FusedMoEImpl::FusedMoEImpl(const ModelArgs& model_args,
|
||||
const FusedMoEArgs& moe_args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options)
|
||||
: num_total_experts_(model_args.n_routed_experts()),
|
||||
topk_(model_args.num_experts_per_tok()),
|
||||
hidden_size_(model_args.hidden_size()),
|
||||
n_shared_experts_(model_args.n_shared_experts()),
|
||||
is_gated_(moe_args.is_gated),
|
||||
renormalize_(model_args.norm_topk_prob() ? 1 : 0),
|
||||
hidden_act_(model_args.hidden_act()),
|
||||
is_smoothquant_(false),
|
||||
quant_args_(quant_args),
|
||||
parallel_args_(parallel_args),
|
||||
options_(options),
|
||||
tp_pg_(parallel_args.tp_group_) {
|
||||
const int64_t num_experts = num_total_experts_;
|
||||
const int64_t intermediate_size =
|
||||
static_cast<int64_t>(model_args.moe_intermediate_size());
|
||||
const std::string& topk_method = model_args.topk_method();
|
||||
int64_t ep_size = parallel_args.ep_size();
|
||||
int64_t ep_rank = 0;
|
||||
if (ep_size > 1) {
|
||||
ep_rank = parallel_args.moe_ep_group_->rank();
|
||||
tp_pg_ = parallel_args.moe_tp_group_;
|
||||
}
|
||||
|
||||
// smoothquant check: If quant_method is not empty, only w8a8 smoothquant is
|
||||
// supported
|
||||
if (!quant_args.quant_method().empty()) {
|
||||
if (quant_args.quant_method() != "smoothquant" || quant_args.bits() != 8 ||
|
||||
!quant_args.activation_dynamic()) {
|
||||
LOG(FATAL) << "FusedMoE only supports w8a8 smoothquant quantization when "
|
||||
"quant_method is set. "
|
||||
<< "Got quant_method=" << quant_args.quant_method()
|
||||
<< ", bits=" << quant_args.bits()
|
||||
<< ", activation_dynamic=" << quant_args.activation_dynamic();
|
||||
}
|
||||
// If confirmed as smoothquant w8a8, set is_smoothquant_ to true
|
||||
is_smoothquant_ = true;
|
||||
} else {
|
||||
is_smoothquant_ = false;
|
||||
}
|
||||
|
||||
// calculate the number of experts per rank
|
||||
num_experts_per_rank_ = num_experts / ep_size;
|
||||
start_expert_id_ = ep_rank * num_experts_per_rank_;
|
||||
|
||||
if (topk_method == "noaux_tc") {
|
||||
e_score_correction_bias_ = register_parameter(
|
||||
"e_score_correction_bias", torch::empty({num_experts}, options), false);
|
||||
}
|
||||
|
||||
gate_ = register_module(
|
||||
"gate_proj",
|
||||
ReplicatedLinear(hidden_size_, num_experts, false, quant_args, options));
|
||||
if (n_shared_experts_ > 0) {
|
||||
/*
|
||||
The shared_experts are usually implemented using the RowParallelLinear
|
||||
layer. Typically, this output serves as the enable_result_reduction results
|
||||
for the module. If only tensor parallelism is applied, immediate
|
||||
reduction of the shared_experts output isn't necessary; instead, we perform
|
||||
the reduction once at the end of the MoE operation.
|
||||
*/
|
||||
shared_experts_ =
|
||||
register_module("shared_experts",
|
||||
DenseMLP(hidden_size_,
|
||||
intermediate_size * n_shared_experts_,
|
||||
is_gated_,
|
||||
false,
|
||||
hidden_act_,
|
||||
/*enable_result_reduction=*/false,
|
||||
quant_args,
|
||||
tp_pg_,
|
||||
options));
|
||||
shared_expert_gate_ = register_module(
|
||||
"shared_expert_gate",
|
||||
torch::nn::Linear(
|
||||
torch::nn::LinearOptions(hidden_size_, 1).bias(false)));
|
||||
shared_expert_gate_->weight.set_data(
|
||||
shared_expert_gate_->weight.to(options));
|
||||
}
|
||||
|
||||
// create weight buffer
|
||||
const int64_t world_size = tp_pg_->world_size();
|
||||
int64_t local_intermediate_size = intermediate_size / world_size;
|
||||
if (is_smoothquant_) {
|
||||
auto quant_option = options_.dtype(torch::kInt8);
|
||||
auto fp_option = options_.dtype(torch::kFloat32);
|
||||
w13_ = register_parameter(
|
||||
"w13",
|
||||
torch::empty(
|
||||
{num_experts_per_rank_, local_intermediate_size * 2, hidden_size_},
|
||||
quant_option),
|
||||
false);
|
||||
w13_scale_ = register_parameter(
|
||||
"w13_scale",
|
||||
torch::empty({num_experts_per_rank_, local_intermediate_size * 2},
|
||||
fp_option),
|
||||
false);
|
||||
input_smooth_ = register_parameter(
|
||||
"input_smooth",
|
||||
torch::empty({num_experts_per_rank_, hidden_size_}, fp_option),
|
||||
false);
|
||||
w2_ = register_parameter(
|
||||
"w2",
|
||||
torch::empty(
|
||||
{num_experts_per_rank_, hidden_size_, local_intermediate_size},
|
||||
quant_option),
|
||||
false);
|
||||
w2_scale_ = register_parameter(
|
||||
"w2_scale",
|
||||
torch::empty({num_experts_per_rank_, hidden_size_}, fp_option),
|
||||
false);
|
||||
act_smooth_ = register_parameter(
|
||||
"act_smooth",
|
||||
torch::empty({num_experts_per_rank_, local_intermediate_size},
|
||||
fp_option),
|
||||
false);
|
||||
|
||||
} else {
|
||||
w13_ = register_parameter(
|
||||
"w13",
|
||||
torch::empty(
|
||||
{num_experts_per_rank_, local_intermediate_size * 2, hidden_size_},
|
||||
options_),
|
||||
false);
|
||||
w2_ = register_parameter(
|
||||
"w2",
|
||||
torch::empty(
|
||||
{num_experts_per_rank_, hidden_size_, local_intermediate_size},
|
||||
options_),
|
||||
false);
|
||||
}
|
||||
}
|
||||
|
||||
torch::Tensor FusedMoEImpl::select_experts(
|
||||
const torch::Tensor& hidden_states_2d,
|
||||
const torch::Tensor& router_logits_2d,
|
||||
SelectedExpertInfo& selected_expert_info) {
|
||||
// prepare the parameters for select_experts
|
||||
xllm::kernel::MoeFusedTopkParams moe_active_topk_params;
|
||||
moe_active_topk_params.input = router_logits_2d;
|
||||
moe_active_topk_params.finished = torch::Tensor();
|
||||
moe_active_topk_params.topk = topk_;
|
||||
moe_active_topk_params.scoring_func = "softmax";
|
||||
auto [topk_weights, topk_ids] =
|
||||
xllm::kernel::moe_active_topk(moe_active_topk_params);
|
||||
topk_ids = topk_ids.to(torch::kInt32);
|
||||
if (renormalize_) {
|
||||
topk_weights = topk_weights / (topk_weights.sum(-1, true) + 1e-6);
|
||||
}
|
||||
|
||||
xllm::kernel::MoeInitRoutingV2Params moe_init_routing_params;
|
||||
moe_init_routing_params.x = hidden_states_2d;
|
||||
moe_init_routing_params.expert_idx = topk_ids;
|
||||
moe_init_routing_params.scale = std::nullopt;
|
||||
moe_init_routing_params.offset = std::nullopt;
|
||||
moe_init_routing_params.active_num = hidden_states_2d.size(0) * topk_;
|
||||
moe_init_routing_params.expert_capacity = 0;
|
||||
moe_init_routing_params.expert_num = num_experts_per_rank_;
|
||||
moe_init_routing_params.drop_pad_mode = 0;
|
||||
moe_init_routing_params.expert_tokens_num_type = 1;
|
||||
moe_init_routing_params.expert_tokens_num_flag = true;
|
||||
moe_init_routing_params.row_idx_type = 0;
|
||||
std::vector<int64_t> expert_range = {
|
||||
start_expert_id_, start_expert_id_ + num_experts_per_rank_};
|
||||
moe_init_routing_params.active_expert_range = expert_range;
|
||||
moe_init_routing_params.quant_mode = -1;
|
||||
// TODO: NPU moe_init_routing_v2 is equivalent to moe_gen_idx +
|
||||
// moe_expand_input (and the token_count/cusum outputs) on other backends.
|
||||
auto [expand_hidden_states, expand_row_ids, group_list, dynamic_scale] =
|
||||
xllm::kernel::moe_init_routing_v2(moe_init_routing_params);
|
||||
(void)dynamic_scale;
|
||||
|
||||
// collect the selected tensor
|
||||
selected_expert_info.reduce_weight = topk_weights;
|
||||
selected_expert_info.combine_idx = expand_row_ids;
|
||||
selected_expert_info.token_count_slice = group_list;
|
||||
selected_expert_info.cusum_token_count = group_list;
|
||||
return expand_hidden_states;
|
||||
}
|
||||
|
||||
torch::Tensor FusedMoEImpl::forward_expert(
|
||||
const torch::Tensor& hidden_states,
|
||||
const torch::Tensor& router_logits,
|
||||
const std::optional<torch::Tensor>& shared_output) {
|
||||
// prepare the parameters for MoE computation
|
||||
torch::IntArrayRef hidden_states_shape = hidden_states.sizes();
|
||||
torch::ScalarType hidden_states_dtype = hidden_states.dtype().toScalarType();
|
||||
torch::Tensor hidden_states_2d =
|
||||
hidden_states.reshape({-1, hidden_states.size(-1)});
|
||||
torch::Tensor router_logits_2d =
|
||||
router_logits.reshape({-1, router_logits.size(-1)});
|
||||
|
||||
// Step 1-3: select experts
|
||||
SelectedExpertInfo selected_expert_info;
|
||||
torch::Tensor expand_hidden_states =
|
||||
select_experts(hidden_states_2d, router_logits_2d, selected_expert_info);
|
||||
|
||||
// Step 4: group gemm 1
|
||||
torch::Tensor gemm1_out =
|
||||
create_group_gemm_output(expand_hidden_states,
|
||||
w13_,
|
||||
selected_expert_info.token_count_slice,
|
||||
hidden_states_dtype);
|
||||
|
||||
{
|
||||
xllm::kernel::GroupGemmParams group_gemm_params;
|
||||
group_gemm_params.a = expand_hidden_states;
|
||||
if (w13_.size(1) != expand_hidden_states.size(1)) {
|
||||
w13_ = w13_.transpose(1, 2);
|
||||
}
|
||||
group_gemm_params.b = w13_;
|
||||
group_gemm_params.group_list = selected_expert_info.token_count_slice;
|
||||
group_gemm_params.split_item = 2;
|
||||
group_gemm_params.group_type = 0;
|
||||
group_gemm_params.group_list_type = 1;
|
||||
gemm1_out = xllm::kernel::group_gemm(group_gemm_params);
|
||||
}
|
||||
|
||||
// Step 5: activation
|
||||
torch::Tensor act_out;
|
||||
|
||||
xllm::kernel::ActivationParams activation_params;
|
||||
activation_params.input = gemm1_out;
|
||||
activation_params.output = act_out;
|
||||
activation_params.act_mode = hidden_act_;
|
||||
activation_params.is_gated = is_gated_;
|
||||
xllm::kernel::active(activation_params);
|
||||
act_out = activation_params.output;
|
||||
// Step 6: group gemm 2
|
||||
torch::Tensor gemm2_out =
|
||||
create_group_gemm_output(act_out,
|
||||
w2_,
|
||||
selected_expert_info.token_count_slice,
|
||||
hidden_states_dtype);
|
||||
|
||||
{
|
||||
xllm::kernel::GroupGemmParams group_gemm_params;
|
||||
group_gemm_params.a = act_out;
|
||||
if (w2_.size(1) != act_out.size(1)) {
|
||||
w2_ = w2_.transpose(1, 2);
|
||||
}
|
||||
group_gemm_params.b = w2_;
|
||||
group_gemm_params.group_list = selected_expert_info.token_count_slice;
|
||||
group_gemm_params.split_item = 2;
|
||||
group_gemm_params.group_type = 0;
|
||||
group_gemm_params.group_list_type = 1;
|
||||
gemm2_out = xllm::kernel::group_gemm(group_gemm_params);
|
||||
}
|
||||
|
||||
// Step 7: combine the intermediate results and get the final hidden states
|
||||
torch::Tensor final_hidden_states;
|
||||
xllm::kernel::MoeCombineResultParams moe_combine_params;
|
||||
moe_combine_params.input = gemm2_out;
|
||||
moe_combine_params.reduce_weight = selected_expert_info.reduce_weight;
|
||||
moe_combine_params.gather_ids = selected_expert_info.combine_idx;
|
||||
final_hidden_states = xllm::kernel::moe_combine_result(moe_combine_params);
|
||||
if (shared_output.has_value()) {
|
||||
final_hidden_states = final_hidden_states + shared_output.value();
|
||||
}
|
||||
// reshape the final hidden states to the original shape
|
||||
final_hidden_states = final_hidden_states.reshape(hidden_states_shape);
|
||||
|
||||
if (tp_pg_->world_size() > 1) {
|
||||
final_hidden_states = parallel_state::reduce(final_hidden_states, tp_pg_);
|
||||
}
|
||||
if (parallel_args_.ep_size() > 1) {
|
||||
final_hidden_states = parallel_state::reduce(final_hidden_states,
|
||||
parallel_args_.moe_ep_group_);
|
||||
}
|
||||
return final_hidden_states;
|
||||
}
|
||||
|
||||
torch::Tensor FusedMoEImpl::forward(const torch::Tensor& hidden_states,
|
||||
const ModelInputParams& input_params) {
|
||||
auto input = hidden_states;
|
||||
bool need_slice = false;
|
||||
if (parallel_args_.dp_size() > 1 && parallel_args_.ep_size() > 1) {
|
||||
input = parallel_state::gather(input,
|
||||
parallel_args_.dp_local_process_group_,
|
||||
input_params.dp_global_token_nums);
|
||||
need_slice = true;
|
||||
}
|
||||
|
||||
std::optional<torch::Tensor> shared_output = std::nullopt;
|
||||
if (n_shared_experts_ > 0) {
|
||||
shared_output = shared_experts_(input);
|
||||
if (shared_expert_gate_) {
|
||||
auto gate = torch::sigmoid(shared_expert_gate_->forward(input));
|
||||
if (shared_output.has_value()) {
|
||||
torch::Tensor res = gate * shared_output.value();
|
||||
shared_output = res;
|
||||
}
|
||||
}
|
||||
}
|
||||
auto router_logits = gate_(input);
|
||||
auto output = forward_expert(input, router_logits, shared_output);
|
||||
|
||||
if (need_slice) {
|
||||
const auto& dp_tokens = input_params.dp_global_token_nums;
|
||||
const int64_t dp_rank = parallel_args_.dp_local_process_group_->rank();
|
||||
auto start =
|
||||
std::accumulate(dp_tokens.begin(), dp_tokens.begin() + dp_rank, 0);
|
||||
auto end = start + dp_tokens[dp_rank];
|
||||
output = output.slice(0, start, end);
|
||||
}
|
||||
return output;
|
||||
}
|
||||
|
||||
void FusedMoEImpl::load_e_score_correction_bias(const StateDict& state_dict) {
|
||||
if (e_score_correction_bias_.defined() &&
|
||||
!e_score_correction_bias_is_loaded_) {
|
||||
LOAD_WEIGHT(e_score_correction_bias);
|
||||
}
|
||||
}
|
||||
|
||||
void FusedMoEImpl::load_experts(const StateDict& state_dict) {
|
||||
const int64_t rank = tp_pg_->rank();
|
||||
const int64_t world_size = tp_pg_->world_size();
|
||||
const int64_t start_expert_id = start_expert_id_;
|
||||
const int64_t num_experts_per_rank = num_experts_per_rank_;
|
||||
std::vector<std::string> prefixes = {"gate_proj.", "up_proj."};
|
||||
if (is_smoothquant_) {
|
||||
LOAD_MOE_FUSED_WEIGHT("qweight", w1, w3, w13);
|
||||
LOAD_MOE_FUSED_WEIGHT("per_channel_scale", w1_scale, w3_scale, w13_scale);
|
||||
LOAD_MOE_WEIGHT("up_proj.", "smooth", input_smooth, -1);
|
||||
LOAD_MOE_WEIGHT("down_proj.", "qweight", w2, 1);
|
||||
LOAD_MOE_WEIGHT("down_proj.", "per_channel_scale", w2_scale, -1);
|
||||
LOAD_MOE_WEIGHT("down_proj.", "smooth", act_smooth, 0);
|
||||
} else {
|
||||
LOAD_MOE_FUSED_WEIGHT("weight", w1, w3, w13);
|
||||
LOAD_MOE_WEIGHT("down_proj.", "weight", w2, 1);
|
||||
|
||||
// Some Qwen3.5-MoE checkpoints store expert weights in fused tensors
|
||||
// (gate_up_proj / down_proj). Fall back to this format when split
|
||||
// gate_proj/up_proj tensors are absent.
|
||||
if (!w13_is_loaded_) {
|
||||
w13_is_loaded_ = load_fused_gate_up_fallback(state_dict,
|
||||
rank,
|
||||
world_size,
|
||||
start_expert_id,
|
||||
num_experts_per_rank,
|
||||
w13_);
|
||||
}
|
||||
|
||||
if (!w2_is_loaded_) {
|
||||
w2_is_loaded_ = load_fused_down_fallback(state_dict,
|
||||
rank,
|
||||
world_size,
|
||||
start_expert_id,
|
||||
num_experts_per_rank,
|
||||
w2_);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void FusedMoEImpl::load_state_dict(const StateDict& state_dict) {
|
||||
if (state_dict.size() == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (n_shared_experts_ > 0) {
|
||||
shared_experts_->load_state_dict(
|
||||
state_dict.get_dict_with_prefix("shared_expert."));
|
||||
auto weight = state_dict.get_tensor("shared_expert_gate.weight");
|
||||
if (weight.defined()) {
|
||||
weight = weight.reshape({weight.size(0), -1});
|
||||
DCHECK_EQ(shared_expert_gate_->weight.sizes(), weight.sizes())
|
||||
<< "proj weight size mismatch for " << name();
|
||||
shared_expert_gate_->weight.data().copy_(weight);
|
||||
}
|
||||
}
|
||||
|
||||
gate_->load_state_dict(state_dict.get_dict_with_prefix("gate."));
|
||||
load_e_score_correction_bias(state_dict.get_dict_with_prefix("gate."));
|
||||
load_experts(state_dict.get_dict_with_prefix("experts."));
|
||||
}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
113
ex_engine/xllm_layers/npu_torch/fused_moe.h
Normal file
113
ex_engine/xllm_layers/npu_torch/fused_moe.h
Normal file
@@ -0,0 +1,113 @@
|
||||
/* 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 <torch/torch.h>
|
||||
|
||||
#include <optional>
|
||||
|
||||
#include "framework/model/model_args.h"
|
||||
#include "framework/model/model_input_params.h"
|
||||
#include "framework/parallel_state/parallel_args.h"
|
||||
#include "framework/quant_args.h"
|
||||
#include "framework/state_dict/state_dict.h"
|
||||
#include "framework/state_dict/utils.h"
|
||||
#include "layers/common/dense_mlp.h"
|
||||
#include "layers/common/fused_moe_base.h"
|
||||
#include "layers/common/linear.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class FusedMoEImpl : public torch::nn::Module {
|
||||
public:
|
||||
FusedMoEImpl() = default;
|
||||
FusedMoEImpl(const ModelArgs& model_args,
|
||||
const FusedMoEArgs& moe_args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options);
|
||||
|
||||
torch::Tensor forward_expert(
|
||||
const torch::Tensor& hidden_states,
|
||||
const torch::Tensor& router_logits,
|
||||
const std::optional<torch::Tensor>& shared_output);
|
||||
torch::Tensor forward(const torch::Tensor& hidden_states,
|
||||
const ModelInputParams& input_params);
|
||||
void load_state_dict(const StateDict& state_dict);
|
||||
|
||||
private:
|
||||
// struct to store the selected expert info
|
||||
struct SelectedExpertInfo {
|
||||
torch::Tensor reduce_weight;
|
||||
torch::Tensor combine_idx;
|
||||
torch::Tensor token_count_slice;
|
||||
torch::Tensor cusum_token_count;
|
||||
std::optional<torch::Tensor> input_scale;
|
||||
};
|
||||
|
||||
// initial steps for MoE computation, select the experts for each token
|
||||
torch::Tensor select_experts(const torch::Tensor& hidden_states_2d,
|
||||
const torch::Tensor& router_logits_2d,
|
||||
SelectedExpertInfo& selected_expert_info);
|
||||
|
||||
private:
|
||||
int64_t num_total_experts_;
|
||||
int64_t topk_;
|
||||
int64_t num_expert_group_;
|
||||
int64_t topk_group_;
|
||||
double route_scale_;
|
||||
int64_t hidden_size_;
|
||||
int64_t n_shared_experts_;
|
||||
bool is_gated_;
|
||||
bool has_score_bias_;
|
||||
bool has_bias_;
|
||||
bool skip_bias_add_;
|
||||
int64_t renormalize_;
|
||||
std::string hidden_act_;
|
||||
std::string scoring_func_;
|
||||
bool is_smoothquant_;
|
||||
|
||||
int64_t num_experts_per_rank_;
|
||||
int64_t start_expert_id_;
|
||||
|
||||
ReplicatedLinear gate_{nullptr};
|
||||
DenseMLP shared_experts_{nullptr};
|
||||
torch::nn::Linear shared_expert_gate_{nullptr};
|
||||
QuantArgs quant_args_;
|
||||
ParallelArgs parallel_args_;
|
||||
torch::TensorOptions options_;
|
||||
ProcessGroup* tp_pg_;
|
||||
|
||||
DEFINE_WEIGHT(w13);
|
||||
DEFINE_FUSED_WEIGHT(w1);
|
||||
DEFINE_FUSED_WEIGHT(w3);
|
||||
DEFINE_FUSED_WEIGHT(w2);
|
||||
DEFINE_WEIGHT(e_score_correction_bias);
|
||||
DEFINE_WEIGHT(w13_scale);
|
||||
DEFINE_FUSED_WEIGHT(w1_scale);
|
||||
DEFINE_FUSED_WEIGHT(w3_scale);
|
||||
DEFINE_FUSED_WEIGHT(w2_scale);
|
||||
DEFINE_FUSED_WEIGHT(input_smooth);
|
||||
DEFINE_FUSED_WEIGHT(act_smooth);
|
||||
|
||||
void load_e_score_correction_bias(const StateDict& state_dict);
|
||||
void load_experts(const StateDict& state_dict);
|
||||
};
|
||||
TORCH_MODULE(FusedMoE);
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
@@ -0,0 +1,32 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "qwen3_5_decoder_layer_impl.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
Qwen3_5DecoderLayerImpl::Qwen3_5DecoderLayerImpl(const ModelContext& context,
|
||||
int32_t layer_id)
|
||||
: Qwen3NextDecoderLayerImpl(context,
|
||||
layer_id,
|
||||
std::make_shared<Qwen3_5GatedDeltaNetImpl>(
|
||||
context.get_model_args(),
|
||||
context.get_quant_args(),
|
||||
context.get_parallel_args(),
|
||||
context.get_tensor_options())) {}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
32
ex_engine/xllm_layers/npu_torch/qwen3_5_decoder_layer_impl.h
Normal file
32
ex_engine/xllm_layers/npu_torch/qwen3_5_decoder_layer_impl.h
Normal file
@@ -0,0 +1,32 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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 "layers/npu_torch/qwen3_5_gated_delta_net.h"
|
||||
#include "layers/npu_torch/qwen3_next_decoder_layer_impl.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3_5DecoderLayerImpl : public Qwen3NextDecoderLayerImpl {
|
||||
public:
|
||||
explicit Qwen3_5DecoderLayerImpl(const ModelContext& context,
|
||||
int32_t layer_id);
|
||||
};
|
||||
TORCH_MODULE(Qwen3_5DecoderLayer);
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
219
ex_engine/xllm_layers/npu_torch/qwen3_5_gated_delta_net.cpp
Normal file
219
ex_engine/xllm_layers/npu_torch/qwen3_5_gated_delta_net.cpp
Normal file
@@ -0,0 +1,219 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "qwen3_5_gated_delta_net.h"
|
||||
|
||||
#include <glog/logging.h>
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
Qwen3_5GatedDeltaNetImpl::Qwen3_5GatedDeltaNetImpl(
|
||||
const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options)
|
||||
: Qwen3NextGatedDeltaNetImpl(args,
|
||||
quant_args,
|
||||
parallel_args,
|
||||
options,
|
||||
/*init_projections=*/false) {
|
||||
in_proj_qkv_ = register_module("in_proj_qkv",
|
||||
ColumnParallelLinear(args.hidden_size(),
|
||||
k_size_ * 2 + v_size_,
|
||||
/*bias=*/false,
|
||||
/*gather_output=*/false,
|
||||
quant_args,
|
||||
parallel_args.tp_group_,
|
||||
options));
|
||||
in_proj_z_ = register_module("in_proj_z",
|
||||
ColumnParallelLinear(args.hidden_size(),
|
||||
v_size_,
|
||||
/*bias=*/false,
|
||||
/*gather_output=*/false,
|
||||
quant_args,
|
||||
parallel_args.tp_group_,
|
||||
options));
|
||||
in_proj_b_ = register_module("in_proj_b",
|
||||
ColumnParallelLinear(args.hidden_size(),
|
||||
num_v_heads_,
|
||||
/*bias=*/false,
|
||||
/*gather_output=*/false,
|
||||
quant_args,
|
||||
parallel_args.tp_group_,
|
||||
options));
|
||||
in_proj_a_ = register_module("in_proj_a",
|
||||
ColumnParallelLinear(args.hidden_size(),
|
||||
num_v_heads_,
|
||||
/*bias=*/false,
|
||||
/*gather_output=*/false,
|
||||
quant_args,
|
||||
parallel_args.tp_group_,
|
||||
options));
|
||||
}
|
||||
|
||||
torch::Tensor Qwen3_5GatedDeltaNetImpl::merge_qkvz_from_split_activations(
|
||||
const torch::Tensor& qkv,
|
||||
const torch::Tensor& z) const {
|
||||
CHECK_EQ(qkv.dim(), 3) << "Expected qkv activation to be 3D, got "
|
||||
<< qkv.sizes();
|
||||
CHECK_EQ(z.dim(), 3) << "Expected z activation to be 3D, got " << z.sizes();
|
||||
CHECK_EQ(qkv.size(0), z.size(0)) << "qkv/z batch size mismatch.";
|
||||
CHECK_EQ(qkv.size(1), z.size(1)) << "qkv/z sequence size mismatch.";
|
||||
CHECK_EQ(qkv.size(2), (2 * k_size_ + v_size_) / tp_size_)
|
||||
<< "Unexpected qkv hidden size for Qwen3.5.";
|
||||
CHECK_EQ(z.size(2), v_size_ / tp_size_)
|
||||
<< "Unexpected z hidden size for Qwen3.5.";
|
||||
CHECK_GT(num_k_heads_, 0) << "linear_num_key_heads must be positive.";
|
||||
CHECK_EQ(num_v_heads_ % num_k_heads_, 0)
|
||||
<< "linear_num_value_heads must be divisible by linear_num_key_heads.";
|
||||
|
||||
const int64_t bs = qkv.size(0);
|
||||
const int64_t seqlen = qkv.size(1);
|
||||
const int64_t local_k_heads = num_k_heads_ / tp_size_;
|
||||
const int64_t local_v_heads = num_v_heads_ / tp_size_;
|
||||
const int64_t num_v_heads_per_k = num_v_heads_ / num_k_heads_;
|
||||
|
||||
auto qkv_split = torch::split(
|
||||
qkv, {k_size_ / tp_size_, k_size_ / tp_size_, v_size_ / tp_size_}, 2);
|
||||
auto q = qkv_split[0].view({bs, seqlen, local_k_heads, head_k_dim_});
|
||||
auto k = qkv_split[1].view({bs, seqlen, local_k_heads, head_k_dim_});
|
||||
auto v = qkv_split[2].view({bs, seqlen, local_v_heads, head_v_dim_});
|
||||
auto z_view = z.view({bs, seqlen, local_v_heads, head_v_dim_});
|
||||
|
||||
v = v.view({bs, seqlen, local_k_heads, num_v_heads_per_k * head_v_dim_});
|
||||
z_view =
|
||||
z_view.view({bs, seqlen, local_k_heads, num_v_heads_per_k * head_v_dim_});
|
||||
|
||||
return torch::cat({q, k, v, z_view}, -1).view({bs, seqlen, -1}).contiguous();
|
||||
}
|
||||
|
||||
torch::Tensor Qwen3_5GatedDeltaNetImpl::merge_ba_from_split_activations(
|
||||
const torch::Tensor& b,
|
||||
const torch::Tensor& a) const {
|
||||
CHECK_EQ(b.dim(), 3) << "Expected b activation to be 3D, got " << b.sizes();
|
||||
CHECK_EQ(a.dim(), 3) << "Expected a activation to be 3D, got " << a.sizes();
|
||||
CHECK_EQ(b.size(0), a.size(0)) << "b/a batch size mismatch.";
|
||||
CHECK_EQ(b.size(1), a.size(1)) << "b/a sequence size mismatch.";
|
||||
CHECK_EQ(b.size(2), num_v_heads_ / tp_size_)
|
||||
<< "Unexpected b hidden size for Qwen3.5.";
|
||||
CHECK_EQ(a.size(2), num_v_heads_ / tp_size_)
|
||||
<< "Unexpected a hidden size for Qwen3.5.";
|
||||
CHECK_GT(num_k_heads_, 0) << "linear_num_key_heads must be positive.";
|
||||
CHECK_EQ(num_v_heads_ % num_k_heads_, 0)
|
||||
<< "linear_num_value_heads must be divisible by linear_num_key_heads.";
|
||||
|
||||
const int64_t bs = b.size(0);
|
||||
const int64_t seqlen = b.size(1);
|
||||
const int64_t local_k_heads = num_k_heads_ / tp_size_;
|
||||
const int64_t num_v_heads_per_k = num_v_heads_ / num_k_heads_;
|
||||
|
||||
auto b_view = b.view({bs, seqlen, local_k_heads, num_v_heads_per_k});
|
||||
auto a_view = a.view({bs, seqlen, local_k_heads, num_v_heads_per_k});
|
||||
return torch::cat({b_view, a_view}, -1).view({bs, seqlen, -1}).contiguous();
|
||||
}
|
||||
|
||||
std::pair<torch::Tensor, torch::Tensor>
|
||||
Qwen3_5GatedDeltaNetImpl::project_decode_inputs(
|
||||
const torch::Tensor& hidden_states) {
|
||||
const auto reshape_projection = [](const torch::Tensor& projection) {
|
||||
return projection.view({projection.size(0), -1, projection.size(-1)});
|
||||
};
|
||||
auto qkv = reshape_projection(in_proj_qkv_->forward(hidden_states));
|
||||
auto z_proj = reshape_projection(in_proj_z_->forward(hidden_states));
|
||||
auto b_proj = reshape_projection(in_proj_b_->forward(hidden_states));
|
||||
auto a_proj = reshape_projection(in_proj_a_->forward(hidden_states));
|
||||
return {merge_qkvz_from_split_activations(qkv, z_proj),
|
||||
merge_ba_from_split_activations(b_proj, a_proj)};
|
||||
}
|
||||
|
||||
std::pair<torch::Tensor, torch::Tensor>
|
||||
Qwen3_5GatedDeltaNetImpl::project_flat_inputs(
|
||||
const torch::Tensor& hidden_states) {
|
||||
auto qkv = in_proj_qkv_->forward(hidden_states).unsqueeze(0);
|
||||
auto z_proj = in_proj_z_->forward(hidden_states).unsqueeze(0);
|
||||
auto b_proj = in_proj_b_->forward(hidden_states).unsqueeze(0);
|
||||
auto a_proj = in_proj_a_->forward(hidden_states).unsqueeze(0);
|
||||
auto qkvz = merge_qkvz_from_split_activations(qkv, z_proj);
|
||||
auto ba = merge_ba_from_split_activations(b_proj, a_proj);
|
||||
return {qkvz.view({hidden_states.size(0), qkvz.size(-1)}).contiguous(),
|
||||
ba.view({hidden_states.size(0), ba.size(-1)}).contiguous()};
|
||||
}
|
||||
|
||||
std::optional<
|
||||
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>>
|
||||
Qwen3_5GatedDeltaNetImpl::project_split_inputs(
|
||||
const torch::Tensor& hidden_states,
|
||||
const AttentionMetadata& attn_metadata) {
|
||||
auto qkv = reshape_projected_tokens_with_pad(
|
||||
attn_metadata, in_proj_qkv_->forward(hidden_states));
|
||||
auto z_proj = reshape_projected_tokens_with_pad(
|
||||
attn_metadata, in_proj_z_->forward(hidden_states));
|
||||
auto b_proj = reshape_projected_tokens_with_pad(
|
||||
attn_metadata, in_proj_b_->forward(hidden_states));
|
||||
auto a_proj = reshape_projected_tokens_with_pad(
|
||||
attn_metadata, in_proj_a_->forward(hidden_states));
|
||||
|
||||
const int64_t batch_size = qkv.size(0);
|
||||
const int64_t seq_len = qkv.size(1);
|
||||
auto z =
|
||||
z_proj.view({batch_size, seq_len, num_v_heads_ / tp_size_, head_v_dim_});
|
||||
auto b = b_proj.view({batch_size, seq_len, num_v_heads_ / tp_size_});
|
||||
auto a = a_proj.view({batch_size, seq_len, num_v_heads_ / tp_size_});
|
||||
return std::make_tuple(qkv, z, b, a);
|
||||
}
|
||||
|
||||
void Qwen3_5GatedDeltaNetImpl::load_projection_state_dict(
|
||||
const StateDict& state_dict) {
|
||||
auto in_proj_qkv_state_dict = state_dict.get_dict_with_prefix("in_proj_qkv.");
|
||||
if (in_proj_qkv_state_dict.size() > 0 && !in_proj_qkv_->is_weight_loaded()) {
|
||||
in_proj_qkv_->load_state_dict(
|
||||
in_proj_qkv_state_dict,
|
||||
/*shard_tensor_count=*/3,
|
||||
/*shard_sizes=*/
|
||||
{k_size_ / tp_size_, k_size_ / tp_size_, v_size_ / tp_size_});
|
||||
}
|
||||
|
||||
auto in_proj_z_state_dict = state_dict.get_dict_with_prefix("in_proj_z.");
|
||||
if (in_proj_z_state_dict.size() > 0 && !in_proj_z_->is_weight_loaded()) {
|
||||
in_proj_z_->load_state_dict(in_proj_z_state_dict);
|
||||
}
|
||||
|
||||
auto in_proj_b_state_dict = state_dict.get_dict_with_prefix("in_proj_b.");
|
||||
if (in_proj_b_state_dict.size() > 0 && !in_proj_b_->is_weight_loaded()) {
|
||||
in_proj_b_->load_state_dict(in_proj_b_state_dict);
|
||||
}
|
||||
|
||||
auto in_proj_a_state_dict = state_dict.get_dict_with_prefix("in_proj_a.");
|
||||
if (in_proj_a_state_dict.size() > 0 && !in_proj_a_->is_weight_loaded()) {
|
||||
in_proj_a_->load_state_dict(in_proj_a_state_dict);
|
||||
}
|
||||
}
|
||||
|
||||
void Qwen3_5GatedDeltaNetImpl::verify_projection_weights(
|
||||
const std::string& prefix) const {
|
||||
CHECK(in_proj_qkv_ && in_proj_qkv_->is_weight_loaded())
|
||||
<< "Missing required weight after all shards loaded: " << prefix
|
||||
<< "in_proj_qkv.weight";
|
||||
CHECK(in_proj_z_ && in_proj_z_->is_weight_loaded())
|
||||
<< "Missing required weight after all shards loaded: " << prefix
|
||||
<< "in_proj_z.weight";
|
||||
CHECK(in_proj_b_ && in_proj_b_->is_weight_loaded())
|
||||
<< "Missing required weight after all shards loaded: " << prefix
|
||||
<< "in_proj_b.weight";
|
||||
CHECK(in_proj_a_ && in_proj_a_->is_weight_loaded())
|
||||
<< "Missing required weight after all shards loaded: " << prefix
|
||||
<< "in_proj_a.weight";
|
||||
}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
66
ex_engine/xllm_layers/npu_torch/qwen3_5_gated_delta_net.h
Normal file
66
ex_engine/xllm_layers/npu_torch/qwen3_5_gated_delta_net.h
Normal file
@@ -0,0 +1,66 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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 <optional>
|
||||
#include <string>
|
||||
#include <tuple>
|
||||
#include <utility>
|
||||
|
||||
#include "qwen3_next_gated_delta_net.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3_5GatedDeltaNetImpl : public Qwen3NextGatedDeltaNetImpl {
|
||||
public:
|
||||
Qwen3_5GatedDeltaNetImpl() = default;
|
||||
Qwen3_5GatedDeltaNetImpl(const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options);
|
||||
|
||||
protected:
|
||||
std::pair<torch::Tensor, torch::Tensor> project_decode_inputs(
|
||||
const torch::Tensor& hidden_states) override;
|
||||
std::pair<torch::Tensor, torch::Tensor> project_flat_inputs(
|
||||
const torch::Tensor& hidden_states) override;
|
||||
std::optional<
|
||||
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>>
|
||||
project_split_inputs(const torch::Tensor& hidden_states,
|
||||
const AttentionMetadata& attn_metadata) override;
|
||||
bool use_fla_ssm_state_layout() const override { return true; }
|
||||
|
||||
void load_projection_state_dict(const StateDict& state_dict) override;
|
||||
void verify_projection_weights(const std::string& prefix) const override;
|
||||
|
||||
private:
|
||||
torch::Tensor merge_qkvz_from_split_activations(const torch::Tensor& qkv,
|
||||
const torch::Tensor& z) const;
|
||||
torch::Tensor merge_ba_from_split_activations(const torch::Tensor& b,
|
||||
const torch::Tensor& a) const;
|
||||
|
||||
ColumnParallelLinear in_proj_qkv_{nullptr};
|
||||
ColumnParallelLinear in_proj_z_{nullptr};
|
||||
ColumnParallelLinear in_proj_b_{nullptr};
|
||||
ColumnParallelLinear in_proj_a_{nullptr};
|
||||
};
|
||||
TORCH_MODULE(Qwen3_5GatedDeltaNet);
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
1164
ex_engine/xllm_layers/npu_torch/qwen3_gated_delta_net_base.cpp
Normal file
1164
ex_engine/xllm_layers/npu_torch/qwen3_gated_delta_net_base.cpp
Normal file
File diff suppressed because it is too large
Load Diff
112
ex_engine/xllm_layers/npu_torch/qwen3_gated_delta_net_base.h
Normal file
112
ex_engine/xllm_layers/npu_torch/qwen3_gated_delta_net_base.h
Normal file
@@ -0,0 +1,112 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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 <optional>
|
||||
#include <string>
|
||||
#include <tuple>
|
||||
#include <utility>
|
||||
|
||||
#include "attention.h"
|
||||
#include "framework/kv_cache/kv_cache.h"
|
||||
#include "framework/model/model_args.h"
|
||||
#include "framework/parallel_state/parallel_args.h"
|
||||
#include "framework/quant_args.h"
|
||||
#include "framework/state_dict/state_dict.h"
|
||||
#include "framework/state_dict/utils.h"
|
||||
#include "layers/common/linear.h"
|
||||
#include "layers/common/rms_norm_gated.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3GatedDeltaNetBaseImpl : public torch::nn::Module {
|
||||
public:
|
||||
Qwen3GatedDeltaNetBaseImpl() = default;
|
||||
Qwen3GatedDeltaNetBaseImpl(const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options);
|
||||
|
||||
virtual void load_state_dict(const StateDict& state_dict) = 0;
|
||||
virtual void verify_loaded_weights(const std::string& prefix) const = 0;
|
||||
|
||||
torch::Tensor forward(const torch::Tensor& hidden_states,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache,
|
||||
const ModelInputParams& input_params);
|
||||
|
||||
protected:
|
||||
virtual std::pair<torch::Tensor, torch::Tensor> project_decode_inputs(
|
||||
const torch::Tensor& hidden_states) = 0;
|
||||
virtual std::pair<torch::Tensor, torch::Tensor> project_flat_inputs(
|
||||
const torch::Tensor& hidden_states) = 0;
|
||||
// Qwen3.5 overrides this to project and reshape its separate qkv/z/b/a
|
||||
// weights in every forward mode. Qwen3Next keeps qkvz/ba packed and returns
|
||||
// nullopt to select the fused-split fallback.
|
||||
virtual std::optional<
|
||||
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>>
|
||||
project_split_inputs(const torch::Tensor& hidden_states,
|
||||
const AttentionMetadata& attn_metadata) {
|
||||
return std::nullopt;
|
||||
}
|
||||
virtual bool use_fla_ssm_state_layout() const { return false; }
|
||||
|
||||
void load_common_state_dict(const StateDict& state_dict);
|
||||
void verify_common_loaded_weights(const std::string& prefix) const;
|
||||
|
||||
torch::Tensor get_linear_state_indices(const ModelInputParams& input_params,
|
||||
const torch::Device& device) const;
|
||||
|
||||
std::pair<torch::Tensor, torch::Tensor> project_padded_inputs(
|
||||
const torch::Tensor& hidden_states,
|
||||
const AttentionMetadata& attn_metadata);
|
||||
|
||||
torch::Tensor reshape_qkvz_unpad(const AttentionMetadata& attn_metadata,
|
||||
const torch::Tensor& padded_qkvz) const;
|
||||
|
||||
// Projection outputs are packed as [total_tokens, dim], while GDN kernels
|
||||
// consume dense [batch, max_query_len, dim] tensors. Split the packed tokens
|
||||
// by query length and pad each sequence before entering the kernels.
|
||||
torch::Tensor reshape_projected_tokens_with_pad(
|
||||
const AttentionMetadata& attn_metadata,
|
||||
const torch::Tensor& projected_tokens) const;
|
||||
|
||||
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> process_mixed_qkv(
|
||||
torch::Tensor& mixed_qkv) const;
|
||||
|
||||
int64_t num_k_heads_ = 0;
|
||||
int64_t num_v_heads_ = 0;
|
||||
int64_t head_k_dim_ = 0;
|
||||
int64_t head_v_dim_ = 0;
|
||||
int64_t k_size_ = 0;
|
||||
int64_t v_size_ = 0;
|
||||
int64_t tp_size_ = 1;
|
||||
int64_t rank_ = 0;
|
||||
int32_t conv_kernel_size_ = 0;
|
||||
|
||||
ColumnParallelLinear conv1d_{nullptr};
|
||||
RowParallelLinear o_proj_{nullptr};
|
||||
RmsNormGated norm_{nullptr};
|
||||
|
||||
DEFINE_WEIGHT(dt_bias);
|
||||
DEFINE_WEIGHT(A_log);
|
||||
};
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
291
ex_engine/xllm_layers/npu_torch/qwen3_next_attention.cpp
Normal file
291
ex_engine/xllm_layers/npu_torch/qwen3_next_attention.cpp
Normal file
@@ -0,0 +1,291 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "qwen3_next_attention.h"
|
||||
|
||||
#include <glog/logging.h>
|
||||
|
||||
#include <tuple>
|
||||
#include <vector>
|
||||
|
||||
#include "common/flash_comm1_context.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
Qwen3NextAttentionImpl::Qwen3NextAttentionImpl(
|
||||
const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options,
|
||||
int32_t layer_id) {
|
||||
const int64_t tp_size = parallel_args.tp_group_->world_size();
|
||||
const int64_t total_num_heads = args.n_heads();
|
||||
const int64_t total_num_kv_heads = args.n_kv_heads().value_or(args.n_heads());
|
||||
layer_id_ = layer_id;
|
||||
rank_ = parallel_args.tp_group_->rank();
|
||||
CHECK(total_num_heads % tp_size == 0);
|
||||
num_heads_ = total_num_heads / tp_size;
|
||||
|
||||
if (total_num_kv_heads >= tp_size) {
|
||||
CHECK(total_num_kv_heads % tp_size == 0);
|
||||
num_kv_heads_ = total_num_kv_heads / tp_size;
|
||||
num_kv_head_replicas_ = 1;
|
||||
} else {
|
||||
CHECK(tp_size % total_num_kv_heads == 0);
|
||||
num_kv_heads_ = 1;
|
||||
num_kv_head_replicas_ = tp_size / total_num_kv_heads;
|
||||
}
|
||||
|
||||
head_dim_ = args.head_dim();
|
||||
q_size_ = num_heads_ * head_dim_;
|
||||
kv_size_ = num_kv_heads_ * head_dim_;
|
||||
scaling_ = 1.0f / std::sqrt(static_cast<float>(head_dim_));
|
||||
attn_output_gate_ = args.attn_output_gate();
|
||||
// 1. QKV linear
|
||||
qkv_proj_ = register_module(
|
||||
"qkv_proj",
|
||||
QKVParallelLinear(args.hidden_size(),
|
||||
attn_output_gate_ ? num_heads_ * 2 : num_heads_,
|
||||
num_kv_heads_,
|
||||
args.head_dim(),
|
||||
num_kv_head_replicas_,
|
||||
/*bias=*/args.attention_bias(),
|
||||
/*gather_output=*/false,
|
||||
parallel_args,
|
||||
options,
|
||||
quant_args));
|
||||
|
||||
// 2. O proj
|
||||
o_proj_ = register_module("o_proj",
|
||||
RowParallelLinear(total_num_heads * head_dim_,
|
||||
args.hidden_size(),
|
||||
/*bias=*/false,
|
||||
/*input_is_parallelized=*/true,
|
||||
/*if_reduce_results=*/true,
|
||||
quant_args,
|
||||
parallel_args.tp_group_,
|
||||
options));
|
||||
|
||||
// 3. Q norm
|
||||
q_norm_ = register_module(
|
||||
"q_norm", Qwen3NextRMSNorm(head_dim_, args.rms_norm_eps(), options));
|
||||
|
||||
// 4. K norm
|
||||
k_norm_ = register_module(
|
||||
"k_norm", Qwen3NextRMSNorm(head_dim_, args.rms_norm_eps(), options));
|
||||
|
||||
// 5. Rotary embedding
|
||||
const int rotary_dim =
|
||||
static_cast<int>(head_dim_ * args.partial_rotary_factor());
|
||||
rotary_emb_ =
|
||||
register_module("rotary_emb",
|
||||
PartialRotaryEmbedding(rotary_dim,
|
||||
args.max_position_embeddings(),
|
||||
args.rope_theta(),
|
||||
head_dim_,
|
||||
true,
|
||||
false,
|
||||
options));
|
||||
|
||||
// 6. Attention
|
||||
attn_ = register_module("attn",
|
||||
Attention(num_heads_,
|
||||
head_dim_,
|
||||
scaling_,
|
||||
num_kv_heads_,
|
||||
args.sliding_window()));
|
||||
|
||||
// 7. Fused split_qkv_rmsnorm_mrope kernel setup
|
||||
rotary_dim_ = static_cast<int64_t>(head_dim_ * args.partial_rotary_factor());
|
||||
rms_norm_eps_ = args.rms_norm_eps();
|
||||
mrope_section_ = args.rope_scaling_mrope_section();
|
||||
is_interleaved_ = args.rope_scaling_mrope_interleaved();
|
||||
use_fused_qkv_ = false;
|
||||
if (attn_output_gate_ && !mrope_section_.empty() &&
|
||||
mrope_section_.size() == 3 && rotary_dim_ > 0 &&
|
||||
xllm::kernel::has_split_qkv_rmsnorm_mrope_specialization(
|
||||
num_heads_, num_kv_heads_, head_dim_)) {
|
||||
mrope_gather_pattern_ =
|
||||
xllm::kernel::build_split_qkv_rmsnorm_mrope_gather_pattern(
|
||||
rotary_dim_, mrope_section_, is_interleaved_, options.device());
|
||||
use_fused_qkv_ = true;
|
||||
LOG(INFO) << "Qwen3NextAttention layer " << layer_id_
|
||||
<< ": using fused split_qkv_rmsnorm_mrope kernel";
|
||||
}
|
||||
}
|
||||
|
||||
torch::Tensor Qwen3NextAttentionImpl::build_mrope_cos_sin(
|
||||
const torch::Tensor& positions) const {
|
||||
auto cos_sin_cache = rotary_emb_->get_cos_sin_cache();
|
||||
if (positions.dim() == 1) {
|
||||
return cos_sin_cache.index_select(0, positions).repeat({1, 3});
|
||||
}
|
||||
// positions is [3, T] for mRoPE (graph mode or VL)
|
||||
// transpose from [3, T] to [T, 3]
|
||||
auto positions_t = positions.permute({1, 0}).contiguous();
|
||||
auto gathered = cos_sin_cache.index_select(0, positions_t.view({-1}));
|
||||
// [T, 3, rope_dim]
|
||||
return gathered.view({positions.size(1), -1});
|
||||
}
|
||||
|
||||
torch::Tensor Qwen3NextAttentionImpl::forward(
|
||||
const torch::Tensor& positions,
|
||||
const torch::Tensor& hidden_states,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache,
|
||||
const torch::Tensor& mrope_cos_sin) {
|
||||
const FlashComm1Context* fc1_ctx = get_current_flash_comm1_context();
|
||||
torch::Tensor h = hidden_states;
|
||||
|
||||
if (fc1_ctx && is_sequence_sharded(*fc1_ctx)) {
|
||||
h = gather_sequence(hidden_states, *fc1_ctx);
|
||||
}
|
||||
|
||||
auto qkv = qkv_proj_->forward(h);
|
||||
|
||||
if (use_fused_qkv_) {
|
||||
const int64_t T = qkv.size(0);
|
||||
xllm::kernel::SplitQkvRmsnormMropeParams params;
|
||||
params.qkvg = qkv;
|
||||
params.q_weight = q_norm_->weight();
|
||||
params.k_weight = k_norm_->weight();
|
||||
params.cos_sin = mrope_cos_sin;
|
||||
params.gather_pattern = mrope_gather_pattern_;
|
||||
params.eps = rms_norm_eps_;
|
||||
params.num_q_heads = num_heads_;
|
||||
params.num_kv_heads = num_kv_heads_;
|
||||
params.head_size = head_dim_;
|
||||
|
||||
auto [q, k, v, gate] = xllm::kernel::split_qkv_rmsnorm_mrope(params);
|
||||
|
||||
auto q_flat = q.view({T, q_size_});
|
||||
auto k_flat = k.view({T, kv_size_});
|
||||
auto v_flat = v.view({T, kv_size_});
|
||||
|
||||
auto out = std::get<0>(
|
||||
attn_->forward(attn_metadata, q_flat, k_flat, v_flat, kv_cache));
|
||||
out = out * torch::sigmoid(gate.view({T, q_size_}));
|
||||
|
||||
if (fc1_ctx && is_sequence_sharded(*fc1_ctx)) {
|
||||
return o_proj_->forward(out, row_parallel_reduce_mode_for_fc1(*fc1_ctx));
|
||||
}
|
||||
return o_proj_->forward(out);
|
||||
}
|
||||
|
||||
// Fallback path: weight-reordered layout [Q | G | K | V]
|
||||
torch::Tensor q, k, v;
|
||||
torch::Tensor gate;
|
||||
|
||||
if (attn_output_gate_) {
|
||||
q = qkv.slice(-1, 0, q_size_);
|
||||
gate = qkv.slice(-1, q_size_, q_size_ * 2);
|
||||
k = qkv.slice(-1, q_size_ * 2, q_size_ * 2 + kv_size_);
|
||||
v = qkv.slice(-1, q_size_ * 2 + kv_size_, q_size_ * 2 + kv_size_ * 2);
|
||||
} else {
|
||||
q = qkv.slice(-1, 0, q_size_);
|
||||
k = qkv.slice(-1, q_size_, q_size_ + kv_size_);
|
||||
v = qkv.slice(-1, q_size_ + kv_size_, q_size_ + 2 * kv_size_);
|
||||
}
|
||||
|
||||
const int64_t T = q.size(0);
|
||||
auto q_3d = q.view({T, num_heads_, head_dim_});
|
||||
q = std::get<0>(q_norm_->forward(q_3d)).view({T, q_size_});
|
||||
auto k_3d = k.view({T, num_kv_heads_, head_dim_});
|
||||
k = std::get<0>(k_norm_->forward(k_3d)).view({T, kv_size_});
|
||||
|
||||
rotary_emb_->forward(positions, q, k);
|
||||
auto out = std::get<0>(attn_->forward(attn_metadata, q, k, v, kv_cache));
|
||||
|
||||
if (attn_output_gate_) {
|
||||
out = out * torch::sigmoid(gate);
|
||||
}
|
||||
|
||||
if (fc1_ctx && is_sequence_sharded(*fc1_ctx)) {
|
||||
return o_proj_->forward(out, row_parallel_reduce_mode_for_fc1(*fc1_ctx));
|
||||
}
|
||||
return o_proj_->forward(out);
|
||||
}
|
||||
|
||||
void Qwen3NextAttentionImpl::load_state_dict(const StateDict& state_dict) {
|
||||
qkv_proj_->load_state_dict(state_dict, {"q_proj.", "k_proj.", "v_proj."});
|
||||
|
||||
if (attn_output_gate_ && qkv_proj_->is_weight_loaded() &&
|
||||
!qkv_weight_reordered_) {
|
||||
// Rearrange q_proj rows from per-head interleaved [q0,g0,q1,g1,...]
|
||||
// to grouped [q0,q1,...,g0,g1,...] so forward output is [Q|G|K|V].
|
||||
auto w = qkv_proj_->weight();
|
||||
auto qg_rows = w.slice(0, 0, q_size_ * 2);
|
||||
const int64_t hidden = w.size(1);
|
||||
auto qg_3d = qg_rows.view({num_heads_, 2 * head_dim_, hidden});
|
||||
auto q_part = qg_3d.slice(1, 0, head_dim_);
|
||||
auto g_part = qg_3d.slice(1, head_dim_, 2 * head_dim_);
|
||||
auto reordered = torch::cat(
|
||||
{q_part.reshape({q_size_, hidden}), g_part.reshape({q_size_, hidden})},
|
||||
0);
|
||||
qg_rows.copy_(reordered);
|
||||
|
||||
// Reorder weight_scale and weight_offset for W8A8 dynamic quantization.
|
||||
// These are per-channel (per output row) tensors that must match the
|
||||
// reordered weight layout for correct dequantization.
|
||||
const int64_t qg_size = q_size_ * 2;
|
||||
auto reorder_per_channel = [this, qg_size](torch::Tensor tensor) {
|
||||
if (!tensor.defined() || tensor.numel() == 0) {
|
||||
return;
|
||||
}
|
||||
auto qg_part = tensor.slice(0, 0, qg_size);
|
||||
auto qg_2d = qg_part.view({num_heads_, 2 * head_dim_});
|
||||
auto q_scale = qg_2d.slice(1, 0, head_dim_);
|
||||
auto g_scale = qg_2d.slice(1, head_dim_, 2 * head_dim_);
|
||||
auto reordered_scale = torch::cat(
|
||||
{q_scale.reshape({q_size_}), g_scale.reshape({q_size_})}, 0);
|
||||
qg_part.copy_(reordered_scale);
|
||||
};
|
||||
|
||||
if (qkv_proj_->is_weight_scale_loaded()) {
|
||||
reorder_per_channel(qkv_proj_->weight_scale());
|
||||
}
|
||||
if (qkv_proj_->is_weight_offset_loaded()) {
|
||||
reorder_per_channel(qkv_proj_->weight_offset());
|
||||
}
|
||||
|
||||
qkv_weight_reordered_ = true;
|
||||
}
|
||||
|
||||
o_proj_->load_state_dict(state_dict.get_dict_with_prefix("o_proj."));
|
||||
if (auto w = state_dict.get_tensor("q_norm.weight"); w.defined()) {
|
||||
q_norm_->load_state_dict(StateDict({{"weight", w}}));
|
||||
}
|
||||
if (auto w = state_dict.get_tensor("k_norm.weight"); w.defined()) {
|
||||
k_norm_->load_state_dict(StateDict({{"weight", w}}));
|
||||
}
|
||||
|
||||
// Gemma RMSNorm uses (1 + w) as the scale factor, but the fused kernel
|
||||
// uses standard RMSNorm (w only). Pre-add 1 so the fused kernel produces
|
||||
// the same result as Qwen3NextRMSNorm (gemma_rms_norm).
|
||||
if (use_fused_qkv_) {
|
||||
if (q_norm_->is_weight_loaded() && !q_norm_weight_adjusted_) {
|
||||
q_norm_->weight().add_(1.0);
|
||||
q_norm_weight_adjusted_ = true;
|
||||
}
|
||||
if (k_norm_->is_weight_loaded() && !k_norm_weight_adjusted_) {
|
||||
k_norm_->weight().add_(1.0);
|
||||
k_norm_weight_adjusted_ = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
88
ex_engine/xllm_layers/npu_torch/qwen3_next_attention.h
Normal file
88
ex_engine/xllm_layers/npu_torch/qwen3_next_attention.h
Normal file
@@ -0,0 +1,88 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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 <vector>
|
||||
|
||||
#include "attention.h"
|
||||
#include "framework/kv_cache/kv_cache.h"
|
||||
#include "framework/model/model_args.h"
|
||||
#include "framework/parallel_state/parallel_args.h"
|
||||
#include "framework/quant_args.h"
|
||||
#include "framework/state_dict/state_dict.h"
|
||||
#include "kernels/ops_api.h"
|
||||
#include "layers/common/linear.h"
|
||||
#include "layers/common/partial_rotary_embedding.h"
|
||||
#include "layers/common/qwen3_next_rms_norm.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3NextAttentionImpl : public torch::nn::Module {
|
||||
public:
|
||||
Qwen3NextAttentionImpl() = default;
|
||||
Qwen3NextAttentionImpl(const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options,
|
||||
int32_t layer_id);
|
||||
|
||||
torch::Tensor forward(const torch::Tensor& positions,
|
||||
const torch::Tensor& hidden_states,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache,
|
||||
const torch::Tensor& mrope_cos_sin);
|
||||
|
||||
torch::Tensor build_mrope_cos_sin(const torch::Tensor& positions) const;
|
||||
|
||||
void load_state_dict(const StateDict& state_dict);
|
||||
|
||||
private:
|
||||
int64_t num_heads_;
|
||||
int64_t num_kv_heads_;
|
||||
int64_t num_kv_head_replicas_;
|
||||
int64_t head_dim_;
|
||||
int64_t q_size_;
|
||||
int64_t kv_size_;
|
||||
float scaling_;
|
||||
bool attn_output_gate_;
|
||||
int32_t layer_id_;
|
||||
int32_t rank_;
|
||||
int64_t rotary_dim_;
|
||||
float rms_norm_eps_;
|
||||
bool use_fused_qkv_;
|
||||
bool is_interleaved_;
|
||||
bool qkv_weight_reordered_ = false;
|
||||
bool q_norm_weight_adjusted_ = false;
|
||||
bool k_norm_weight_adjusted_ = false;
|
||||
std::vector<int64_t> mrope_section_;
|
||||
torch::Tensor mrope_gather_pattern_;
|
||||
|
||||
QKVParallelLinear qkv_proj_{nullptr};
|
||||
RowParallelLinear o_proj_{nullptr};
|
||||
|
||||
Qwen3NextRMSNorm q_norm_{nullptr};
|
||||
Qwen3NextRMSNorm k_norm_{nullptr};
|
||||
|
||||
Attention attn_{nullptr};
|
||||
PartialRotaryEmbedding rotary_emb_{nullptr};
|
||||
};
|
||||
TORCH_MODULE(Qwen3NextAttention);
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
@@ -0,0 +1,41 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "qwen3_next_decoder_layer_impl.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
Qwen3NextDecoderLayerImpl::Qwen3NextDecoderLayerImpl(
|
||||
const ModelContext& context,
|
||||
int32_t layer_id)
|
||||
: Qwen3NextDecoderLayerImpl(context,
|
||||
layer_id,
|
||||
std::make_shared<Qwen3NextGatedDeltaNetImpl>(
|
||||
context.get_model_args(),
|
||||
context.get_quant_args(),
|
||||
context.get_parallel_args(),
|
||||
context.get_tensor_options())) {}
|
||||
|
||||
Qwen3NextDecoderLayerImpl::Qwen3NextDecoderLayerImpl(
|
||||
const ModelContext& context,
|
||||
int32_t layer_id,
|
||||
std::shared_ptr<Qwen3GatedDeltaNetBaseImpl> linear_attention_module)
|
||||
: Qwen3HybridDecoderLayerImplBase(context,
|
||||
layer_id,
|
||||
std::move(linear_attention_module)) {}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
@@ -0,0 +1,38 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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 "layers/npu_torch/qwen3_next_gated_delta_net.h"
|
||||
#include "layers/npu_torch/qwen3_next_hybrid_decoder_layer_base.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3NextDecoderLayerImpl : public Qwen3HybridDecoderLayerImplBase {
|
||||
public:
|
||||
explicit Qwen3NextDecoderLayerImpl(const ModelContext& context,
|
||||
int32_t layer_id);
|
||||
|
||||
protected:
|
||||
Qwen3NextDecoderLayerImpl(
|
||||
const ModelContext& context,
|
||||
int32_t layer_id,
|
||||
std::shared_ptr<Qwen3GatedDeltaNetBaseImpl> linear_attention_module);
|
||||
};
|
||||
TORCH_MODULE(Qwen3NextDecoderLayer);
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
118
ex_engine/xllm_layers/npu_torch/qwen3_next_gated_delta_net.cpp
Normal file
118
ex_engine/xllm_layers/npu_torch/qwen3_next_gated_delta_net.cpp
Normal file
@@ -0,0 +1,118 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "qwen3_next_gated_delta_net.h"
|
||||
|
||||
#include <glog/logging.h>
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
Qwen3NextGatedDeltaNetImpl::Qwen3NextGatedDeltaNetImpl(
|
||||
const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options)
|
||||
: Qwen3NextGatedDeltaNetImpl(args,
|
||||
quant_args,
|
||||
parallel_args,
|
||||
options,
|
||||
/*init_projections=*/true) {}
|
||||
|
||||
Qwen3NextGatedDeltaNetImpl::Qwen3NextGatedDeltaNetImpl(
|
||||
const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options,
|
||||
bool init_projections)
|
||||
: Qwen3GatedDeltaNetBaseImpl(args, quant_args, parallel_args, options) {
|
||||
if (init_projections) {
|
||||
init_next_projections(args, quant_args, parallel_args, options);
|
||||
}
|
||||
}
|
||||
|
||||
void Qwen3NextGatedDeltaNetImpl::init_next_projections(
|
||||
const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options) {
|
||||
// QKVZ projection used by Qwen3-Next linear attention.
|
||||
qkvz_proj_ = register_module("in_proj_qkvz",
|
||||
ColumnParallelLinear(args.hidden_size(),
|
||||
k_size_ * 2 + v_size_ * 2,
|
||||
/*bias=*/false,
|
||||
/*gather_output=*/false,
|
||||
quant_args,
|
||||
parallel_args.tp_group_,
|
||||
options));
|
||||
// BA projection used to derive gating and beta terms.
|
||||
ba_proj_ = register_module("in_proj_ba",
|
||||
ColumnParallelLinear(args.hidden_size(),
|
||||
num_v_heads_ * 2,
|
||||
/*bias=*/false,
|
||||
/*gather_output=*/false,
|
||||
quant_args,
|
||||
parallel_args.tp_group_,
|
||||
options));
|
||||
}
|
||||
|
||||
std::pair<torch::Tensor, torch::Tensor>
|
||||
Qwen3NextGatedDeltaNetImpl::project_decode_inputs(
|
||||
const torch::Tensor& hidden_states) {
|
||||
auto qkvz = qkvz_proj_->forward(hidden_states);
|
||||
auto ba = ba_proj_->forward(hidden_states);
|
||||
return {qkvz.view({qkvz.size(0), -1, qkvz.size(-1)}),
|
||||
ba.view({ba.size(0), -1, ba.size(-1)})};
|
||||
}
|
||||
|
||||
std::pair<torch::Tensor, torch::Tensor>
|
||||
Qwen3NextGatedDeltaNetImpl::project_flat_inputs(
|
||||
const torch::Tensor& hidden_states) {
|
||||
return {qkvz_proj_->forward(hidden_states), ba_proj_->forward(hidden_states)};
|
||||
}
|
||||
|
||||
void Qwen3NextGatedDeltaNetImpl::load_state_dict(const StateDict& state_dict) {
|
||||
load_projection_state_dict(state_dict);
|
||||
load_common_state_dict(state_dict);
|
||||
}
|
||||
|
||||
void Qwen3NextGatedDeltaNetImpl::load_projection_state_dict(
|
||||
const StateDict& state_dict) {
|
||||
auto qkvz_state_dict = state_dict.get_dict_with_prefix("in_proj_qkvz.");
|
||||
if (qkvz_state_dict.size() > 0 && !qkvz_proj_->is_weight_loaded()) {
|
||||
qkvz_proj_->load_state_dict(qkvz_state_dict);
|
||||
}
|
||||
|
||||
auto ba_state_dict = state_dict.get_dict_with_prefix("in_proj_ba.");
|
||||
if (ba_state_dict.size() > 0 && !ba_proj_->is_weight_loaded()) {
|
||||
ba_proj_->load_state_dict(ba_state_dict);
|
||||
}
|
||||
}
|
||||
|
||||
void Qwen3NextGatedDeltaNetImpl::verify_loaded_weights(
|
||||
const std::string& prefix) const {
|
||||
verify_projection_weights(prefix);
|
||||
verify_common_loaded_weights(prefix);
|
||||
}
|
||||
|
||||
void Qwen3NextGatedDeltaNetImpl::verify_projection_weights(
|
||||
const std::string& prefix) const {
|
||||
CHECK(qkvz_proj_ && qkvz_proj_->is_weight_loaded())
|
||||
<< "Missing required weight after all shards loaded: " << prefix
|
||||
<< "in_proj_qkvz.weight";
|
||||
CHECK(ba_proj_ && ba_proj_->is_weight_loaded())
|
||||
<< "Missing required weight after all shards loaded: " << prefix
|
||||
<< "in_proj_ba.weight";
|
||||
}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
66
ex_engine/xllm_layers/npu_torch/qwen3_next_gated_delta_net.h
Normal file
66
ex_engine/xllm_layers/npu_torch/qwen3_next_gated_delta_net.h
Normal file
@@ -0,0 +1,66 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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 <string>
|
||||
#include <utility>
|
||||
|
||||
#include "qwen3_gated_delta_net_base.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3NextGatedDeltaNetImpl : public Qwen3GatedDeltaNetBaseImpl {
|
||||
public:
|
||||
Qwen3NextGatedDeltaNetImpl() = default;
|
||||
Qwen3NextGatedDeltaNetImpl(const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options);
|
||||
|
||||
void load_state_dict(const StateDict& state_dict) override;
|
||||
void verify_loaded_weights(const std::string& prefix) const override;
|
||||
|
||||
protected:
|
||||
Qwen3NextGatedDeltaNetImpl(const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options,
|
||||
bool init_projections);
|
||||
|
||||
std::pair<torch::Tensor, torch::Tensor> project_decode_inputs(
|
||||
const torch::Tensor& hidden_states) override;
|
||||
std::pair<torch::Tensor, torch::Tensor> project_flat_inputs(
|
||||
const torch::Tensor& hidden_states) override;
|
||||
|
||||
virtual void load_projection_state_dict(const StateDict& state_dict);
|
||||
virtual void verify_projection_weights(const std::string& prefix) const;
|
||||
|
||||
void init_next_projections(const ModelArgs& args,
|
||||
const QuantArgs& quant_args,
|
||||
const ParallelArgs& parallel_args,
|
||||
const torch::TensorOptions& options);
|
||||
|
||||
private:
|
||||
ColumnParallelLinear qkvz_proj_{nullptr};
|
||||
ColumnParallelLinear ba_proj_{nullptr};
|
||||
};
|
||||
TORCH_MODULE(Qwen3NextGatedDeltaNet);
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
@@ -0,0 +1,176 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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.
|
||||
==============================================================================*/
|
||||
|
||||
#include "qwen3_next_hybrid_decoder_layer_base.h"
|
||||
|
||||
#include <algorithm>
|
||||
#include <optional>
|
||||
#include <tuple>
|
||||
|
||||
#include "common/flash_comm1_context.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
Qwen3HybridDecoderLayerImplBase::Qwen3HybridDecoderLayerImplBase(
|
||||
const ModelContext& context,
|
||||
int32_t layer_id,
|
||||
std::shared_ptr<Qwen3GatedDeltaNetBaseImpl> linear_attention_module) {
|
||||
const auto& model_args = context.get_model_args();
|
||||
const auto& quant_args = context.get_quant_args();
|
||||
const auto& parallel_args = context.get_parallel_args();
|
||||
const auto& options = context.get_tensor_options();
|
||||
const bool use_full_attention = is_full_attention_layer(model_args, layer_id);
|
||||
|
||||
// Initialize attention layers
|
||||
if (use_full_attention) {
|
||||
attention_ = register_module(
|
||||
"self_attn",
|
||||
Qwen3NextAttention(
|
||||
model_args, quant_args, parallel_args, options, layer_id));
|
||||
} else {
|
||||
linear_attention_ =
|
||||
register_module("linear_attn", std::move(linear_attention_module));
|
||||
}
|
||||
|
||||
// Initialize norm layers
|
||||
input_norm_ = register_module(
|
||||
"input_layernorm",
|
||||
Qwen3NextRMSNorm(
|
||||
model_args.hidden_size(), model_args.rms_norm_eps(), options));
|
||||
|
||||
post_norm_ = register_module(
|
||||
"post_attention_layernorm",
|
||||
Qwen3NextRMSNorm(
|
||||
model_args.hidden_size(), model_args.rms_norm_eps(), options));
|
||||
|
||||
// Initialize mlp
|
||||
auto mlp_only_layers = model_args.mlp_only_layers();
|
||||
if ((std::count(mlp_only_layers.begin(), mlp_only_layers.end(), layer_id) ==
|
||||
0) &&
|
||||
model_args.n_routed_experts() > 0 &&
|
||||
(layer_id + 1) % model_args.decoder_sparse_step() == 0) {
|
||||
moe_mlp_ = register_module("mlp",
|
||||
FusedMoE(model_args,
|
||||
FusedMoEArgs{.is_gated = true},
|
||||
quant_args,
|
||||
parallel_args,
|
||||
options));
|
||||
} else {
|
||||
mlp_ = register_module("mlp",
|
||||
DenseMLP(model_args.hidden_size(),
|
||||
model_args.intermediate_size(),
|
||||
true,
|
||||
false,
|
||||
model_args.hidden_act(),
|
||||
/*enable_result_reduction=*/true,
|
||||
quant_args,
|
||||
parallel_args.tp_group_,
|
||||
options));
|
||||
}
|
||||
}
|
||||
|
||||
void Qwen3HybridDecoderLayerImplBase::load_state_dict(
|
||||
const StateDict& state_dict) {
|
||||
if (attention_) {
|
||||
attention_->load_state_dict(state_dict.get_dict_with_prefix("self_attn."));
|
||||
} else {
|
||||
linear_attention_->load_state_dict(
|
||||
state_dict.get_dict_with_prefix("linear_attn."));
|
||||
}
|
||||
input_norm_->load_state_dict(
|
||||
state_dict.get_dict_with_prefix("input_layernorm."));
|
||||
post_norm_->load_state_dict(
|
||||
state_dict.get_dict_with_prefix("post_attention_layernorm."));
|
||||
if (moe_mlp_) {
|
||||
moe_mlp_->load_state_dict(state_dict.get_dict_with_prefix("mlp."));
|
||||
} else {
|
||||
mlp_->load_state_dict(state_dict.get_dict_with_prefix("mlp."));
|
||||
}
|
||||
}
|
||||
|
||||
void Qwen3HybridDecoderLayerImplBase::verify_loaded_weights(
|
||||
const std::string& prefix) const {
|
||||
if (linear_attention_) {
|
||||
linear_attention_->verify_loaded_weights(prefix + "linear_attn.");
|
||||
}
|
||||
}
|
||||
|
||||
torch::Tensor Qwen3HybridDecoderLayerImplBase::forward(
|
||||
torch::Tensor& x,
|
||||
std::optional<torch::Tensor>& residual,
|
||||
torch::Tensor& positions,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache,
|
||||
const ModelInputParams& input_params,
|
||||
const torch::Tensor& mrope_cos_sin) {
|
||||
const FlashComm1Context* fc1_ctx = get_current_flash_comm1_context();
|
||||
// Pre-attention norm
|
||||
if (!residual.has_value()) {
|
||||
residual = x;
|
||||
x = std::get<0>(input_norm_->forward(x));
|
||||
} else {
|
||||
if (fc1_ctx && is_sequence_sharded(*fc1_ctx) &&
|
||||
residual.value().size(0) != x.size(0)) {
|
||||
residual = maybe_shard_residual(residual.value(), *fc1_ctx);
|
||||
}
|
||||
if (fc1_ctx && is_sequence_sharded(*fc1_ctx)) {
|
||||
CHECK_EQ(residual.value().size(0), x.size(0))
|
||||
<< "FC1 input residual and hidden states must share the same "
|
||||
<< "padded local sequence layout.";
|
||||
}
|
||||
std::tie(x, residual) = input_norm_->forward(x, residual);
|
||||
}
|
||||
|
||||
// Attention
|
||||
if (attention_) {
|
||||
x = attention_->forward(
|
||||
positions, x, attn_metadata, kv_cache, mrope_cos_sin);
|
||||
} else {
|
||||
x = linear_attention_->forward(x, attn_metadata, kv_cache, input_params);
|
||||
}
|
||||
|
||||
// Post-attention norm
|
||||
// Ensure the residual layout matches the attention output before post_norm.
|
||||
if (fc1_ctx && is_sequence_sharded(*fc1_ctx) && residual.has_value() &&
|
||||
residual.value().size(0) != x.size(0)) {
|
||||
residual = maybe_shard_residual(residual.value(), *fc1_ctx);
|
||||
CHECK_EQ(residual.value().size(0), x.size(0))
|
||||
<< "FC1 post-attention residual and hidden states must share the same "
|
||||
<< "padded local sequence layout.";
|
||||
}
|
||||
|
||||
std::tie(x, residual) = post_norm_->forward(x, residual);
|
||||
|
||||
// MLP forward
|
||||
if (moe_mlp_) {
|
||||
x = moe_mlp_(x, input_params);
|
||||
} else {
|
||||
x = mlp_(x);
|
||||
}
|
||||
|
||||
return x;
|
||||
}
|
||||
|
||||
torch::Tensor Qwen3HybridDecoderLayerImplBase::build_mrope_cos_sin(
|
||||
const torch::Tensor& positions) const {
|
||||
if (attention_) {
|
||||
return attention_->build_mrope_cos_sin(positions);
|
||||
}
|
||||
return {};
|
||||
}
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
@@ -0,0 +1,90 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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 <memory>
|
||||
#include <string>
|
||||
|
||||
#include "framework/kv_cache/kv_cache.h"
|
||||
#include "framework/model/model_input_params.h"
|
||||
#include "framework/model_context.h"
|
||||
#include "framework/state_dict/state_dict.h"
|
||||
#include "layers/common/dense_mlp.h"
|
||||
#include "layers/common/qwen3_next_rms_norm.h"
|
||||
#include "layers/npu_torch/fused_moe.h"
|
||||
#include "layers/npu_torch/qwen3_gated_delta_net_base.h"
|
||||
#include "layers/npu_torch/qwen3_next_attention.h"
|
||||
|
||||
namespace xllm {
|
||||
namespace layer {
|
||||
|
||||
class Qwen3HybridDecoderLayerModule : public torch::nn::Module {
|
||||
public:
|
||||
virtual void load_state_dict(const StateDict& state_dict) = 0;
|
||||
virtual void verify_loaded_weights(const std::string& prefix) const = 0;
|
||||
virtual torch::Tensor forward(torch::Tensor& x,
|
||||
std::optional<torch::Tensor>& residual,
|
||||
torch::Tensor& positions,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache,
|
||||
const ModelInputParams& input_params,
|
||||
const torch::Tensor& mrope_cos_sin = {}) = 0;
|
||||
virtual torch::Tensor build_mrope_cos_sin(
|
||||
const torch::Tensor& positions) const {
|
||||
return {};
|
||||
}
|
||||
};
|
||||
|
||||
using Qwen3HybridDecoderLayerModulePtr =
|
||||
std::shared_ptr<Qwen3HybridDecoderLayerModule>;
|
||||
|
||||
class Qwen3HybridDecoderLayerImplBase : public Qwen3HybridDecoderLayerModule {
|
||||
public:
|
||||
explicit Qwen3HybridDecoderLayerImplBase(
|
||||
const ModelContext& context,
|
||||
int32_t layer_id,
|
||||
std::shared_ptr<Qwen3GatedDeltaNetBaseImpl> linear_attention_module);
|
||||
|
||||
void load_state_dict(const StateDict& state_dict) override;
|
||||
|
||||
void verify_loaded_weights(const std::string& prefix) const override;
|
||||
|
||||
torch::Tensor forward(torch::Tensor& x,
|
||||
std::optional<torch::Tensor>& residual,
|
||||
torch::Tensor& positions,
|
||||
const AttentionMetadata& attn_metadata,
|
||||
KVCache& kv_cache,
|
||||
const ModelInputParams& input_params,
|
||||
const torch::Tensor& mrope_cos_sin = {}) override;
|
||||
|
||||
torch::Tensor build_mrope_cos_sin(
|
||||
const torch::Tensor& positions) const override;
|
||||
|
||||
protected:
|
||||
Qwen3NextAttention attention_{nullptr};
|
||||
std::shared_ptr<Qwen3GatedDeltaNetBaseImpl> linear_attention_;
|
||||
|
||||
DenseMLP mlp_{nullptr};
|
||||
FusedMoE moe_mlp_{nullptr};
|
||||
|
||||
Qwen3NextRMSNorm input_norm_{nullptr};
|
||||
Qwen3NextRMSNorm post_norm_{nullptr};
|
||||
};
|
||||
|
||||
} // namespace layer
|
||||
} // namespace xllm
|
||||
231
ex_engine/xllm_models/llm/qwen3_5.h
Normal file
231
ex_engine/xllm_models/llm/qwen3_5.h
Normal file
@@ -0,0 +1,231 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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 <memory>
|
||||
#include <string>
|
||||
#include <unordered_set>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "models/model_registry.h"
|
||||
#if defined(USE_NPU) || defined(USE_MLU) || defined(USE_MUSA) || \
|
||||
defined(USE_DCU)
|
||||
#include "core/layers/qwen3_5_decoder_layer.h"
|
||||
#include "qwen3_next.h"
|
||||
#endif
|
||||
|
||||
namespace xllm {
|
||||
|
||||
#if defined(USE_NPU) || defined(USE_MLU) || defined(USE_MUSA) || \
|
||||
defined(USE_DCU)
|
||||
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::Tensor get_input_embeddings(torch::Tensor input_ids) {
|
||||
return get_word_embedding()(input_ids);
|
||||
}
|
||||
|
||||
void load_model(std::unique_ptr<ModelLoader> loader) {
|
||||
Qwen3NextForCausalLMImpl::load_model(
|
||||
std::move(loader), "model.language_model.", "lm_head.");
|
||||
}
|
||||
|
||||
void load_model(std::unique_ptr<ModelLoader> loader,
|
||||
const std::string& model_prefix) {
|
||||
Qwen3NextForCausalLMImpl::load_model(
|
||||
std::move(loader), model_prefix, "lm_head.");
|
||||
}
|
||||
};
|
||||
TORCH_MODULE(Qwen3_5ForCausalLM);
|
||||
#endif
|
||||
|
||||
#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(), 248046})); \
|
||||
LOAD_ARG_TEXT_OR_ROOT(mamba_ssm_dtype, "mamba_ssm_dtype", "float32")
|
||||
|
||||
#define LOAD_QWEN3_5_TEXT_TYPE_AND_DTYPE(default_model_type) \
|
||||
SET_ARG(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_MODEL_BACKEND(qwen3_5_text, "llm");
|
||||
#if defined(USE_NPU) || defined(USE_MLU) || defined(USE_MUSA) || \
|
||||
defined(USE_DCU)
|
||||
REGISTER_CAUSAL_MODEL(qwen3_5_text, Qwen3_5ForCausalLM);
|
||||
#endif
|
||||
REGISTER_MODEL_ARGS(qwen3_5_text, [&] {
|
||||
LOAD_QWEN3_5_TEXT_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_MODEL_BACKEND(qwen3_5_moe_text, "llm");
|
||||
#if defined(USE_NPU) || defined(USE_MLU) || defined(USE_MUSA) || \
|
||||
defined(USE_DCU)
|
||||
REGISTER_CAUSAL_MODEL(qwen3_5_moe_text, Qwen3_5ForCausalLM);
|
||||
#endif
|
||||
REGISTER_MODEL_ARGS(qwen3_5_moe_text, [&] {
|
||||
LOAD_QWEN3_5_TEXT_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_TEXT_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
|
||||
59
ex_engine/xllm_models/llm/qwen3_5_mtp.h
Normal file
59
ex_engine/xllm_models/llm/qwen3_5_mtp.h
Normal file
@@ -0,0 +1,59 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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 <memory>
|
||||
|
||||
#include "models/llm/qwen3_5.h"
|
||||
#include "models/llm/qwen3_5_mtp_base.h"
|
||||
#include "models/model_registry.h"
|
||||
|
||||
namespace xllm {
|
||||
|
||||
class Qwen3_5MtpModelImpl final : public Qwen3_5MtpModelImplBase {
|
||||
public:
|
||||
explicit Qwen3_5MtpModelImpl(const ModelContext& context)
|
||||
: Qwen3_5MtpModelImplBase(context) {}
|
||||
};
|
||||
|
||||
class Qwen3_5MtpForCausalLMImpl final : public Qwen3_5MtpForCausalLMImplBase {
|
||||
public:
|
||||
explicit Qwen3_5MtpForCausalLMImpl(const ModelContext& context)
|
||||
: Qwen3_5MtpForCausalLMImplBase(
|
||||
context,
|
||||
std::make_shared<Qwen3_5MtpModelImpl>(context)) {}
|
||||
};
|
||||
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 qwen3_5_mtp::load_model_args(
|
||||
json, args, "qwen3_5_text", "qwen3_5_mtp");
|
||||
});
|
||||
|
||||
REGISTER_MODEL_ARGS_LOADER(qwen3_5_moe_mtp,
|
||||
[](const JsonReader& json, ModelArgs* args) {
|
||||
return qwen3_5_mtp::load_model_args(
|
||||
json,
|
||||
args,
|
||||
"qwen3_5_moe_text",
|
||||
"qwen3_5_moe_mtp");
|
||||
});
|
||||
|
||||
} // namespace xllm
|
||||
299
ex_engine/xllm_models/llm/qwen3_5_mtp_base.h
Normal file
299
ex_engine/xllm_models/llm/qwen3_5_mtp_base.h
Normal file
@@ -0,0 +1,299 @@
|
||||
/* 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 <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "core/layers/common/linear.h"
|
||||
#include "core/layers/qwen3_5_decoder_layer.h"
|
||||
#include "models/llm/qwen3_next_hybrid_base.h"
|
||||
#include "models/model_registry.h"
|
||||
|
||||
namespace xllm {
|
||||
|
||||
namespace qwen3_5_mtp {
|
||||
|
||||
inline StateDict get_lm_head_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 std::string& prefix : kLmHeadPrefixes) {
|
||||
StateDict 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({}, "");
|
||||
}
|
||||
|
||||
inline bool load_model_args(const JsonReader& json,
|
||||
ModelArgs* args,
|
||||
const std::string& base_type,
|
||||
const std::string& mtp_type) {
|
||||
ModelArgsLoader base_loader = ModelRegistry::get_model_args_loader(base_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_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 qwen3_5_mtp
|
||||
|
||||
class Qwen3_5MtpModelImplBase : public Qwen3HybridModelImplBase {
|
||||
public:
|
||||
explicit Qwen3_5MtpModelImplBase(const ModelContext& context)
|
||||
: Qwen3HybridModelImplBase(context) {
|
||||
const torch::TensorOptions& 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_);
|
||||
}
|
||||
|
||||
layer::AttentionMetadata attn_metadata =
|
||||
layer::AttentionMetadataBuilder::build(
|
||||
input_params,
|
||||
model_args_.enable_mla(),
|
||||
build_attention_mask(input_params),
|
||||
/*device=*/device_);
|
||||
prepare_mrope(positions, attn_metadata);
|
||||
|
||||
torch::Tensor embedding = embed_tokens_(tokens);
|
||||
torch::Tensor hidden = input_params.embedding.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());
|
||||
torch::Tensor mrope_cos_sin;
|
||||
for (const layer::Qwen3HybridDecoderLayerModulePtr& 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) {
|
||||
if (!input_params.synchronize_layer(static_cast<uint32_t>(i))) {
|
||||
return ModelOutput();
|
||||
}
|
||||
mtp_hidden = layers_[i]->forward(mtp_hidden,
|
||||
residual,
|
||||
positions,
|
||||
attn_metadata,
|
||||
kv_caches[i],
|
||||
input_params,
|
||||
mrope_cos_sin);
|
||||
#if defined(USE_NPU)
|
||||
if (input_params.parallel.layer_synchronizer != nullptr &&
|
||||
!input_params.parallel.layer_synchronizer->record_event(
|
||||
static_cast<int64_t>(i), device_.index())) {
|
||||
return ModelOutput();
|
||||
}
|
||||
#endif
|
||||
}
|
||||
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) {
|
||||
StateDict 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) +
|
||||
".");
|
||||
}
|
||||
}
|
||||
|
||||
protected:
|
||||
virtual void prepare_mrope(const torch::Tensor& positions,
|
||||
layer::AttentionMetadata& attn_metadata) const {
|
||||
UNUSED_PARAMETER(positions);
|
||||
UNUSED_PARAMETER(attn_metadata);
|
||||
}
|
||||
|
||||
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_5MtpForCausalLMImplBase : public Qwen3HybridForCausalLMImplBase {
|
||||
public:
|
||||
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 std::unique_ptr<StateDict>& state_dict :
|
||||
loader->get_state_dicts()) {
|
||||
StateDict shared_embedding_state_dict =
|
||||
state_dict->get_dict_with_prefix(kEmbeddingPrefixes);
|
||||
StateDict 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 {
|
||||
StateDict lm_head_state_dict =
|
||||
qwen3_5_mtp::get_lm_head_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.");
|
||||
}
|
||||
|
||||
protected:
|
||||
Qwen3_5MtpForCausalLMImplBase(
|
||||
const ModelContext& context,
|
||||
std::shared_ptr<Qwen3_5MtpModelImplBase> mtp_model)
|
||||
: Qwen3HybridForCausalLMImplBase(context),
|
||||
mtp_model_(std::move(mtp_model)) {
|
||||
set_model_module(mtp_model_);
|
||||
}
|
||||
|
||||
private:
|
||||
std::shared_ptr<Qwen3_5MtpModelImplBase> mtp_model_;
|
||||
};
|
||||
|
||||
} // namespace xllm
|
||||
126
ex_engine/xllm_models/llm/qwen3_next.h
Normal file
126
ex_engine/xllm_models/llm/qwen3_next.h
Normal file
@@ -0,0 +1,126 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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
|
||||
364
ex_engine/xllm_models/llm/qwen3_next_hybrid_base.h
Normal file
364
ex_engine/xllm_models/llm/qwen3_next_hybrid_base.h
Normal file
@@ -0,0 +1,364 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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/flash_comm1_context.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/framework/parallel_state/parallel_args.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"
|
||||
#if defined(USE_NPU)
|
||||
#include "core/layers/npu_torch/qwen3_next_hybrid_decoder_layer_base.h"
|
||||
#elif defined(USE_MLU)
|
||||
#include "core/layers/mlu/qwen3_5/qwen3_5_hybrid_decoder_layer_base.h"
|
||||
#endif
|
||||
|
||||
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()),
|
||||
parallel_args_(context.get_parallel_args()),
|
||||
flash_comm1_options_(context.get_flash_comm1_options()) {
|
||||
if (model_args_.n_routed_experts() > 0) {
|
||||
flash_comm1_options_.enable_flashcomm1 = false;
|
||||
flash_comm1_options_.enable_mmrs_fusion = false;
|
||||
}
|
||||
|
||||
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));
|
||||
attn_mask_ = layer::AttentionMask(options.device(),
|
||||
options.dtype().toScalarType(),
|
||||
/*mask_value=*/-9984);
|
||||
dense_attn_mask_ = layer::AttentionMask(options.device(),
|
||||
options.dtype().toScalarType(),
|
||||
/*mask_value=*/1);
|
||||
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),
|
||||
/*device=*/device_);
|
||||
const int32_t num_tokens = static_cast<int32_t>(tokens.size(0));
|
||||
const auto& batch_forward_type = input_params.meta.batch_forward_type;
|
||||
const bool is_prefill_side = batch_forward_type.no_decode();
|
||||
FlashComm1Context fc1_ctx = build_flash_comm1_context(
|
||||
num_tokens, is_prefill_side, parallel_args_, flash_comm1_options_);
|
||||
FlashComm1ContextScope fc1_scope(&fc1_ctx);
|
||||
|
||||
torch::Tensor h;
|
||||
if (input_params.embedding.input_embedding.defined()) {
|
||||
h = input_params.embedding.input_embedding;
|
||||
} else {
|
||||
h = embed_tokens_(tokens);
|
||||
}
|
||||
|
||||
if (is_sequence_sharded(fc1_ctx)) {
|
||||
h = shard_sequence(h, fc1_ctx);
|
||||
}
|
||||
|
||||
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);
|
||||
#if defined(USE_NPU)
|
||||
if (input_params.parallel.layer_synchronizer != nullptr &&
|
||||
!input_params.parallel.layer_synchronizer->record_event(
|
||||
static_cast<int64_t>(i), device_.index())) {
|
||||
return ModelOutput();
|
||||
}
|
||||
#endif
|
||||
}
|
||||
auto [hidden_states, residual_out] = norm_->forward(h, residual);
|
||||
h = hidden_states;
|
||||
if (is_sequence_sharded(fc1_ctx)) {
|
||||
h = gather_sequence(h, fc1_ctx);
|
||||
}
|
||||
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) {
|
||||
#if defined(USE_NPU)
|
||||
// On NPU the hybrid path never consumes attn_metadata.attn_mask: full
|
||||
// attention runs through the fused-infer / paged-attention kernels (which
|
||||
// carry their own fixed fia_attn_mask or need no mask at all) and linear
|
||||
// attention is mask-free by construction. Materializing a dense
|
||||
// [seq_len, seq_len] mask here is pure waste and, for long sequences,
|
||||
// triggers an NPU OOM. Hand the kernels an empty mask unless a graph buffer
|
||||
// already supplies one.
|
||||
if (input_params.graph.attn_mask.defined()) {
|
||||
return input_params.graph.attn_mask;
|
||||
}
|
||||
return torch::Tensor();
|
||||
#else
|
||||
if (input_params.graph.attn_mask.defined()) {
|
||||
return input_params.graph.attn_mask;
|
||||
}
|
||||
max_seq_len_ = std::max(input_params.meta.kv_max_seq_len, max_seq_len_);
|
||||
const bool use_append_mask =
|
||||
input_params.is_spec_verify ||
|
||||
input_params.meta.batch_forward_type.is_mixed() ||
|
||||
input_params.meta.batch_forward_type.is_chunked_prefill();
|
||||
if (!use_append_mask) {
|
||||
return dense_attn_mask_.get_attn_mask(max_seq_len_, dtype_, device_);
|
||||
}
|
||||
|
||||
const int32_t num_sequences = input_params.meta.num_sequences;
|
||||
if (num_sequences <= 0) {
|
||||
return dense_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.attention.host.q_seq_lens[j],
|
||||
input_params.attention.host.kv_seq_lens[j],
|
||||
max_seq_len_,
|
||||
dtype_,
|
||||
device_));
|
||||
}
|
||||
return torch::cat(req_mask_vec, 0);
|
||||
#endif
|
||||
}
|
||||
|
||||
ModelArgs model_args_;
|
||||
torch::nn::ModuleList blocks_{nullptr};
|
||||
std::vector<layer::Qwen3HybridDecoderLayerModulePtr> layers_;
|
||||
int32_t max_seq_len_ = 0;
|
||||
int32_t dp_size_ = 1;
|
||||
ParallelArgs parallel_args_;
|
||||
FlashComm1Options flash_comm1_options_;
|
||||
torch::Device device_;
|
||||
torch::ScalarType dtype_ = torch::kFloat;
|
||||
layer::Qwen3NextRMSNorm norm_{nullptr};
|
||||
layer::AttentionMask attn_mask_;
|
||||
layer::AttentionMask dense_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) {
|
||||
load_model(std::move(loader), "model.", "lm_head.");
|
||||
}
|
||||
|
||||
void load_model(std::unique_ptr<ModelLoader> loader,
|
||||
const std::string& model_prefix) {
|
||||
load_model(std::move(loader), model_prefix, "lm_head.");
|
||||
}
|
||||
|
||||
void load_model(std::unique_ptr<ModelLoader> loader,
|
||||
const std::string& model_prefix,
|
||||
const std::string& lm_head_prefix) {
|
||||
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_prefix);
|
||||
model_->load_state_dict(model_state_dict);
|
||||
|
||||
auto lm_head_state_dict =
|
||||
state_dict->get_dict_with_prefix(lm_head_prefix);
|
||||
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_prefix);
|
||||
}
|
||||
|
||||
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; }
|
||||
|
||||
bool is_hybrid_linear_attention() { return true; }
|
||||
|
||||
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
|
||||
440
ex_engine/xllm_models/vlm/qwen3_5.h
Normal file
440
ex_engine/xllm_models/vlm/qwen3_5.h
Normal file
@@ -0,0 +1,440 @@
|
||||
/* Copyright 2025-2026 The xLLM Authors.
|
||||
|
||||
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/rotary_embedding_util.h"
|
||||
#include "models/model_registry.h"
|
||||
#include "models/vlm/mposition/mposition.h"
|
||||
#include "models/vlm/qwen3_vl_base.h"
|
||||
#include "processors/multimodal_processor.h"
|
||||
#include "processors/qwen2_vl_image_processor.h"
|
||||
#include "processors/qwen3_vl_prompt_processor.h"
|
||||
#include "processors/qwen3_vl_video_processor.h"
|
||||
|
||||
#if defined(USE_NPU)
|
||||
#include "models/llm/qwen3_5.h"
|
||||
#include "models/vlm/npu/qwen3_vl.h"
|
||||
#elif defined(USE_MLU) || defined(USE_DCU)
|
||||
#include "core/layers/common/qwen3_next_rms_norm.h"
|
||||
#include "core/layers/common/rms_norm.h"
|
||||
#include "core/layers/qwen3_5_decoder_layer.h"
|
||||
#include "core/layers/qwen3_vision_layer.h"
|
||||
#include "models/llm/llm_model_base.h"
|
||||
#include "qwen3_vl.h"
|
||||
#endif
|
||||
|
||||
namespace xllm {
|
||||
#if !defined(USE_NPU)
|
||||
|
||||
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 {
|
||||
return layer::rotary::apply_mrope(cos_sin_, positions, mrope_section_);
|
||||
}
|
||||
|
||||
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.parallel.dp_global_token_nums;
|
||||
std::replace(dp_token_nums.begin(), dp_token_nums.end(), 0, 1);
|
||||
}
|
||||
|
||||
auto inputs_embeds = input_params.embedding.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);
|
||||
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,
|
||||
/*enable_mla=*/false,
|
||||
/*attn_mask=*/{},
|
||||
h.device());
|
||||
// Init batch and token_block_offset for GDN attention
|
||||
if (attn_metadata.is_prefill || attn_metadata.is_chunked_prefill) {
|
||||
constexpr int32_t kBlockM = 64;
|
||||
constexpr int64_t pad_slot_id = -1;
|
||||
constexpr int64_t default_max_num_programs = 1024;
|
||||
constexpr int64_t chunk_size = 64;
|
||||
auto seqlens = attn_metadata.q_cu_seq_lens.diff();
|
||||
auto nums = (seqlens + kBlockM - 1) / kBlockM;
|
||||
nums = nums.to(torch::kLong);
|
||||
int32_t tot = nums.sum().item<int32_t>();
|
||||
torch::Tensor range_batch = torch::arange(nums.size(0), nums.options());
|
||||
torch::Tensor mlist_tensor = torch::repeat_interleave(range_batch, nums);
|
||||
int64_t mlist_len = mlist_tensor.size(0);
|
||||
int64_t max_num_programs =
|
||||
std::max(default_max_num_programs, mlist_len) * 2;
|
||||
torch::Tensor batch_ptr =
|
||||
torch::full({max_num_programs},
|
||||
pad_slot_id,
|
||||
torch::dtype(torch::kInt32).device(seqlens.device()));
|
||||
torch::Tensor token_block_offset_ptr =
|
||||
torch::full({max_num_programs},
|
||||
pad_slot_id,
|
||||
torch::dtype(torch::kInt32).device(seqlens.device()));
|
||||
|
||||
std::vector<torch::Tensor> vec;
|
||||
vec.reserve(nums.size(0));
|
||||
for (int64_t i = 0; i < nums.size(0); ++i) {
|
||||
vec.emplace_back(
|
||||
torch::arange(nums[i].item<int64_t>(), nums.options()));
|
||||
}
|
||||
torch::Tensor offsetlist_tensor = torch::cat(vec, -1).to(torch::kInt32);
|
||||
batch_ptr.narrow(0, 0, mlist_len).copy_(mlist_tensor);
|
||||
token_block_offset_ptr.narrow(0, 0, mlist_len).copy_(offsetlist_tensor);
|
||||
|
||||
// Compute chunk indices for the chunked GDN kernel
|
||||
{
|
||||
torch::Tensor lengths = seqlens;
|
||||
torch::Tensor num_chunks = (lengths + chunk_size - 1) / chunk_size;
|
||||
num_chunks = num_chunks.to(torch::kLong);
|
||||
torch::Tensor cumsum = torch::cumsum(num_chunks, 0);
|
||||
int64_t total_chunks = cumsum[-1].item<int64_t>();
|
||||
torch::Tensor arange_total =
|
||||
torch::arange(total_chunks, attn_metadata.q_cu_seq_lens.options());
|
||||
torch::Tensor zeros = torch::zeros({1}, cumsum.options());
|
||||
torch::Tensor prefix = torch::cat(
|
||||
{zeros, cumsum.slice(/*dim=*/0, /*start=*/0, /*end=*/-1)});
|
||||
torch::Tensor repeats_prefix =
|
||||
torch::repeat_interleave(prefix, num_chunks);
|
||||
torch::Tensor indices = arange_total - repeats_prefix;
|
||||
torch::Tensor mask = indices == 0;
|
||||
torch::Tensor col0 = mask.cumsum(0) - 1;
|
||||
attn_metadata.chunk_indices = torch::stack({col0, indices}, /*dim=*/1)
|
||||
.to(attn_metadata.q_cu_seq_lens)
|
||||
.to(torch::kInt32);
|
||||
}
|
||||
attn_metadata.tot = tot;
|
||||
attn_metadata.batch = batch_ptr;
|
||||
attn_metadata.token_block_offset = token_block_offset_ptr;
|
||||
}
|
||||
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);
|
||||
|
||||
#endif // !defined(USE_NPU)
|
||||
|
||||
#if defined(USE_NPU)
|
||||
using Qwen3_5_VisionTransformer = npu::model::Qwen3_VisionTransformer;
|
||||
#else
|
||||
using Qwen3_5_VisionTransformer = Qwen3_VisionTransformer;
|
||||
#endif
|
||||
|
||||
using Qwen3_5ForConditionalGenerationImpl =
|
||||
Qwen3VLForConditionalGenerationBase<Qwen3_5_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(bos_token_id, "text_config.bos_token_id", 151643); \
|
||||
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", \
|
||||
static_cast<int32_t>(args->n_heads() * 2)); \
|
||||
LOAD_ARG_OR( \
|
||||
full_attention_interval, "text_config.full_attention_interval", 4); \
|
||||
LOAD_ARG_OR(attn_output_gate, "text_config.attn_output_gate", true); \
|
||||
LOAD_ARG_OR( \
|
||||
num_nextn_predict_layers, "text_config.mtp_num_hidden_layers", 0); \
|
||||
LOAD_ARG_OR(num_nextn_predict_layers, \
|
||||
"text_config.num_nextn_predict_layers", \
|
||||
args->num_nextn_predict_layers()); \
|
||||
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_OR(rope_scaling_mrope_section, \
|
||||
"text_config.rope_parameters.mrope_section", \
|
||||
std::vector<int64_t>({11, 11, 10})); \
|
||||
LOAD_ARG_OR(rope_scaling_mrope_interleaved, \
|
||||
"text_config.rope_parameters.mrope_interleaved", \
|
||||
true); \
|
||||
LOAD_ARG_OR(rope_scaling_rope_type, \
|
||||
"text_config.rope_parameters.rope_type", \
|
||||
"default"); \
|
||||
if (args->rope_scaling_rope_type() == "default") { \
|
||||
args->rope_scaling_rope_type() = "mrope"; \
|
||||
} \
|
||||
LOAD_ARG_OR(partial_rotary_factor, \
|
||||
"text_config.rope_parameters.partial_rotary_factor", \
|
||||
0.25f); \
|
||||
LOAD_ARG_OR(mamba_ssm_dtype, "text_config.mamba_ssm_dtype", "float32")
|
||||
|
||||
#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_OR(mm_deepstack_visual_indexes, \
|
||||
"vision_config.deepstack_visual_indexes", \
|
||||
std::vector<int64_t>()); \
|
||||
if (!args->mm_deepstack_visual_indexes().empty()) { \
|
||||
LOG(FATAL) << "qwen3_5 VLM does not support 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", \
|
||||
args->hidden_size()); \
|
||||
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(); \
|
||||
})
|
||||
|
||||
// qwen3_5/qwen3_5_moe are multimodal entry points. On NPU, text-only serving
|
||||
// uses qwen3_5_text/qwen3_5_moe_text from llm/qwen3_5.h because the VLM
|
||||
// request protocol currently requires array-form chat content.
|
||||
REGISTER_CAUSAL_VLM_MODEL(qwen3_5, Qwen3_5ForConditionalGeneration);
|
||||
REGISTER_MPOSITION_GENERATOR(qwen3_5, Qwen3VLMPositionGenerator);
|
||||
using Qwen35MultimodalProcessor = MultimodalProcessor<Qwen3VLPromptProcessor,
|
||||
Qwen2VLImageProcessor,
|
||||
Qwen3VLVideoProcessor>;
|
||||
REGISTER_MULTIMODAL_PROCESSOR(qwen3_5, Qwen35MultimodalProcessor);
|
||||
REGISTER_MODEL_ARGS(qwen3_5, [&] {
|
||||
LOAD_QWEN3_5_COMMON_ARGS();
|
||||
LOAD_QWEN3_5_VISION_ARGS();
|
||||
|
||||
SET_ARG(num_experts, 0);
|
||||
SET_ARG(n_routed_experts, 0);
|
||||
SET_ARG(n_shared_experts, 0);
|
||||
|
||||
SET_ARG(stop_token_ids,
|
||||
std::unordered_set<int32_t>({args->eos_token_id(), 248046}));
|
||||
});
|
||||
|
||||
REGISTER_CAUSAL_VLM_MODEL(qwen3_5_moe, Qwen3_5ForConditionalGeneration);
|
||||
REGISTER_MPOSITION_GENERATOR(qwen3_5_moe, Qwen3VLMPositionGenerator);
|
||||
REGISTER_MULTIMODAL_PROCESSOR(qwen3_5_moe, Qwen35MultimodalProcessor);
|
||||
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(), 248046}));
|
||||
});
|
||||
|
||||
// Text-only model registrations. On NPU these are handled by llm/qwen3_5.h.
|
||||
#if !defined(USE_NPU)
|
||||
// qwen3_5 without vision config (text-only serving).
|
||||
// Model args are already registered by the VLM registration above.
|
||||
REGISTER_CAUSAL_MODEL_WITH_VARNAME(qwen3_5_lm, qwen3_5, Qwen3_5ForCausalLM);
|
||||
REGISTER_CAUSAL_MODEL_WITH_VARNAME(qwen3_5_moe_lm,
|
||||
qwen3_5_moe,
|
||||
Qwen3_5ForCausalLM);
|
||||
|
||||
REGISTER_CAUSAL_MODEL(qwen3_5_text, Qwen3_5ForCausalLM);
|
||||
REGISTER_MODEL_ARGS(qwen3_5_text, [&] {
|
||||
LOAD_QWEN3_5_COMMON_ARGS();
|
||||
SET_ARG(num_experts, 0);
|
||||
SET_ARG(n_routed_experts, 0);
|
||||
SET_ARG(n_shared_experts, 0);
|
||||
SET_ARG(decoder_sparse_step, 1);
|
||||
SET_ARG(stop_token_ids,
|
||||
std::unordered_set<int32_t>({args->eos_token_id(), 248046}));
|
||||
});
|
||||
|
||||
REGISTER_CAUSAL_MODEL(qwen3_5_moe_text, Qwen3_5ForCausalLM);
|
||||
REGISTER_MODEL_ARGS(qwen3_5_moe_text, [&] {
|
||||
LOAD_QWEN3_5_COMMON_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(), 248046}));
|
||||
});
|
||||
#endif // !defined(USE_NPU)
|
||||
|
||||
#undef LOAD_QWEN3_5_VISION_ARGS
|
||||
#undef LOAD_QWEN3_5_COMMON_ARGS
|
||||
|
||||
} // namespace xllm
|
||||
@@ -410,6 +410,16 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
||||
|
||||
return None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def fold_max_completion_tokens(cls, data):
|
||||
"""OpenAI newer API: max_completion_tokens → max_tokens alias."""
|
||||
if isinstance(data, dict):
|
||||
mct = data.pop("max_completion_tokens", None)
|
||||
if mct is not None and data.get("max_tokens") is None:
|
||||
data["max_tokens"] = mct
|
||||
return data
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def normalize_messages(cls, data):
|
||||
|
||||
Reference in New Issue
Block a user