Skip to content
Merged
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
15 changes: 13 additions & 2 deletions gemma/gemma4_moe.cc
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@
#include "hwy/highway.h"
// After highway.h
#include "gemma/attention.h" // includes highway.h
#include "gemma/tiled_attention.h"
#include "gemma/gemma-inl.h"
#include "ops/ops-inl.h"

Expand Down Expand Up @@ -454,8 +455,18 @@ void Gemma4MoETransformerLayer(size_t num_tokens, size_t layer_idx,
HWY_DASSERT(layer.layer_config.type == LayerAttentionType::kGemma);
HWY_DASSERT(qbatch.PrefixEnd(0) == 0); // expect causal attention
int flags = 0;
GemmaAttention(num_tokens, kv_cache_layer_idx, layer, activations.attention,
qbatch, env, activations.attention_impl, flags);
if (activations.attention_impl == AttentionImpl::kFlashTransposedQs ||
activations.attention_impl == AttentionImpl::kFlashTransposedQsBF16 ||
activations.attention_impl == AttentionImpl::kFlashTransposedQsInt16 ||
activations.attention_impl == AttentionImpl::kFlashTransposedQsInt8 ||
activations.attention_impl == AttentionImpl::kInt8MatrixAccumulation ||
activations.attention_impl == AttentionImpl::kFlashMatrixAccumulation) {
TiledAttention(activations.attention_impl, num_tokens, kv_cache_layer_idx, layer,
activations.attention, qbatch, env, flags);
} else {
GemmaAttention(num_tokens, kv_cache_layer_idx, layer, activations.attention,
qbatch, env, activations.attention_impl, flags);
}

post_norm(layer.layer_config.post_norm, layer.post_attention_norm_scale,
activations.attention.att_sums);
Expand Down
19 changes: 18 additions & 1 deletion gemma/kv_cache.cc
Original file line number Diff line number Diff line change
Expand Up @@ -78,14 +78,18 @@ KVCache::KVCache(const Extents2D& kv_extents, size_t num_layers,
allocator_(allocator) {
layer_flat_offsets.resize(num_layers, 0);
layer_k_v_offsets.resize(num_layers, 0);
layer_kv_head_offsets.resize(num_layers, 0);
rounded_qkv_dims.resize(num_layers, static_cast<uint32_t>(rounded_qkv_dim));
size_t flat_accum = 0;
size_t k_v_accum = 0;
size_t kv_head_accum = 0;
for (size_t i = 0; i < num_layers; ++i) {
layer_flat_offsets[i] = static_cast<uint32_t>(flat_accum);
flat_accum += 2 * kv_heads * qkv_dim;
layer_k_v_offsets[i] = static_cast<uint32_t>(k_v_accum);
k_v_accum += kv_heads * rounded_qkv_dim;
layer_kv_head_offsets[i] = static_cast<uint32_t>(kv_head_accum);
kv_head_accum += kv_heads;
}
// NOTE: k_v_cols is intentionally left at 0 (default). It serves as a
// sentinel for MaybeReshapeCache: when k_v_cols == cache.Cols(), the reshape
Expand Down Expand Up @@ -140,10 +144,12 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args,
// 1. Build non-uniform offset tables dynamically
layer_flat_offsets.resize(num_layers, 0);
layer_k_v_offsets.resize(num_layers, 0);
layer_kv_head_offsets.resize(num_layers, 0);
rounded_qkv_dims.resize(num_layers, 0);

size_t flat_accum = 0;
size_t k_v_accum = 0;
size_t kv_head_accum = 0;

for (size_t i = 0; i < num_layers; ++i) {
layer_flat_offsets[i] = static_cast<uint32_t>(flat_accum);
Expand All @@ -154,6 +160,9 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args,
hwy::RoundUpTo(kv_layer_configs[i].qkv_dim, kMaxBF16PerVector);
rounded_qkv_dims[i] = static_cast<uint32_t>(rounded_dim);
k_v_accum += kv_layer_configs[i].kv_heads * rounded_dim;

layer_kv_head_offsets[i] = static_cast<uint32_t>(kv_head_accum);
kv_head_accum += config.layer_configs[i].kv_heads;
}
k_v_cols = static_cast<uint32_t>(k_v_accum);

Expand Down Expand Up @@ -195,10 +204,12 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args,
// 1. Build non-uniform offset tables dynamically
layer_flat_offsets.resize(num_layers, 0);
layer_k_v_offsets.resize(num_layers, 0);
layer_kv_head_offsets.resize(num_layers, 0);
rounded_qkv_dims.resize(num_layers, 0);

size_t flat_accum = 0;
size_t k_v_accum = 0;
size_t kv_head_accum = 0;
size_t max_qkv_dim = 0;
size_t max_kv_heads = 0;

Expand All @@ -212,6 +223,9 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args,
rounded_qkv_dims[i] = static_cast<uint32_t>(rounded_dim);
k_v_accum += kv_layer_configs[i].kv_heads * rounded_dim;

layer_kv_head_offsets[i] = static_cast<uint32_t>(kv_head_accum);
kv_head_accum += config.layer_configs[i].kv_heads;

max_qkv_dim = HWY_MAX(max_qkv_dim, kv_layer_configs[i].qkv_dim);
max_kv_heads = HWY_MAX(max_kv_heads, kv_layer_configs[i].kv_heads);
}
Expand Down Expand Up @@ -370,7 +384,7 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args,
size_t local_tiles_processed = 0;
size_t global_tiles_processed = 0;
kv_head_ptrs.clear();
kv_head_ptrs.reserve(num_layers * max_kv_heads);
kv_head_ptrs.reserve(kv_head_accum);
for (size_t i = 0; i < num_layers; ++i) {
size_t layer_tile_length = 2 * kv_layer_configs[i].qkv_dim * kTileSize;
if (kv_cache_type == Type::kInt8) {
Expand Down Expand Up @@ -446,6 +460,9 @@ KVCache KVCache::Copy() {
copy.ds_state_offsets = ds_state_offsets;
}
copy.layer_flat_offsets = layer_flat_offsets;
copy.layer_k_v_offsets = layer_k_v_offsets;
copy.rounded_qkv_dims = rounded_qkv_dims;
copy.layer_kv_head_offsets = layer_kv_head_offsets;
return copy;
}

Expand Down
7 changes: 4 additions & 3 deletions gemma/kv_cache.h
Original file line number Diff line number Diff line change
Expand Up @@ -74,12 +74,12 @@ struct KVCache {
// layers start_pos might be in a middle of the first tile. At start_pos %
// kTileSize
std::vector<MatPtr> GetPointers(int layer_idx, int kv_head_idx,
int num_kv_heads, int start_pos,
int start_pos,
bool is_global_layer) {
if (!IsTiled()) {
HWY_ABORT("This function is only meant to be used with tiled KV caches.");
}
MatPtr& source_ptr = kv_head_ptrs[layer_idx * num_kv_heads + kv_head_idx];
MatPtr& source_ptr = kv_head_ptrs[layer_kv_head_offsets[layer_idx] + kv_head_idx];
if (is_global_layer) {
return {source_ptr};
}
Expand Down Expand Up @@ -137,6 +137,7 @@ struct KVCache {
std::vector<uint32_t> layer_flat_offsets;
std::vector<uint32_t> layer_k_v_offsets;
std::vector<uint32_t> rounded_qkv_dims;
std::vector<uint32_t> layer_kv_head_offsets;

// DeepSeek V4 per-query incremental compressor state (kv_state/score_state
// per layer, plus the indexer compressor's on CSA layers), f32. One row;
Expand Down Expand Up @@ -189,7 +190,7 @@ struct KVCache {
// number of tiles in storage. All pointers point into compact_kv_cache.

// To access the tiles of (layer_idx, head_idx), index the array with
// layer_idx * num_kv_heads + kv_head_idx.
// layer_kv_head_offsets[layer_idx] + kv_head_idx.
// Or use GetPointers function.

// The returned MatPtr will have one tile per row. The number of rows for
Expand Down
11 changes: 6 additions & 5 deletions gemma/tiled_attention.cc
Original file line number Diff line number Diff line change
Expand Up @@ -115,9 +115,6 @@ static HWY_INLINE void ComputeQKVTransposedTile(
const bool skip_kv =
(layer_config.kv_share_layer_idx >= 0) || (flags & kSkipKV);

// The original qkv_einsum_w has shape [(heads + kv_heads * 2), qkv_dim,
// model_dim], which we reshaped to (heads + kv_heads * 2) * qkv_dim rows.
// This computes Q and stores it in activations.q.
// The original qkv_einsum_w has shape [(heads + kv_heads * 2), qkv_dim,
// model_dim], which we reshaped to (heads + kv_heads * 2) * qkv_dim rows.
// This computes Q and stores it in activations.q.
Expand Down Expand Up @@ -162,7 +159,7 @@ static HWY_INLINE void ComputeQKVTransposedTile(
const bool is_global_layer =
activations.config.IsGlobalLayer(layer_idx);
std::vector<MatPtr> kv_ptrs = qbatch.KV(query_idx).cache->GetPointers(
kv_layer_idx, kv_head, kv_heads, start_pos, is_global_layer);
kv_layer_idx, kv_head, start_pos, is_global_layer);
const size_t v_offset = qkv_dim * KVCache::kTileSize;
const size_t tile_span_size = 2 * qkv_dim * KVCache::kTileSize;
const size_t k_size = qkv_dim * KVCache::kTileSize;
Expand Down Expand Up @@ -939,7 +936,7 @@ void LocalAttentionForAllHeadsTokensAndBatch(
std::vector<MatPtr> kv_ptrs =
qbatch.KV(current_qbatch_idx)
.cache->GetPointers(
layer_idx, kv_head_idx, layer.layer_config.kv_heads,
layer_idx, kv_head_idx,
global_start_context_pos,
activations.config.IsGlobalLayer(layer_idx));

Expand Down Expand Up @@ -1150,6 +1147,10 @@ void TiledAttention(AttentionImpl attention_impl, size_t num_tokens,
"query heads must be a multiple of key-value heads");
(void)layer_config; // only used in HWY_DASSERT

const size_t active_qkv_dim = layer_config.heads * layer_config.qkv_dim;
activations.q.OverrideCols(active_qkv_dim);
activations.att_out.OverrideCols(active_qkv_dim);

const Type kv_type = qbatch.KV(0).cache->compact_kv_cache_ptr.GetType();
if (kv_type == Type::kBF16) {
ComputeQKVTransposedTile<BF16>(num_tokens, layer_idx, layer, attention_impl,
Expand Down
Loading