Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
131 changes: 114 additions & 17 deletions csrc/layers/linear/fused_linear.cpp
Original file line number Diff line number Diff line change
@@ -1,8 +1,105 @@
#include "fused_linear.hpp"

#include <spdlog/spdlog.h>
#include <utility>

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<MergedLinearSplit> &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<MergedLinearSplit> &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<infinilm::quantization::NoneQuantization>(),
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 &param : params) {
register_fn_(param.full_name, std::move(param.param));
}
}

// ---------------------------------------------------------
// QKV Parallel Linear
// ---------------------------------------------------------
Expand Down Expand Up @@ -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<infinilm::quantization::NoneQuantization>() : 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<infinilm::quantization::NoneQuantization>() : 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),
Expand Down Expand Up @@ -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<infinilm::quantization::NoneQuantization>() : 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<infinilm::quantization::NoneQuantization>() : 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_) {
Expand Down
40 changes: 40 additions & 0 deletions csrc/layers/linear/fused_linear.hpp
Original file line number Diff line number Diff line change
@@ -1,12 +1,52 @@
#pragma once

#include "../../engine/distributed/communication_group.hpp"
#include "../quantization/quantization.hpp"
#include "linear.hpp"

#include <cstddef>
#include <functional>
#include <memory>
#include <stdexcept>
#include <string>
#include <tuple>
#include <vector>

namespace infinilm::layers::linear {
using RegisterParamFn = std::function<void(const std::string &, infinicore::nn::Parameter)>;

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<MergedLinearSplit> &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<MergedLinearSplit> &splits,
size_t tp_size);
void register_split_parameters();

RegisterParamFn register_fn_;
std::vector<infinilm::quantization::SplitInfo> split_infos_;
};

class QKVParallelLinear : public infinilm::nn::ColumnParallelLinear {
public:
explicit QKVParallelLinear(size_t hidden_size,
Expand Down
8 changes: 7 additions & 1 deletion csrc/layers/quantization/none_quantization.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -82,11 +82,17 @@ std::vector<SplitParam> NoneQuantization::split_params(
std::vector<SplitParam> 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<const infinicore::Tensor &>(weight_it->second);

for (const auto &s : splits) {
result.push_back({s.prefix + ".weight",
infinicore::nn::Parameter(
weight_it->second->narrow({{static_cast<size_t>(narrow_dim), s.start, s.size}}),
weight->narrow({{static_cast<size_t>(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",
Expand Down
75 changes: 62 additions & 13 deletions csrc/models/qwen3_next/qwen3_next_gated_deltanet.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
#include <cstdint>
#include <optional>
#include <stdexcept>
#include <tuple>
#include <vector>

namespace infinilm::models::qwen3_next {
Expand Down Expand Up @@ -114,6 +115,7 @@ Qwen3NextGatedDeltaNet::Qwen3NextGatedDeltaNet(std::shared_ptr<infinilm::config:
local_num_key_heads_ = linear_num_key_heads / tp_size;
key_head_dim_ = linear_key_head_dim;
value_head_dim_ = linear_value_head_dim;
size_t key_dim = linear_key_head_dim * linear_num_key_heads;
size_t value_dim = linear_value_head_dim * linear_num_value_heads;
local_key_dim_ = key_head_dim_ * local_num_key_heads_;
local_value_dim_ = value_head_dim_ * local_num_value_heads_;
Expand All @@ -122,17 +124,44 @@ Qwen3NextGatedDeltaNet::Qwen3NextGatedDeltaNet(std::shared_ptr<infinilm::config:

conv1d_ = this->register_module<Qwen3NextCausalConv1D>("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<layers::linear::QKVParallelLinear>(
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<infinilm::layers::linear::ColumnParallelLinear>("in_proj_z", hidden_size, value_dim, false, dtype, device, tp_rank, tp_size);
in_proj_a_ = this->register_module<infinilm::layers::linear::ColumnParallelLinear>("in_proj_a", hidden_size, linear_num_value_heads, false, dtype, device, tp_rank, tp_size);
in_proj_b_ = this->register_module<infinilm::layers::linear::ColumnParallelLinear>("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<layers::linear::MergedColumnParallelLinear>(
hidden_size,
std::vector<layers::linear::MergedLinearSplit>{
{"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<layers::linear::GateUpParallelLinear>(
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<layers::linear::QKVParallelLinear>(
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<infinilm::layers::linear::ColumnParallelLinear>("in_proj_z", hidden_size, value_dim, false, dtype, device, tp_rank, tp_size);
in_proj_a_ = this->register_module<infinilm::layers::linear::ColumnParallelLinear>("in_proj_a", hidden_size, linear_num_value_heads, false, dtype, device, tp_rank, tp_size);
in_proj_b_ = this->register_module<infinilm::layers::linear::ColumnParallelLinear>("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));
Expand All @@ -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;
Expand Down Expand Up @@ -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
3 changes: 3 additions & 0 deletions csrc/models/qwen3_next/qwen3_next_gated_deltanet.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<layers::linear::MergedColumnParallelLinear> in_proj_qkvz_;
std::shared_ptr<layers::linear::GateUpParallelLinear> in_proj_ba_;
std::shared_ptr<layers::linear::QKVParallelLinear> in_proj_qkv_;
std::shared_ptr<layers::linear::ColumnParallelLinear> in_proj_z_;
std::shared_ptr<layers::linear::ColumnParallelLinear> in_proj_a_;
Expand Down