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
Original file line number Diff line number Diff line change
@@ -0,0 +1,138 @@
#include "granitemoehybrid_allocate_kv_cache_tensors.hpp"

#include "../../global_state/global_state.hpp"

#include "infinicore/context/context.hpp"

#include <algorithm>
#include <stdexcept>
#include <string>
#include <utility>
#include <vector>

namespace infinilm::models::granitemoehybrid {

GraniteMoeHybridAllocatedCache granitemoehybrid_allocate_cache_tensors(
const cache::CacheConfig *cache_config,
const std::shared_ptr<infinilm::config::ModelConfig> &text_config,
const backends::AttentionBackend &attention_backend) {
if (nullptr == cache_config) {
return {};
}
if (nullptr == text_config) {
throw std::runtime_error("infinilm::models::granitemoehybrid::granitemoehybrid_allocate_cache_tensors: text_config is null");
}

const size_t num_hidden_layers = text_config->get<size_t>("num_hidden_layers");
const size_t head_dim = text_config->get_head_dim();
const size_t num_key_value_heads = text_config->get<size_t>("num_key_value_heads");

const size_t hidden_size = text_config->get<size_t>("hidden_size");
const size_t mamba_expand = text_config->get_or<size_t>("mamba_expand", 2);
const size_t mamba_n_groups = text_config->get_or<size_t>("mamba_n_groups", 1);
const size_t mamba_d_state = text_config->get<size_t>("mamba_d_state");
const size_t mamba_d_conv = text_config->get<size_t>("mamba_d_conv");

const auto &dtype = text_config->get_dtype();
const auto &kv_cache_dtype = text_config->get_kv_cache_dtype();
const std::vector<std::string> layer_types = text_config->get<std::vector<std::string>>("layer_types");

std::vector<infinicore::Tensor> kv_cache_vec;
std::vector<infinicore::Tensor> conv_state_vec;
kv_cache_vec.reserve(num_hidden_layers);
conv_state_vec.reserve(num_hidden_layers);

const auto &rank_info = infinilm::global_state::get_tensor_model_parallel_rank_info();
const size_t local_mamba_groups = mamba_n_groups >= static_cast<size_t>(rank_info.tp_size)
? mamba_n_groups / rank_info.tp_size
: 1;
const size_t mamba_conv_dim = mamba_expand * hidden_size / rank_info.tp_size + 2 * local_mamba_groups * mamba_d_state;

auto allocate_mamba_cache = [&](size_t pool_size) {
auto conv_state = infinicore::Tensor::zeros(
{pool_size, mamba_conv_dim, mamba_d_conv - 1},
dtype,
rank_info.device);
kv_cache_vec.emplace_back();
conv_state_vec.push_back(std::move(conv_state));
};

auto allocate_static_attention_cache = [&](const cache::StaticKVCacheConfig &config) {
auto kv_cache = cache::StaticKVCache::create_layer_kv_cache(
head_dim,
head_dim,
num_key_value_heads,
num_key_value_heads,
text_config->get<size_t>("max_position_embeddings"),
kv_cache_dtype,
config);

kv_cache_vec.push_back(std::move(kv_cache));
conv_state_vec.emplace_back();
};

auto allocate_paged_attention_cache = [&](const cache::PagedKVCacheConfig &config) {
auto kv_cache = cache::PagedKVCache::create_layer_kv_cache(
head_dim,
head_dim,
num_key_value_heads,
num_key_value_heads,
kv_cache_dtype,
config);

kv_cache_vec.push_back(std::move(kv_cache));
conv_state_vec.emplace_back();
};

switch (attention_backend) {
case backends::AttentionBackend::STATIC_ATTN: {
const auto *static_kv_cache_config = dynamic_cast<const cache::StaticKVCacheConfig *>(cache_config);
if (nullptr == static_kv_cache_config) {
throw std::runtime_error("infinilm::models::granitemoehybrid::granitemoehybrid_allocate_cache_tensors: invalid static kv cache config type");
}

const size_t mamba_pool_size = std::max<size_t>(2, static_kv_cache_config->max_batch_size() + 1);
for (size_t layer_idx = 0; layer_idx < num_hidden_layers; ++layer_idx) {
const std::string &layer_type = layer_types[layer_idx];
if ("mamba" == layer_type) {
allocate_mamba_cache(mamba_pool_size);
} else if ("attention" == layer_type) {
allocate_static_attention_cache(*static_kv_cache_config);
} else {
throw std::runtime_error("infinilm::models::granitemoehybrid::granitemoehybrid_allocate_cache_tensors: unsupported layer_type '" + layer_type + "' for layer " + std::to_string(layer_idx));
}
}
break;
}
case backends::AttentionBackend::FLASH_ATTN: {
;
}
case backends::AttentionBackend::PAGED_ATTN: {
const auto *paged_kv_cache_config = dynamic_cast<const cache::PagedKVCacheConfig *>(cache_config);
if (nullptr == paged_kv_cache_config) {
throw std::runtime_error("infinilm::models::granitemoehybrid::granitemoehybrid_allocate_cache_tensors: invalid paged kv cache config type");
}
const size_t mamba_pool_size = std::max<size_t>(2, paged_kv_cache_config->num_blocks() / 4);

for (size_t layer_idx = 0; layer_idx < num_hidden_layers; ++layer_idx) {
const std::string &layer_type = layer_types[layer_idx];
if ("mamba" == layer_type) {
allocate_mamba_cache(mamba_pool_size);
} else if ("attention" == layer_type) {
allocate_paged_attention_cache(*paged_kv_cache_config);
} else {
throw std::runtime_error("infinilm::models::granitemoehybrid::granitemoehybrid_allocate_cache_tensors: unsupported layer_type '" + layer_type + "' for layer " + std::to_string(layer_idx));
}
}
break;
}
default:
throw std::runtime_error("infinilm::models::granitemoehybrid::granitemoehybrid_allocate_cache_tensors: unsupported attention backend " + std::to_string(static_cast<int>(attention_backend)));
}
infinicore::context::syncStream();
return GraniteMoeHybridAllocatedCache{
std::move(kv_cache_vec),
std::move(conv_state_vec)};
}

} // namespace infinilm::models::granitemoehybrid
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
#pragma once

#include "../../backends/attention_backends.hpp"
#include "../../cache/kv_cache.hpp"
#include "../../config/model_config.hpp"

#include <memory>
#include <vector>

namespace infinilm::models::granitemoehybrid {

struct GraniteMoeHybridAllocatedCache {
std::vector<infinicore::Tensor> kv_cache_tensors;
std::vector<infinicore::Tensor> conv_state_tensors;
};

GraniteMoeHybridAllocatedCache granitemoehybrid_allocate_cache_tensors(
const cache::CacheConfig *cache_config,
const std::shared_ptr<infinilm::config::ModelConfig> &text_config,
const backends::AttentionBackend &attention_backend);

} // namespace infinilm::models::granitemoehybrid
177 changes: 177 additions & 0 deletions csrc/models/granitemoehybrid/granitemoehybrid_attention.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,177 @@
#include "granitemoehybrid_attention.hpp"

#include "../../global_state/global_state.hpp"
#include "../../layers/attention/attention.hpp"
#include "../../layers/rotary_embedding/rotary_embedding.hpp"
#include "../../utils.hpp"

#include <stdexcept>
#include <string>
#include <utility>

namespace infinilm::models::granitemoehybrid {

GraniteMoeHybridAttention::GraniteMoeHybridAttention(
std::shared_ptr<infinilm::config::ModelConfig> model_config,
size_t layer_idx,
const infinicore::Device &device) {
layer_idx_ = layer_idx;
hidden_size_ = model_config->get<size_t>("hidden_size");
head_dim_ = model_config->get_head_dim();

const auto &dtype{model_config->get_dtype()};
const size_t total_num_heads = model_config->get<size_t>("num_attention_heads");
const size_t total_num_kv_heads = model_config->get<size_t>("num_key_value_heads");
const bool use_bias = model_config->get_or<bool>("attention_bias", false);
const bool use_output_bias = model_config->get_or<bool>("attention_output_bias", use_bias);

attention_backend_ = infinilm::global_state::get_infinilm_config().attention_backend;
const engine::distributed::RankInfo &rank_info = infinilm::global_state::get_tensor_model_parallel_rank_info();
const int tp_rank = rank_info.tp_rank;
const int tp_size = rank_info.tp_size;
if (tp_size <= 0 || total_num_heads % tp_size != 0) {
throw std::runtime_error(
"infinilm::models::granitemoehybrid::GraniteMoeHybridAttention: "
"num_attention_heads must be divisible by tp_size");
}
if (total_num_kv_heads < static_cast<size_t>(tp_size) || total_num_kv_heads % tp_size != 0) {
throw std::runtime_error(
"infinilm::models::granitemoehybrid::GraniteMoeHybridAttention: "
"num_key_value_heads must be divisible by tp_size");
}

num_attention_heads_ = total_num_heads / tp_size;
num_key_value_heads_ = total_num_kv_heads / tp_size;

auto quantization_method = model_config->get_quantization_method();
auto register_fn = [this](const std::string &name, infinicore::nn::Parameter parameter) {
this->register_parameter(name, std::move(parameter));
};

qkv_proj_ = std::make_shared<infinilm::layers::linear::QKVParallelLinear>(
hidden_size_,
head_dim_,
total_num_heads,
total_num_kv_heads,
"q_proj",
"k_proj",
"v_proj",
register_fn,
quantization_method,
use_bias,
dtype,
device,
rank_info);
o_proj_ = this->register_module<infinilm::layers::linear::RowParallelLinear>(
"o_proj",
total_num_heads * head_dim_,
hidden_size_,
quantization_method,
use_output_bias,
dtype,
device,
tp_rank,
tp_size,
rank_info.comm);
o_proj_->set_alpha(model_config->get_or<float>("residual_multiplier", 1.0f));

const std::string position_embedding_type = model_config->get_or<std::string>("position_embedding_type", "nope");
if ("rope" == position_embedding_type) {
rotary_emb_ = infinilm::layers::rotary_embedding::get_rope(model_config, device);
}

const float attention_multiplier = model_config->get_or<float>("attention_multiplier", 1.0f);
infinilm::layers::attention::init_kv_cache_quant_params(
register_fn,
device,
kv_cache_k_scale_,
kv_cache_v_scale_);
attn_ = std::make_shared<infinilm::layers::attention::AttentionLayer>(
num_attention_heads_,
head_dim_,
attention_multiplier,
num_key_value_heads_,
layer_idx_,
kv_cache_k_scale_,
kv_cache_v_scale_,
attention_backend_);
}

infinicore::Tensor GraniteMoeHybridAttention::forward(
const infinicore::Tensor &positions,
const infinicore::Tensor &hidden_states) const {
if (::infinilm::backends::AttentionBackend::STATIC_ATTN == attention_backend_) {
return forward_static_(positions, hidden_states);
}
return forward_paged_(positions, hidden_states);
}

infinicore::Tensor GraniteMoeHybridAttention::forward_static_(
const infinicore::Tensor &position_ids,
const infinicore::Tensor &hidden_states) const {
auto hidden_states_mutable = hidden_states;
const auto &shape = hidden_states->shape();
const size_t batch_size = shape[0];
const size_t seq_len = shape[1];

auto [query, key, value] = qkv_proj_->forward_split(hidden_states_mutable);
query = query->view({batch_size, seq_len, num_attention_heads_, head_dim_});
key = key->view({batch_size, seq_len, num_key_value_heads_, head_dim_});
value = value->view({batch_size, seq_len, num_key_value_heads_, head_dim_});

if (rotary_emb_) {
const auto &position_shape = position_ids->shape();
infinicore::Tensor rope_positions;
if (position_shape.size() == 2) {
rope_positions = position_ids->narrow({{0, 0, 1}})->contiguous()->view({position_shape[1]});
} else if (position_shape.size() == 1) {
rope_positions = position_ids->contiguous();
} else {
throw std::runtime_error(
"infinilm::models::granitemoehybrid::GraniteMoeHybridAttention: "
"unexpected position_ids shape");
}

rotary_emb_->forward(query, rope_positions, true);
rotary_emb_->forward(key, rope_positions, true);
}

auto attention_output = attn_->forward(query, key, value);
return o_proj_->forward(attention_output);
}

infinicore::Tensor GraniteMoeHybridAttention::forward_paged_(
const infinicore::Tensor &position_ids,
const infinicore::Tensor &hidden_states) const {
auto hidden_states_mutable = hidden_states;
const auto &shape = hidden_states->shape();
const size_t batch_size = shape[0];
const size_t seq_len = shape[1];
ASSERT_EQ(batch_size, 1);

auto [query, key, value] = qkv_proj_->forward_split(hidden_states_mutable);
query = query->view({seq_len, num_attention_heads_, head_dim_});
key = key->view({seq_len, num_key_value_heads_, head_dim_});
value = value->view({seq_len, num_key_value_heads_, head_dim_});

if (rotary_emb_) {
const auto &position_shape = position_ids->shape();
infinicore::Tensor rope_positions;
if (position_shape.size() == 2) {
rope_positions = position_ids->narrow({{0, 0, 1}})->view({position_shape[1]});
} else if (position_shape.size() == 1) {
rope_positions = position_ids;
} else {
throw std::runtime_error(
"infinilm::models::granitemoehybrid::GraniteMoeHybridAttention: "
"unexpected position_ids shape");
}
rotary_emb_->forward(query, rope_positions, true);
rotary_emb_->forward(key, rope_positions, true);
}

auto attention_output = attn_->forward(query, key, value);
return o_proj_->forward(attention_output);
}

} // namespace infinilm::models::granitemoehybrid
56 changes: 56 additions & 0 deletions csrc/models/granitemoehybrid/granitemoehybrid_attention.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,56 @@
#pragma once

#include "../../layers/common_modules.hpp"

#include <memory>

namespace infinilm::models::granitemoehybrid {

class GraniteMoeHybridAttention : public infinicore::nn::Module {
public:
GraniteMoeHybridAttention(std::shared_ptr<infinilm::config::ModelConfig> model_config,
size_t layer_idx,
const infinicore::Device &device);

infinicore::Tensor forward(const infinicore::Tensor &positions,
const infinicore::Tensor &hidden_states) const;

void process_weights_after_loading() override {
qkv_proj_->process_weights_after_loading();
}

void reset_runtime_state() const override {
qkv_proj_->reset_runtime_state();
}

size_t layer_idx() const { return layer_idx_; }
size_t num_heads() const { return num_attention_heads_; }
size_t num_kv_heads() const { return num_key_value_heads_; }
size_t head_dim() const { return head_dim_; }
size_t hidden_size() const { return hidden_size_; }

private:
infinicore::Tensor forward_static_(const infinicore::Tensor &positions,
const infinicore::Tensor &hidden_states) const;

infinicore::Tensor forward_paged_(const infinicore::Tensor &positions,
const infinicore::Tensor &hidden_states) const;

protected:
std::shared_ptr<infinilm::layers::linear::QKVParallelLinear> qkv_proj_;
std::shared_ptr<infinilm::layers::linear::RowParallelLinear> o_proj_;
std::shared_ptr<infinicore::nn::RoPE> rotary_emb_;
std::shared_ptr<infinilm::layers::attention::AttentionLayer> attn_;
::infinilm::backends::AttentionBackend attention_backend_;

size_t layer_idx_;
size_t num_attention_heads_;
size_t num_key_value_heads_;
size_t hidden_size_;
size_t head_dim_;

INFINICORE_NN_PARAMETER(kv_cache_k_scale);
INFINICORE_NN_PARAMETER(kv_cache_v_scale);
};

} // namespace infinilm::models::granitemoehybrid
Loading