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
24 changes: 24 additions & 0 deletions gemma/activations.h
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
7 changes: 5 additions & 2 deletions gemma/kv_cache.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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) {
Expand All @@ -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()) {
Expand All @@ -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) {
Expand Down
Loading