From c93e0724bb35d8d77ca8e2ca6c8891efcb717917 Mon Sep 17 00:00:00 2001 From: Krzysztof Rymski Date: Mon, 21 Sep 2026 05:10:53 -0700 Subject: [PATCH] Reduce pre allocation of KV caches, and run ZeroInit before running code PiperOrigin-RevId: 985203840 --- gemma/activations.h | 24 ++++++++++++++++++++++++ gemma/kv_cache.cc | 7 +++++-- 2 files changed, 29 insertions(+), 2 deletions(-) diff --git a/gemma/activations.h b/gemma/activations.h index 341875ae..1d2e60d7 100644 --- a/gemma/activations.h +++ b/gemma/activations.h @@ -144,6 +144,30 @@ struct AttentionActivations { flash_params.reserve(batch_size * layer_config.heads); split_flash_params.reserve(batch_size * layer_config.heads); + const size_t kv_out_elems = + batch_size * layer_config.kv_heads * 2 * max_qkv_dim; + kv_out_mem.resize(kv_out_elems, 0.0f); + const size_t total_queries = batch_size * layer_config.heads; + const size_t query_elems = total_queries * max_qkv_dim; + float_queries.resize(query_elems, 0.0f); + bf16_queries.resize(query_elems); + hwy::ZeroBytes(bf16_queries.data(), bf16_queries.size() * sizeof(BF16)); + constexpr size_t kSubtaskQueries = 128; + const size_t num_sub_tasks_init = + hwy::DivCeil(total_queries, kSubtaskQueries) * rep_factor; + const size_t max_q_per_subtask = std::min(total_queries, kSubtaskQueries); + const size_t max_q_rounded_8 = hwy::RoundUpTo(max_q_per_subtask, 8); + sub_task_att_out.resize(num_sub_tasks_init); + sub_task_exp_denominator_sums.resize(num_sub_tasks_init); + sub_task_max_logits.resize(num_sub_tasks_init); + for (size_t t = 0; t < num_sub_tasks_init; ++t) { + sub_task_att_out[t] = MatFactory("att_out", max_q_per_subtask, + max_qkv_dim, allocator); + ZeroInit(sub_task_att_out[t]); + sub_task_exp_denominator_sums[t].resize(max_q_rounded_8, 0.0f); + sub_task_max_logits[t].resize(max_q_rounded_8, 0.0f); + } + // For MatMul outputs, precompute their row pointers. // If we forget any MatMul outputs here, debug builds print a warning but // fill them in each MatMul call. diff --git a/gemma/kv_cache.cc b/gemma/kv_cache.cc index 6362fde3..21e5a531 100644 --- a/gemma/kv_cache.cc +++ b/gemma/kv_cache.cc @@ -343,10 +343,11 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args, size_t local_tile_length = 0; size_t global_tile_length = 0; + const size_t capped_seq_len = CappedSeqLen(config, inference_args); for (size_t i = 0; i < num_layers; ++i) { size_t num_tiles = num_tiles_per_head(kv_attention_window_sizes[i], runtime_config.prefill_tbatch_size, - config.max_seq_len) * + capped_seq_len) * kv_layer_configs[i].kv_heads; size_t tile_len = 2 * kv_layer_configs[i].qkv_dim * kTileSize; @@ -383,6 +384,7 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args, } compact_local_kv_cache.AllocateFor(compact_local_kv_cache_ptr, allocator, MatPadding::kPacked); + ZeroInit(compact_local_kv_cache_ptr); } if (total_global_num_tiles > 0) { @@ -401,6 +403,7 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args, compact_global_kv_cache.AllocateFor(compact_global_kv_cache_ptr, allocator, MatPadding::kPacked); + ZeroInit(compact_global_kv_cache_ptr); } if (compact_global_kv_cache_ptr.HasPtr()) { @@ -427,7 +430,7 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args, for (size_t kv = 0; kv < kv_layer_configs[i].kv_heads; ++kv) { size_t num_tiles_per_kv_head = num_tiles_per_head( kv_attention_window_sizes[i], runtime_config.prefill_tbatch_size, - config.max_seq_len); + CappedSeqLen(config, inference_args)); MatPtr kv_ptr("kv_ptr", kv_cache_type, Extents2D(num_tiles_per_kv_head, layer_tile_length)); if (is_global) {