diff --git a/csrc/layers/linear/fused_linear.cpp b/csrc/layers/linear/fused_linear.cpp index 1bcb96c94..321ff61dc 100644 --- a/csrc/layers/linear/fused_linear.cpp +++ b/csrc/layers/linear/fused_linear.cpp @@ -1,8 +1,105 @@ #include "fused_linear.hpp" -#include +#include namespace infinilm::layers::linear { +size_t MergedColumnParallelLinear::calculate_local_output_size( + const MergedLinearSplit &split, + size_t tp_size) { + if (tp_size == 0) { + throw std::runtime_error("`tp_size` must be positive"); + } + if (split.output_size == 0) { + throw std::runtime_error("`output_size` must be positive"); + } + if (split.num_shards == 0) { + if (split.output_size % tp_size != 0) { + throw std::runtime_error( + "`output_size` must be divisible by `tp_size`"); + } + return split.output_size / tp_size; + } + if (split.output_size % split.num_shards != 0) { + throw std::runtime_error( + "`output_size` must be divisible by `num_shards`"); + } + if (split.num_shards % tp_size == 0) { + return split.output_size / tp_size; + } + if (tp_size % split.num_shards == 0) { + return split.output_size / split.num_shards; + } + throw std::runtime_error( + "`num_shards` and `tp_size` are incompatible"); +} + +size_t MergedColumnParallelLinear::calculate_total_output_size( + const std::vector &splits, + size_t tp_size) { + if (splits.empty()) { + throw std::runtime_error("`splits` must not be empty"); + } + size_t local_output_size = 0; + for (size_t i = 0; i < splits.size(); ++i) { + const auto &split = splits[i]; + if (split.name.empty()) { + throw std::runtime_error("`split.name` must not be empty"); + } + for (size_t j = 0; j < i; ++j) { + if (splits[j].name == split.name) { + throw std::runtime_error("duplicate split name `" + split.name + "`"); + } + } + local_output_size += calculate_local_output_size(split, tp_size); + } + return local_output_size * tp_size; +} + +MergedColumnParallelLinear::MergedColumnParallelLinear( + size_t hidden_size, + const std::vector &splits, + RegisterParamFn register_fn, + bool bias, + const infinicore::DataType &dtype, + const infinicore::Device &device, + engine::distributed::RankInfo rank_info) + : infinilm::nn::ColumnParallelLinear( + hidden_size, + calculate_total_output_size(splits, rank_info.tp_size), + std::make_shared(), + bias, + dtype, + device, + rank_info.tp_rank, + rank_info.tp_size), + register_fn_(std::move(register_fn)) { + if (!register_fn_) { + throw std::runtime_error("`register_fn` must not be empty"); + } + split_infos_.reserve(splits.size()); + size_t start = 0; + for (const auto &split : splits) { + const size_t local_output_size = calculate_local_output_size( + split, rank_info.tp_size); + split_infos_.push_back( + {split.name, start, local_output_size, split.num_shards}); + start += local_output_size; + } + register_split_parameters(); +} + +void MergedColumnParallelLinear::process_weights_after_loading() { + BaseLinear::process_weights_after_loading(); + register_split_parameters(); +} + +void MergedColumnParallelLinear::register_split_parameters() { + auto params = this->split_params(split_infos_, tp_rank_, tp_size_, -1); + for (auto ¶m : params) { + register_fn_(param.full_name, std::move(param.param)); + } +} + // --------------------------------------------------------- // QKV Parallel Linear // --------------------------------------------------------- @@ -31,14 +128,14 @@ QKVParallelLinear::QKVParallelLinear(size_t hidden_size, const infinicore::Device &device, engine::distributed::RankInfo rank_info) : infinilm::nn::ColumnParallelLinear( - hidden_size, - calculate_out_feature_size(num_q_head, q_dim, num_k_head, k_dim, num_v_head, v_dim, rank_info), - quantization == nullptr ? std::make_shared() : quantization, - (q_bias || k_bias || v_bias), - dtype, - device, - rank_info.tp_rank, - rank_info.tp_size), + hidden_size, + calculate_out_feature_size(num_q_head, q_dim, num_k_head, k_dim, num_v_head, v_dim, rank_info), + quantization == nullptr ? std::make_shared() : quantization, + (q_bias || k_bias || v_bias), + dtype, + device, + rank_info.tp_rank, + rank_info.tp_size), q_dim_(q_dim), k_dim_(k_dim), v_dim_(v_dim), @@ -134,14 +231,14 @@ GateUpParallelLinear::GateUpParallelLinear(size_t hidden_size, size_t intermedia const infinicore::DataType &dtype, const infinicore::Device &device, engine::distributed::RankInfo rank_info) : infinilm::nn::ColumnParallelLinear( - hidden_size, - intermediate_size * 2, - quantization == nullptr ? std::make_shared() : quantization, - gate_bias || up_bias, - dtype, - device, - rank_info.tp_rank, - rank_info.tp_size), + hidden_size, + intermediate_size * 2, + quantization == nullptr ? std::make_shared() : quantization, + gate_bias || up_bias, + dtype, + device, + rank_info.tp_rank, + rank_info.tp_size), gate_bias_(gate_bias), up_bias_(up_bias) { if (gate_bias_ != up_bias_) { diff --git a/csrc/layers/linear/fused_linear.hpp b/csrc/layers/linear/fused_linear.hpp index 8773a081c..9ec0d2a44 100644 --- a/csrc/layers/linear/fused_linear.hpp +++ b/csrc/layers/linear/fused_linear.hpp @@ -1,12 +1,52 @@ #pragma once + #include "../../engine/distributed/communication_group.hpp" #include "../quantization/quantization.hpp" #include "linear.hpp" + +#include #include +#include +#include +#include +#include +#include namespace infinilm::layers::linear { using RegisterParamFn = std::function; +struct MergedLinearSplit { + std::string name; + size_t output_size; + size_t num_shards = 0; +}; + +class MergedColumnParallelLinear : public infinilm::nn::ColumnParallelLinear { +public: + MergedColumnParallelLinear( + size_t hidden_size, + const std::vector &splits, + RegisterParamFn register_fn, + bool bias = false, + const infinicore::DataType &dtype = infinicore::DataType::F32, + const infinicore::Device &device = infinicore::Device(), + engine::distributed::RankInfo rank_info = engine::distributed::RankInfo()); + + void process_weights_after_loading() override; + +private: + static size_t calculate_local_output_size( + const MergedLinearSplit &split, + size_t tp_size); + static size_t calculate_total_output_size( + const std::vector &splits, + size_t tp_size); + void register_split_parameters(); + + RegisterParamFn register_fn_; + std::vector split_infos_; +}; + class QKVParallelLinear : public infinilm::nn::ColumnParallelLinear { public: explicit QKVParallelLinear(size_t hidden_size, diff --git a/csrc/layers/quantization/none_quantization.cpp b/csrc/layers/quantization/none_quantization.cpp index 3184da991..6f9808fa4 100644 --- a/csrc/layers/quantization/none_quantization.cpp +++ b/csrc/layers/quantization/none_quantization.cpp @@ -82,11 +82,17 @@ std::vector NoneQuantization::split_params( std::vector result; auto weight_it = params.find("weight"); auto bias_it = params.find("bias"); + // Keep named parameters in checkpoint layout [OC, IC], even when the + // backing weight is packed as [IC, OC]. This view still aliases the packed + // storage, so TP loading through the named parameters updates the GEMM weight. + auto weight = weight_prepacked_ + ? weight_it->second->permute({1, 0}) + : static_cast(weight_it->second); for (const auto &s : splits) { result.push_back({s.prefix + ".weight", infinicore::nn::Parameter( - weight_it->second->narrow({{static_cast(narrow_dim), s.start, s.size}}), + weight->narrow({{static_cast(narrow_dim), s.start, s.size}}), narrow_dim, tp_rank, tp_size, s.num_shards)}); if (bias_it != params.end()) { result.push_back({s.prefix + ".bias", diff --git a/csrc/models/qwen3_next/qwen3_next_gated_deltanet.cpp b/csrc/models/qwen3_next/qwen3_next_gated_deltanet.cpp index 022454247..8cac86c3f 100644 --- a/csrc/models/qwen3_next/qwen3_next_gated_deltanet.cpp +++ b/csrc/models/qwen3_next/qwen3_next_gated_deltanet.cpp @@ -12,6 +12,7 @@ #include #include #include +#include #include namespace infinilm::models::qwen3_next { @@ -114,6 +115,7 @@ Qwen3NextGatedDeltaNet::Qwen3NextGatedDeltaNet(std::shared_ptrregister_module("conv1d", model_config, layer_idx, device); - size_t projection_size_qkv = local_key_dim_ * 2 + local_value_dim_; auto quantization_method = model_config->get_quantization_method(); auto register_fn = [this](const std::string &n, infinicore::nn::Parameter p) { this->register_parameter(n, std::move(p)); }; - in_proj_qkv_ = std::make_shared( - hidden_size, linear_key_head_dim, linear_key_head_dim, linear_value_head_dim, linear_num_key_heads, linear_num_key_heads, linear_num_value_heads, - false, false, false, - "in_proj_q", "in_proj_k", "in_proj_v", register_fn, - quantization_method, dtype, device, rank_info); - in_proj_z_ = this->register_module("in_proj_z", hidden_size, value_dim, false, dtype, device, tp_rank, tp_size); - in_proj_a_ = this->register_module("in_proj_a", hidden_size, linear_num_value_heads, false, dtype, device, tp_rank, tp_size); - in_proj_b_ = this->register_module("in_proj_b", hidden_size, linear_num_value_heads, false, dtype, device, tp_rank, tp_size); + if (quantization_method->get_quant_scheme() + == infinilm::quantization::QuantScheme::NONE) { + in_proj_qkvz_ = std::make_shared( + hidden_size, + std::vector{ + {"in_proj_q", key_dim}, + {"in_proj_k", key_dim, linear_num_key_heads}, + {"in_proj_v", value_dim, linear_num_value_heads}, + {"in_proj_z", value_dim}, + }, + register_fn, + false, + dtype, + device, + rank_info); + in_proj_ba_ = std::make_shared( + hidden_size, + linear_num_value_heads, + "in_proj_b", + "in_proj_a", + register_fn, + nullptr, + false, + dtype, + device, + rank_info); + } else { + in_proj_qkv_ = std::make_shared( + hidden_size, linear_key_head_dim, linear_key_head_dim, linear_value_head_dim, linear_num_key_heads, linear_num_key_heads, linear_num_value_heads, + false, false, false, + "in_proj_q", "in_proj_k", "in_proj_v", register_fn, + quantization_method, dtype, device, rank_info); + in_proj_z_ = this->register_module("in_proj_z", hidden_size, value_dim, false, dtype, device, tp_rank, tp_size); + in_proj_a_ = this->register_module("in_proj_a", hidden_size, linear_num_value_heads, false, dtype, device, tp_rank, tp_size); + in_proj_b_ = this->register_module("in_proj_b", hidden_size, linear_num_value_heads, false, dtype, device, tp_rank, tp_size); + } INFINICORE_NN_PARAMETER_INIT(dt_bias, ({linear_num_value_heads}, dtype, device, 0, tp_rank, tp_size)); INFINICORE_NN_PARAMETER_INIT(A_log, ({linear_num_value_heads}, dtype, device, 0, tp_rank, tp_size)); @@ -150,10 +179,23 @@ infinicore::Tensor Qwen3NextGatedDeltaNet::forward(const infinicore::Tensor &hid size_t batch_size = shape[0]; size_t seq_len = shape[1]; - auto qkv = in_proj_qkv_->forward(hidden_states_mutable); - auto z = in_proj_z_->forward(hidden_states_mutable); - auto a = in_proj_a_->forward(hidden_states_mutable); - auto b = in_proj_b_->forward(hidden_states_mutable); + infinicore::Tensor qkv; + infinicore::Tensor z; + infinicore::Tensor a; + infinicore::Tensor b; + if (in_proj_qkvz_) { + auto qkvz = in_proj_qkvz_->forward(hidden_states_mutable); + const size_t output_dim = qkvz->ndim() - 1; + const size_t qkv_size = local_key_dim_ * 2 + local_value_dim_; + qkv = qkvz->narrow({{output_dim, 0, qkv_size}}); + z = qkvz->narrow({{output_dim, qkv_size, local_value_dim_}}); + std::tie(b, a) = in_proj_ba_->forward_split(hidden_states_mutable); + } else { + qkv = in_proj_qkv_->forward(hidden_states_mutable); + z = in_proj_z_->forward(hidden_states_mutable); + a = in_proj_a_->forward(hidden_states_mutable); + b = in_proj_b_->forward(hidden_states_mutable); + } auto &forward_context = infinilm::global_state::get_forward_context(); auto &mamba_metadata = forward_context.mamba_metadata; @@ -245,4 +287,11 @@ infinicore::Tensor Qwen3NextGatedDeltaNet::forward(const infinicore::Tensor &hid return out_proj_->forward(gated); } +void Qwen3NextGatedDeltaNet::process_weights_after_loading() { + if (in_proj_qkvz_) { + in_proj_qkvz_->process_weights_after_loading(); + in_proj_ba_->process_weights_after_loading(); + } +} + } // namespace infinilm::models::qwen3_next diff --git a/csrc/models/qwen3_next/qwen3_next_gated_deltanet.hpp b/csrc/models/qwen3_next/qwen3_next_gated_deltanet.hpp index e7d64e7a9..e6290cfaf 100644 --- a/csrc/models/qwen3_next/qwen3_next_gated_deltanet.hpp +++ b/csrc/models/qwen3_next/qwen3_next_gated_deltanet.hpp @@ -34,8 +34,11 @@ class Qwen3NextGatedDeltaNet : public infinicore::nn::Module { const infinicore::Device &device); infinicore::Tensor forward(const infinicore::Tensor &hidden_states) const; + void process_weights_after_loading() override; private: + std::shared_ptr in_proj_qkvz_; + std::shared_ptr in_proj_ba_; std::shared_ptr in_proj_qkv_; std::shared_ptr in_proj_z_; std::shared_ptr in_proj_a_;