diff --git a/BUILD.bazel b/BUILD.bazel index 60bc61d6..7d1bf8c9 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -530,6 +530,9 @@ cc_library( cc_library( name = "ops", + srcs = [ + "ops/ops.cc", + ], hdrs = [ "ops/ops.h", ], diff --git a/gemma/gemma.cc b/gemma/gemma.cc index 620dbb79..4012e637 100644 --- a/gemma/gemma.cc +++ b/gemma/gemma.cc @@ -1188,7 +1188,8 @@ static HWY_NOINLINE void PrefillTBatch(const ModelConfig& config, max_tbatch_size >= prefill_this_query); // For each batch of tokens in the query: - for (size_t tbatch_start = 0; tbatch_start < prefill_this_query; + for (size_t tbatch_start = qbatch_1.InitialPos(0); + tbatch_start < prefill_this_query; tbatch_start += max_tbatch_size) { const size_t tbatch_size = HWY_MIN(max_tbatch_size, prefill_this_query - tbatch_start); @@ -1528,6 +1529,7 @@ size_t PrefillTBatchOrQBatch(const ModelConfig& config, MatMulEnv& env, TimingInfo& timing_info) { size_t max_prompt_size = 0; bool all_prefix_end_are_zero = true; + bool all_initial_pos_are_zero = true; size_t total_prefill_tokens = 0; // only for throughput stats. const size_t seq_len = qbatch.KV(0).SeqLen(); for (size_t qi = 0; qi < qbatch.Size(); ++qi) { @@ -1540,9 +1542,13 @@ size_t PrefillTBatchOrQBatch(const ModelConfig& config, // Prefill stops before size - 1 because the last prompt token is the // first input token for generation. - total_prefill_tokens += prompt.size() - 1; + const size_t initial_pos = qbatch.InitialPos(qi); + if (prompt.size() - 1 > initial_pos) { + total_prefill_tokens += (prompt.size() - 1 - initial_pos); + } all_prefix_end_are_zero &= qbatch.PrefixEnd(qi) == 0; + all_initial_pos_are_zero &= initial_pos == 0; // We use a single divisor, so all sequence lengths must be the same. HWY_ASSERT(qbatch.KV(qi).SeqLen() == seq_len); @@ -1559,7 +1565,8 @@ size_t PrefillTBatchOrQBatch(const ModelConfig& config, timing_info.prefill_start = hwy::platform::Now(); // Batch over the larger of prompt length, or queries. - if ((qbatch.Size() > max_prompt_size) && all_prefix_end_are_zero) { + if ((qbatch.Size() > max_prompt_size) && all_prefix_end_are_zero && + all_initial_pos_are_zero) { activations.SetBatchSize(qbatch.Size()); // required before PrefillQBatch PrefillQBatch(max_prompt_size, config, runtime_config, weights, activations, qbatch, env, non_eos); @@ -1586,7 +1593,9 @@ void StreamAndUpdateEOSAfterPrefill(const ModelConfig& config, const RuntimeConfig& runtime_config, QBatch& qbatch, hwy::BitSet4096<>& non_eos, size_t qi) { - const size_t last_pos_in_prompt = qbatch.Pos(qi) - qbatch.InitialPos(qi); + const size_t prompt_size = qbatch.Prompt(qi).size(); + const size_t last_pos_in_prompt = + prompt_size == 0 ? 0 : std::min(qbatch.Pos(qi), prompt_size - 1); const size_t pos = qbatch.Pos(qi); // during prefill, pos is still correct. // In autoregressive mode, we have not prefilled the last token, so do diff --git a/gemma/kv_cache.cc b/gemma/kv_cache.cc index 6362fde3..5e48726b 100644 --- a/gemma/kv_cache.cc +++ b/gemma/kv_cache.cc @@ -54,7 +54,8 @@ static const std::vector& KVAttentionWindowSizes( } KVCache::KVCache(const Extents2D& kv_extents, size_t num_layers, - size_t kv_heads, size_t qkv_dim, const Allocator& allocator) + size_t kv_heads, size_t qkv_dim, size_t k_v_cols, + const Allocator& allocator) : num_layers(num_layers), kv_heads(kv_heads), qkv_dim(qkv_dim), @@ -69,11 +70,11 @@ KVCache::KVCache(const Extents2D& kv_extents, size_t num_layers, // The change is shape is safe only if the padding is kPacked. k_cache("k", Extents2D(hwy::RoundUpTo(kv_extents.rows, kMaxBF16PerVector), - KOrVDefaultCols()), + k_v_cols != 0 ? k_v_cols : KOrVDefaultCols()), allocator, MatPadding::kPacked), v_cache("v", Extents2D(hwy::RoundUpTo(kv_extents.rows, kMaxBF16PerVector), - KOrVDefaultCols()), + k_v_cols != 0 ? k_v_cols : KOrVDefaultCols()), allocator, MatPadding::kPacked), allocator_(allocator) { layer_flat_offsets.resize(num_layers, 0); @@ -399,8 +400,7 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args, MatPtr::Layout::kInt8MatrixAccumulation); } compact_global_kv_cache.AllocateFor(compact_global_kv_cache_ptr, - allocator, - MatPadding::kPacked); + allocator, MatPadding::kPacked); } if (compact_global_kv_cache_ptr.HasPtr()) { @@ -467,19 +467,135 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args, InitDSState(config, allocator, ds_state, ds_state_snapshot, ds_state_offsets); } -KVCache KVCache::Copy() { - KVCache copy(kv_cache.Extents(), num_layers, kv_heads, qkv_dim, allocator_); +KVCache KVCache::Copy() const { + KVCache copy(kv_cache.Extents(), num_layers, kv_heads, qkv_dim, k_v_cols, + allocator_); + + if (kv_cache.HasPtr() && copy.kv_cache.HasPtr()) { + CopyMat(kv_cache, copy.kv_cache); + } + if (k_cache.HasPtr() && copy.k_cache.HasPtr()) { + if (copy.k_cache.Cols() != k_cache.Cols() && + copy.k_cache.Rows() != k_cache.Rows()) { + const size_t factor = k_cache.Cols() / copy.k_cache.Cols(); + if (factor > 0 && copy.k_cache.Cols() * factor == k_cache.Cols()) { + copy.k_cache.ReshapePackedRowsToCols(factor); + } + } + if (copy.k_cache.SameShape(k_cache)) { + CopyMat(k_cache, copy.k_cache); + } + } + if (v_cache.HasPtr() && copy.v_cache.HasPtr()) { + if (copy.v_cache.Cols() != v_cache.Cols() && + copy.v_cache.Rows() != v_cache.Rows()) { + const size_t factor = v_cache.Cols() / copy.v_cache.Cols(); + if (factor > 0 && copy.v_cache.Cols() * factor == v_cache.Cols()) { + copy.v_cache.ReshapePackedRowsToCols(factor); + } + } + if (copy.v_cache.SameShape(v_cache)) { + CopyMat(v_cache, copy.v_cache); + } + } - CopyMat(kv_cache, copy.kv_cache); if (compact_local_kv_cache_ptr.HasPtr()) { + copy.compact_local_kv_cache_ptr = MatPtr( + compact_local_kv_cache_ptr.Name(), compact_local_kv_cache_ptr.GetType(), + compact_local_kv_cache_ptr.Extents()); + copy.compact_local_kv_cache_ptr.SetLayout( + compact_local_kv_cache_ptr.GetLayout()); + copy.compact_local_kv_cache.AllocateFor(copy.compact_local_kv_cache_ptr, + allocator_, MatPadding::kPacked); CopyMat(compact_local_kv_cache_ptr, copy.compact_local_kv_cache_ptr); } + if (compact_global_kv_cache_ptr.HasPtr()) { + copy.compact_global_kv_cache_ptr = + MatPtr(compact_global_kv_cache_ptr.Name(), + compact_global_kv_cache_ptr.GetType(), + compact_global_kv_cache_ptr.Extents()); + copy.compact_global_kv_cache_ptr.SetLayout( + compact_global_kv_cache_ptr.GetLayout()); + copy.compact_global_kv_cache.AllocateFor(copy.compact_global_kv_cache_ptr, + allocator_, MatPadding::kPacked); CopyMat(compact_global_kv_cache_ptr, copy.compact_global_kv_cache_ptr); } - copy.compact_kv_cache_ptr = compact_global_kv_cache_ptr.HasPtr() - ? copy.compact_global_kv_cache_ptr - : copy.compact_local_kv_cache_ptr; + + if (!compact_local_kv_cache_ptr.HasPtr() && + !compact_global_kv_cache_ptr.HasPtr() && compact_kv_cache_ptr.HasPtr()) { + copy.compact_kv_cache_ptr = + MatPtr(compact_kv_cache_ptr.Name(), compact_kv_cache_ptr.GetType(), + compact_kv_cache_ptr.Extents()); + copy.compact_kv_cache_ptr.SetLayout(compact_kv_cache_ptr.GetLayout()); + copy.compact_kv_cache.AllocateFor(copy.compact_kv_cache_ptr, allocator_, + MatPadding::kPacked); + CopyMat(compact_kv_cache_ptr, copy.compact_kv_cache_ptr); + } else { + copy.compact_kv_cache_ptr = copy.compact_global_kv_cache_ptr.HasPtr() + ? copy.compact_global_kv_cache_ptr + : copy.compact_local_kv_cache_ptr; + } + + // Re-base all MatPtr pointers in kv_head_ptrs + copy.kv_head_ptrs.clear(); + copy.kv_head_ptrs.reserve(kv_head_ptrs.size()); + + const uint8_t* orig_local_base = (compact_local_kv_cache_ptr.HasPtr() && + compact_local_kv_cache_ptr.Rows() > 0) + ? compact_local_kv_cache_ptr.RowBytes(0) + : nullptr; + const uint8_t* orig_global_base = + (compact_global_kv_cache_ptr.HasPtr() && + compact_global_kv_cache_ptr.Rows() > 0) + ? compact_global_kv_cache_ptr.RowBytes(0) + : nullptr; + const uint8_t* orig_compact_base = + (compact_kv_cache_ptr.HasPtr() && compact_kv_cache_ptr.Rows() > 0) + ? compact_kv_cache_ptr.RowBytes(0) + : nullptr; + + for (const MatPtr& orig_ptr : kv_head_ptrs) { + MatPtr new_ptr(orig_ptr.Name(), orig_ptr.GetType(), orig_ptr.Extents()); + new_ptr.SetLayout(orig_ptr.GetLayout()); + if (orig_ptr.HasPtr() && orig_ptr.Rows() > 0) { + const uint8_t* orig_addr = orig_ptr.RowBytes(0); + uint8_t* new_base = nullptr; + size_t new_stride = 0; + if (orig_local_base != nullptr && orig_addr >= orig_local_base && + orig_addr < + orig_local_base + compact_local_kv_cache_ptr.Rows() * + compact_local_kv_cache_ptr.Stride() * + compact_local_kv_cache_ptr.ElementBytes()) { + const uintptr_t offset = orig_addr - orig_local_base; + new_base = copy.compact_local_kv_cache_ptr.RowBytes(0) + offset; + new_stride = copy.compact_local_kv_cache_ptr.Stride(); + } else if (orig_global_base != nullptr && orig_addr >= orig_global_base && + orig_addr < + orig_global_base + + compact_global_kv_cache_ptr.Rows() * + compact_global_kv_cache_ptr.Stride() * + compact_global_kv_cache_ptr.ElementBytes()) { + const uintptr_t offset = orig_addr - orig_global_base; + new_base = copy.compact_global_kv_cache_ptr.RowBytes(0) + offset; + new_stride = copy.compact_global_kv_cache_ptr.Stride(); + } else if (orig_compact_base != nullptr && + orig_addr >= orig_compact_base && + orig_addr < orig_compact_base + + compact_kv_cache_ptr.Rows() * + compact_kv_cache_ptr.Stride() * + compact_kv_cache_ptr.ElementBytes()) { + const uintptr_t offset = orig_addr - orig_compact_base; + new_base = copy.compact_kv_cache_ptr.RowBytes(0) + offset; + new_stride = copy.compact_kv_cache_ptr.Stride(); + } else { + HWY_ABORT("Unrecognized buffer in kv_head_ptrs during Copy()"); + } + new_ptr.SetPtr(new_base, new_stride); + } + copy.kv_head_ptrs.push_back(std::move(new_ptr)); + } + copy.tiled_seq_len = tiled_seq_len; if (ds_state.Rows() > 0) { copy.ds_state = MatStorageT("ds_state", ds_state.Extents(), @@ -491,6 +607,7 @@ KVCache KVCache::Copy() { CopyMat(ds_state_snapshot, copy.ds_state_snapshot); copy.ds_state_offsets = ds_state_offsets; } + copy.k_v_cols = k_v_cols; copy.layer_flat_offsets = layer_flat_offsets; copy.layer_k_v_offsets = layer_k_v_offsets; copy.rounded_qkv_dims = rounded_qkv_dims; @@ -498,4 +615,252 @@ KVCache KVCache::Copy() { return copy; } +KVCache KVCache::Copy() { return static_cast(this)->Copy(); } + +KVCache KVCache::ClonePrefix(size_t prefix_len, + const Allocator& allocator) const { + KVCache clone(kv_cache.Extents(), num_layers, kv_heads, qkv_dim, k_v_cols, + allocator); + + clone.tiled_seq_len = tiled_seq_len; + clone.k_v_cols = k_v_cols; + clone.layer_flat_offsets = layer_flat_offsets; + clone.layer_k_v_offsets = layer_k_v_offsets; + clone.rounded_qkv_dims = rounded_qkv_dims; + clone.layer_kv_head_offsets = layer_kv_head_offsets; + + if (compact_local_kv_cache_ptr.HasPtr()) { + clone.compact_local_kv_cache_ptr = MatPtr( + compact_local_kv_cache_ptr.Name(), compact_local_kv_cache_ptr.GetType(), + compact_local_kv_cache_ptr.Extents()); + clone.compact_local_kv_cache_ptr.SetLayout( + compact_local_kv_cache_ptr.GetLayout()); + clone.compact_local_kv_cache.AllocateFor(clone.compact_local_kv_cache_ptr, + allocator, MatPadding::kPacked); + ZeroInit(clone.compact_local_kv_cache_ptr); + } + + if (compact_global_kv_cache_ptr.HasPtr()) { + clone.compact_global_kv_cache_ptr = + MatPtr(compact_global_kv_cache_ptr.Name(), + compact_global_kv_cache_ptr.GetType(), + compact_global_kv_cache_ptr.Extents()); + clone.compact_global_kv_cache_ptr.SetLayout( + compact_global_kv_cache_ptr.GetLayout()); + clone.compact_global_kv_cache.AllocateFor(clone.compact_global_kv_cache_ptr, + allocator, MatPadding::kPacked); + ZeroInit(clone.compact_global_kv_cache_ptr); + } + + if (!compact_local_kv_cache_ptr.HasPtr() && + !compact_global_kv_cache_ptr.HasPtr() && compact_kv_cache_ptr.HasPtr()) { + clone.compact_kv_cache_ptr = + MatPtr(compact_kv_cache_ptr.Name(), compact_kv_cache_ptr.GetType(), + compact_kv_cache_ptr.Extents()); + clone.compact_kv_cache_ptr.SetLayout(compact_kv_cache_ptr.GetLayout()); + clone.compact_kv_cache.AllocateFor(clone.compact_kv_cache_ptr, allocator, + MatPadding::kPacked); + ZeroInit(clone.compact_kv_cache_ptr); + } else { + clone.compact_kv_cache_ptr = clone.compact_global_kv_cache_ptr.HasPtr() + ? clone.compact_global_kv_cache_ptr + : clone.compact_local_kv_cache_ptr; + } + + // Re-base all MatPtr pointers in kv_head_ptrs and copy prefix tiles + clone.kv_head_ptrs.clear(); + clone.kv_head_ptrs.reserve(kv_head_ptrs.size()); + + const uint8_t* orig_local_base = (compact_local_kv_cache_ptr.HasPtr() && + compact_local_kv_cache_ptr.Rows() > 0) + ? compact_local_kv_cache_ptr.RowBytes(0) + : nullptr; + const uint8_t* orig_global_base = + (compact_global_kv_cache_ptr.HasPtr() && + compact_global_kv_cache_ptr.Rows() > 0) + ? compact_global_kv_cache_ptr.RowBytes(0) + : nullptr; + const uint8_t* orig_compact_base = + (compact_kv_cache_ptr.HasPtr() && compact_kv_cache_ptr.Rows() > 0) + ? compact_kv_cache_ptr.RowBytes(0) + : nullptr; + + const size_t prefix_tiles = hwy::DivCeil(prefix_len, kTileSize); + + for (const MatPtr& orig_ptr : kv_head_ptrs) { + MatPtr new_ptr(orig_ptr.Name(), orig_ptr.GetType(), orig_ptr.Extents()); + new_ptr.SetLayout(orig_ptr.GetLayout()); + if (orig_ptr.HasPtr() && orig_ptr.Rows() > 0) { + const uint8_t* orig_addr = orig_ptr.RowBytes(0); + uint8_t* new_base = nullptr; + size_t new_stride = 0; + if (orig_local_base != nullptr && orig_addr >= orig_local_base && + orig_addr < + orig_local_base + compact_local_kv_cache_ptr.Rows() * + compact_local_kv_cache_ptr.Stride() * + compact_local_kv_cache_ptr.ElementBytes()) { + const uintptr_t offset = orig_addr - orig_local_base; + new_base = clone.compact_local_kv_cache_ptr.RowBytes(0) + offset; + new_stride = clone.compact_local_kv_cache_ptr.Stride(); + } else if (orig_global_base != nullptr && orig_addr >= orig_global_base && + orig_addr < + orig_global_base + + compact_global_kv_cache_ptr.Rows() * + compact_global_kv_cache_ptr.Stride() * + compact_global_kv_cache_ptr.ElementBytes()) { + const uintptr_t offset = orig_addr - orig_global_base; + new_base = clone.compact_global_kv_cache_ptr.RowBytes(0) + offset; + new_stride = clone.compact_global_kv_cache_ptr.Stride(); + } else if (orig_compact_base != nullptr && + orig_addr >= orig_compact_base && + orig_addr < orig_compact_base + + compact_kv_cache_ptr.Rows() * + compact_kv_cache_ptr.Stride() * + compact_kv_cache_ptr.ElementBytes()) { + const uintptr_t offset = orig_addr - orig_compact_base; + new_base = clone.compact_kv_cache_ptr.RowBytes(0) + offset; + new_stride = clone.compact_kv_cache_ptr.Stride(); + } else { + HWY_ABORT("Unrecognized buffer in kv_head_ptrs during ClonePrefix()"); + } + new_ptr.SetPtr(new_base, new_stride); + + // Copy only prefix tiles: T_prefix = min(N_tiles, ceil(prefix_len / + // kTileSize)) + const size_t tiles_to_copy = std::min(orig_ptr.Rows(), prefix_tiles); + if (tiles_to_copy > 0) { + hwy::CopyBytes( + orig_ptr.RowBytes(0), new_ptr.RowBytes(0), + tiles_to_copy * orig_ptr.Stride() * orig_ptr.ElementBytes()); + } + } + clone.kv_head_ptrs.push_back(std::move(new_ptr)); + } + + if (clone.kv_cache.HasPtr()) { + ZeroInit(clone.kv_cache); + } + if (clone.k_cache.HasPtr()) { + ZeroInit(clone.k_cache); + } + if (clone.v_cache.HasPtr()) { + ZeroInit(clone.v_cache); + } + + // Non-tiled legacy buffers + if (kv_cache.HasPtr() && clone.kv_cache.HasPtr()) { + const size_t rows_to_copy = std::min(prefix_len, kv_cache.Rows()); + if (rows_to_copy > 0) { + hwy::CopyBytes( + kv_cache.RowBytes(0), clone.kv_cache.RowBytes(0), + rows_to_copy * kv_cache.Stride() * kv_cache.ElementBytes()); + } + } + if (k_cache.HasPtr() && clone.k_cache.HasPtr()) { + if (k_cache.Cols() != clone.k_cache.Cols() && + k_cache.Rows() != clone.k_cache.Rows()) { + const size_t factor = k_cache.Cols() / clone.k_cache.Cols(); + if (factor > 0 && clone.k_cache.Cols() * factor == k_cache.Cols()) { + clone.k_cache.ReshapePackedRowsToCols(factor); + const size_t prefix_rows = hwy::DivCeil(prefix_len, factor); + const size_t rows_to_copy = std::min(prefix_rows, k_cache.Rows()); + if (rows_to_copy > 0) { + hwy::CopyBytes( + k_cache.RowBytes(0), clone.k_cache.RowBytes(0), + rows_to_copy * k_cache.Stride() * k_cache.ElementBytes()); + } + } + } else if (clone.k_cache.SameShape(k_cache)) { + const size_t rows_to_copy = std::min(prefix_len, k_cache.Rows()); + if (rows_to_copy > 0) { + hwy::CopyBytes( + k_cache.RowBytes(0), clone.k_cache.RowBytes(0), + rows_to_copy * k_cache.Stride() * k_cache.ElementBytes()); + } + } + } + if (v_cache.HasPtr() && clone.v_cache.HasPtr()) { + if (v_cache.Cols() != clone.v_cache.Cols() && + v_cache.Rows() != clone.v_cache.Rows()) { + const size_t factor = v_cache.Cols() / clone.v_cache.Cols(); + if (factor > 0 && clone.v_cache.Cols() * factor == v_cache.Cols()) { + clone.v_cache.ReshapePackedRowsToCols(factor); + const size_t prefix_rows = hwy::DivCeil(prefix_len, factor); + const size_t rows_to_copy = std::min(prefix_rows, v_cache.Rows()); + if (rows_to_copy > 0) { + hwy::CopyBytes( + v_cache.RowBytes(0), clone.v_cache.RowBytes(0), + rows_to_copy * v_cache.Stride() * v_cache.ElementBytes()); + } + } + } else if (clone.v_cache.SameShape(v_cache)) { + const size_t rows_to_copy = std::min(prefix_len, v_cache.Rows()); + if (rows_to_copy > 0) { + hwy::CopyBytes( + v_cache.RowBytes(0), clone.v_cache.RowBytes(0), + rows_to_copy * v_cache.Stride() * v_cache.ElementBytes()); + } + } + } + + // DeepSeek V4 compressor state: deep copy ds_state.Row(0) and snapshot into + // ds_state_snapshot.Row(0). + if (ds_state.Rows() > 0) { + clone.ds_state = MatStorageT("ds_state", ds_state.Extents(), + allocator, MatPadding::kPacked); + CopyMat(ds_state, clone.ds_state); + clone.ds_state_snapshot = MatStorageT( + "ds_snap", ds_state_snapshot.Extents(), allocator, MatPadding::kPacked); + ZeroInit(clone.ds_state_snapshot); + hwy::CopyBytes(ds_state.Row(0), clone.ds_state_snapshot.Row(0), + ds_state.Cols() * sizeof(float)); + clone.ds_state_offsets = ds_state_offsets; + } + + return clone; +} + +void KVCache::RestoreDSStateFromSnapshot(size_t snapshot_row) { + if (ds_state.Rows() > 0 && snapshot_row < ds_state_snapshot.Rows()) { + hwy::CopyBytes(ds_state_snapshot.Row(snapshot_row), ds_state.Row(0), + ds_state.Cols() * sizeof(float)); + } +} + +size_t KVCache::TotalByteSize() const { + size_t total = 0; + if (kv_cache.HasPtr()) { + total += kv_cache.Rows() * kv_cache.Stride() * kv_cache.ElementBytes(); + } + if (k_cache.HasPtr()) { + total += k_cache.Rows() * k_cache.Stride() * k_cache.ElementBytes(); + } + if (v_cache.HasPtr()) { + total += v_cache.Rows() * v_cache.Stride() * v_cache.ElementBytes(); + } + if (compact_local_kv_cache_ptr.HasPtr()) { + total += compact_local_kv_cache_ptr.Rows() * + compact_local_kv_cache_ptr.Stride() * + compact_local_kv_cache_ptr.ElementBytes(); + } + if (compact_global_kv_cache_ptr.HasPtr()) { + total += compact_global_kv_cache_ptr.Rows() * + compact_global_kv_cache_ptr.Stride() * + compact_global_kv_cache_ptr.ElementBytes(); + } + if (!compact_local_kv_cache_ptr.HasPtr() && + !compact_global_kv_cache_ptr.HasPtr() && compact_kv_cache_ptr.HasPtr()) { + total += compact_kv_cache_ptr.Rows() * compact_kv_cache_ptr.Stride() * + compact_kv_cache_ptr.ElementBytes(); + } + if (ds_state.HasPtr()) { + total += ds_state.Rows() * ds_state.Stride() * ds_state.ElementBytes(); + } + if (ds_state_snapshot.HasPtr()) { + total += ds_state_snapshot.Rows() * ds_state_snapshot.Stride() * + ds_state_snapshot.ElementBytes(); + } + return total; +} + } // namespace gcpp diff --git a/gemma/kv_cache.h b/gemma/kv_cache.h index 7a10f571..edada43e 100644 --- a/gemma/kv_cache.h +++ b/gemma/kv_cache.h @@ -51,10 +51,23 @@ struct KVCache { const Allocator& allocator); KVCache(const ModelConfig& config, const InferenceArgs& inference_args, const RuntimeConfig& runtime_config, const Allocator& allocator); + KVCache(KVCache&&) noexcept = default; + KVCache& operator=(KVCache&&) noexcept = default; // Returns a deep copy of the KVCache. Use explicit function instead of // copy ctor to make the cost explicit. + KVCache Copy() const; KVCache Copy(); + // Clones prefix of length prefix_len into a cache with full sequence capacity + // allocated via the specified allocator. + KVCache ClonePrefix(size_t prefix_len, const Allocator& allocator) const; + + // Calculates total byte consumption of all owned buffers. + size_t TotalByteSize() const; + + // Restores ds_state from snapshot row. + void RestoreDSStateFromSnapshot(size_t snapshot_row = 0); + size_t SeqLen() const { if (IsTiled()) { return tiled_seq_len.value(); @@ -234,9 +247,9 @@ struct KVCache { private: const Allocator& allocator_; - // For use by other ctor and Copy() + // For use by Copy() and ClonePrefix() KVCache(const Extents2D& kv_extents, size_t num_layers, size_t kv_heads, - size_t qkv_dim, const Allocator& allocator); + size_t qkv_dim, size_t k_v_cols, const Allocator& allocator); }; inline size_t KVCachePtr::SeqLen() const { diff --git a/gemma/kv_cache_test.cc b/gemma/kv_cache_test.cc index bd7036dd..f579993b 100644 --- a/gemma/kv_cache_test.cc +++ b/gemma/kv_cache_test.cc @@ -1,6 +1,8 @@ #include "gemma/kv_cache.h" #include +#include +#include #include #include "gtest/gtest.h" @@ -82,5 +84,278 @@ TEST(KVCacheTest, SharedLayersReserveNoCache) { EXPECT_EQ(cache.kv_cache.Cols(), model_config.KVCacheCols()); } +TEST(KVCacheTest, CopyTiledCacheAllocatesAndRebases) { + ModelConfig model_config; + model_config.max_seq_len = 1024; + model_config.num_layers = 2; + + // Layer 0: local attention layer + model_config.layer_configs.push_back(LayerConfig()); + model_config.layer_configs.back().kv_heads = 2; + model_config.layer_configs.back().qkv_dim = 256; + model_config.attention_window_sizes.push_back(512); + + // Layer 1: global attention layer + model_config.layer_configs.push_back(LayerConfig()); + model_config.layer_configs.back().kv_heads = 2; + model_config.layer_configs.back().qkv_dim = 512; + model_config.attention_window_sizes.push_back(1024); + + InferenceArgs inference_args; + inference_args.seq_len = 1024; + RuntimeConfig runtime_config; + runtime_config.attention_impl = AttentionImpl::kFlashTransposedQs; + ThreadingArgs threading_args; + ThreadingContext ctx(threading_args); + + KVCache orig(model_config, inference_args, runtime_config, ctx.allocator); + ASSERT_TRUE(orig.compact_local_kv_cache_ptr.HasPtr()); + ASSERT_TRUE(orig.compact_global_kv_cache_ptr.HasPtr()); + ASSERT_EQ(orig.kv_head_ptrs.size(), 4); // 2 heads * 2 layers + + // Fill buffers with pattern + std::memset(orig.compact_local_kv_cache_ptr.RowBytes(0), 0xAB, + orig.compact_local_kv_cache_ptr.Rows() * + orig.compact_local_kv_cache_ptr.Stride() * + orig.compact_local_kv_cache_ptr.ElementBytes()); + std::memset(orig.compact_global_kv_cache_ptr.RowBytes(0), 0xCD, + orig.compact_global_kv_cache_ptr.Rows() * + orig.compact_global_kv_cache_ptr.Stride() * + orig.compact_global_kv_cache_ptr.ElementBytes()); + + KVCache copy = orig.Copy(); + + // Verify memory allocation in copy + EXPECT_TRUE(copy.compact_local_kv_cache_ptr.HasPtr()); + EXPECT_TRUE(copy.compact_global_kv_cache_ptr.HasPtr()); + EXPECT_NE(copy.compact_local_kv_cache_ptr.RowBytes(0), + orig.compact_local_kv_cache_ptr.RowBytes(0)); + EXPECT_NE(copy.compact_global_kv_cache_ptr.RowBytes(0), + orig.compact_global_kv_cache_ptr.RowBytes(0)); + + // Verify data copied + EXPECT_EQ(0, std::memcmp(copy.compact_local_kv_cache_ptr.RowBytes(0), + orig.compact_local_kv_cache_ptr.RowBytes(0), + orig.compact_local_kv_cache_ptr.Rows() * + orig.compact_local_kv_cache_ptr.Stride() * + orig.compact_local_kv_cache_ptr.ElementBytes())); + EXPECT_EQ(0, + std::memcmp(copy.compact_global_kv_cache_ptr.RowBytes(0), + orig.compact_global_kv_cache_ptr.RowBytes(0), + orig.compact_global_kv_cache_ptr.Rows() * + orig.compact_global_kv_cache_ptr.Stride() * + orig.compact_global_kv_cache_ptr.ElementBytes())); + + // Verify pointer rebasing for all kv_head_ptrs + EXPECT_EQ(copy.kv_head_ptrs.size(), orig.kv_head_ptrs.size()); + for (size_t i = 0; i < copy.kv_head_ptrs.size(); ++i) { + const MatPtr& orig_hp = orig.kv_head_ptrs[i]; + const MatPtr& copy_hp = copy.kv_head_ptrs[i]; + EXPECT_NE(copy_hp.RowBytes(0), orig_hp.RowBytes(0)); + EXPECT_EQ(copy_hp.Rows(), orig_hp.Rows()); + EXPECT_EQ(copy_hp.Cols(), orig_hp.Cols()); + EXPECT_EQ(copy_hp.Stride(), orig_hp.Stride()); + + if (i < 2) { + // Local layer heads point into compact_local_kv_cache_ptr + const uintptr_t orig_offset = + orig_hp.RowBytes(0) - orig.compact_local_kv_cache_ptr.RowBytes(0); + const uintptr_t copy_offset = + copy_hp.RowBytes(0) - copy.compact_local_kv_cache_ptr.RowBytes(0); + EXPECT_EQ(copy_offset, orig_offset); + } else { + // Global layer heads point into compact_global_kv_cache_ptr + const uintptr_t orig_offset = + orig_hp.RowBytes(0) - orig.compact_global_kv_cache_ptr.RowBytes(0); + const uintptr_t copy_offset = + copy_hp.RowBytes(0) - copy.compact_global_kv_cache_ptr.RowBytes(0); + EXPECT_EQ(copy_offset, orig_offset); + } + } + + // Verify GetPointers works on copy without crash + std::vector local_ptrs = + copy.GetPointers(/*layer_idx=*/0, /*kv_head_idx=*/0, /*start_pos=*/64, + /*is_global_layer=*/false); + EXPECT_FALSE(local_ptrs.empty()); + std::vector global_ptrs = + copy.GetPointers(/*layer_idx=*/1, /*kv_head_idx=*/1, /*start_pos=*/64, + /*is_global_layer=*/true); + EXPECT_EQ(global_ptrs.size(), 1); + + // Metadata and byte size + EXPECT_EQ(copy.k_v_cols, orig.k_v_cols); + EXPECT_EQ(copy.TotalByteSize(), orig.TotalByteSize()); + EXPECT_GT(copy.TotalByteSize(), 0); +} + +TEST(KVCacheTest, ClonePrefixCopiesOnlyPrefix) { + ModelConfig model_config; + model_config.max_seq_len = 1024; + model_config.num_layers = 2; + + // Layer 0: local attention layer + model_config.layer_configs.push_back(LayerConfig()); + model_config.layer_configs.back().kv_heads = 2; + model_config.layer_configs.back().qkv_dim = 256; + model_config.attention_window_sizes.push_back(512); + + // Layer 1: global attention layer + model_config.layer_configs.push_back(LayerConfig()); + model_config.layer_configs.back().kv_heads = 2; + model_config.layer_configs.back().qkv_dim = 512; + model_config.attention_window_sizes.push_back(1024); + + InferenceArgs inference_args; + inference_args.seq_len = 1024; + RuntimeConfig runtime_config; + runtime_config.attention_impl = AttentionImpl::kFlashTransposedQs; + ThreadingArgs threading_args; + ThreadingContext ctx(threading_args); + + KVCache orig(model_config, inference_args, runtime_config, ctx.allocator); + ASSERT_EQ(orig.kv_head_ptrs.size(), 4); + + // Fill head tile rows with distinct byte values per row + for (size_t h = 0; h < orig.kv_head_ptrs.size(); ++h) { + MatPtr& hp = orig.kv_head_ptrs[h]; + for (size_t r = 0; r < hp.Rows(); ++r) { + uint8_t val = static_cast((h + 1) * 10 + r + 1); + std::memset(hp.RowBytes(r), val, hp.Stride() * hp.ElementBytes()); + } + } + + // Clone prefix of length 40. With kTileSize = 32, prefix_tiles = ceil(40/32) + // = 2. + const size_t prefix_len = 40; + const size_t expected_prefix_tiles = 2; + KVCache clone = orig.ClonePrefix(prefix_len, ctx.allocator); + + EXPECT_EQ(clone.SeqLen(), orig.SeqLen()); + EXPECT_EQ(clone.TotalByteSize(), orig.TotalByteSize()); + EXPECT_EQ(clone.kv_head_ptrs.size(), orig.kv_head_ptrs.size()); + + for (size_t h = 0; h < clone.kv_head_ptrs.size(); ++h) { + const MatPtr& orig_hp = orig.kv_head_ptrs[h]; + const MatPtr& clone_hp = clone.kv_head_ptrs[h]; + EXPECT_NE(clone_hp.RowBytes(0), orig_hp.RowBytes(0)); + EXPECT_EQ(clone_hp.Rows(), orig_hp.Rows()); + + // Rows < expected_prefix_tiles must match original + for (size_t r = 0; r < expected_prefix_tiles; ++r) { + EXPECT_EQ(0, std::memcmp(clone_hp.RowBytes(r), orig_hp.RowBytes(r), + clone_hp.Stride() * clone_hp.ElementBytes())); + } + // Rows >= expected_prefix_tiles must be zeroes (not copied) + for (size_t r = expected_prefix_tiles; r < clone_hp.Rows(); ++r) { + const uint8_t* row = clone_hp.RowBytes(r); + const size_t bytes = clone_hp.Stride() * clone_hp.ElementBytes(); + bool all_zero = true; + for (size_t b = 0; b < bytes; ++b) { + if (row[b] != 0) { + all_zero = false; + break; + } + } + EXPECT_TRUE(all_zero) + << "Head " << h << " row " << r << " should be zero"; + } + } +} + +TEST(KVCacheTest, DeepSeekCompressorStatePreservation) { + ModelConfig model_config; + model_config.max_seq_len = 512; + model_config.num_layers = 1; + model_config.layer_configs.push_back(LayerConfig()); + model_config.layer_configs[0].kv_lora_rank = 32; + model_config.layer_configs[0].o_lora_rank = 32; + model_config.layer_configs[0].rope_head_dim = 32; + model_config.layer_configs[0].kv_compression_rate = 16; + model_config.layer_configs[0].heads = 4; + model_config.layer_configs[0].kv_heads = 4; + + InferenceArgs inference_args; + inference_args.seq_len = 512; + ThreadingArgs threading_args; + ThreadingContext ctx(threading_args); + + KVCache orig(model_config, inference_args, ctx.allocator); + ASSERT_GT(orig.ds_state.Rows(), 0); + ASSERT_GT(orig.ds_state.Cols(), 0); + ASSERT_EQ(orig.ds_state_snapshot.Rows(), 32); + + // Fill orig.ds_state.Row(0) with unique values + float* ds_row = orig.ds_state.Row(0); + for (size_t c = 0; c < orig.ds_state.Cols(); ++c) { + ds_row[c] = static_cast(c + 1) * 0.25f; + } + + // Clone prefix + KVCache clone = orig.ClonePrefix(16, ctx.allocator); + ASSERT_EQ(clone.ds_state.Rows(), 1); + ASSERT_EQ(clone.ds_state.Cols(), orig.ds_state.Cols()); + ASSERT_EQ(clone.ds_state_snapshot.Rows(), 32); + ASSERT_EQ(clone.ds_state_offsets, orig.ds_state_offsets); + + // Verify ds_state.Row(0) and ds_state_snapshot.Row(0) match + for (size_t c = 0; c < clone.ds_state.Cols(); ++c) { + EXPECT_FLOAT_EQ(clone.ds_state.Row(0)[c], ds_row[c]); + EXPECT_FLOAT_EQ(clone.ds_state_snapshot.Row(0)[c], ds_row[c]); + } + + // Mutate clone.ds_state.Row(0) + for (size_t c = 0; c < clone.ds_state.Cols(); ++c) { + clone.ds_state.Row(0)[c] = -999.0f; + } + + // Restore from snapshot row 0 + clone.RestoreDSStateFromSnapshot(0); + for (size_t c = 0; c < clone.ds_state.Cols(); ++c) { + EXPECT_FLOAT_EQ(clone.ds_state.Row(0)[c], ds_row[c]); + } +} + +TEST(KVCacheTest, TotalByteSizeCalculation) { + ModelConfig model_config; + model_config.max_seq_len = 256; + model_config.num_layers = 1; + model_config.layer_configs.push_back(LayerConfig()); + model_config.layer_configs.back().kv_heads = 2; + model_config.layer_configs.back().qkv_dim = 128; + model_config.attention_window_sizes.push_back(256); + + InferenceArgs inference_args; + inference_args.seq_len = 256; + RuntimeConfig runtime_config; + runtime_config.attention_impl = AttentionImpl::kFlashTransposedQs; + ThreadingArgs threading_args; + ThreadingContext ctx(threading_args); + + KVCache cache(model_config, inference_args, runtime_config, ctx.allocator); + const size_t total_bytes = cache.TotalByteSize(); + EXPECT_GT(total_bytes, 0); + + size_t expected_bytes = 0; + if (cache.kv_cache.HasPtr()) { + expected_bytes += cache.kv_cache.Rows() * cache.kv_cache.Stride() * + cache.kv_cache.ElementBytes(); + } + if (cache.k_cache.HasPtr()) { + expected_bytes += cache.k_cache.Rows() * cache.k_cache.Stride() * + cache.k_cache.ElementBytes(); + } + if (cache.v_cache.HasPtr()) { + expected_bytes += cache.v_cache.Rows() * cache.v_cache.Stride() * + cache.v_cache.ElementBytes(); + } + if (cache.compact_global_kv_cache_ptr.HasPtr()) { + expected_bytes += cache.compact_global_kv_cache_ptr.Rows() * + cache.compact_global_kv_cache_ptr.Stride() * + cache.compact_global_kv_cache_ptr.ElementBytes(); + } + EXPECT_EQ(total_bytes, expected_bytes); +} + } // namespace } // namespace gcpp diff --git a/ops/ops-inl.h b/ops/ops-inl.h index c36d555f..71ae82d9 100644 --- a/ops/ops-inl.h +++ b/ops/ops-inl.h @@ -684,10 +684,11 @@ static HWY_NOINLINE void GroupedRMSNormInplace( } template -void GroupedRMSNormBatched( - const MatPtrT& activations, const MatPtr& weights, MatPtrT& out, - const size_t num_groups, ThreadingContext& ctx, size_t cluster_idx = 0, - Parallelism parallelism = Parallelism::kFlat) { +void GroupedRMSNormBatched(const MatPtrT& activations, + const MatPtr& weights, MatPtrT& out, + const size_t num_groups, ThreadingContext& ctx, + size_t cluster_idx = 0, + Parallelism parallelism = Parallelism::kFlat) { CallUpcasted(&weights, [&](const auto* weights_t) { ParallelFor(parallelism, activations.Rows(), ctx, cluster_idx, Callers::kOpsGroupedRMSNormBatched, @@ -1111,8 +1112,7 @@ template > HWY_INLINE HWY_MAYBE_UNUSED void MulByConstAndAddTileUpTo8_BF16_Int16( DF df, const float* HWY_RESTRICT scales_old, size_t actual_steps, const int16_t* HWY_RESTRICT step_cs_i16, - const int8_t* const* HWY_RESTRICT step_v_tiles, - MatPtrT& out, + const int8_t* const* HWY_RESTRICT step_v_tiles, MatPtrT& out, const float* const* HWY_RESTRICT step_q_scales_s) { static_assert(N <= 8); namespace hn = hwy::HWY_NAMESPACE; @@ -1829,12 +1829,12 @@ HWY_INLINE void ApplySoftCap(DF df, float att_cap, float one_over_cap, VF& x0, template , typename DU = hn::ScalableTag, class VU = hn::Vec> HWY_INLINE void ApplyMasking(DF df, DU du, size_t position, - const size_t* HWY_RESTRICT first_pos_per_query, - const size_t* HWY_RESTRICT last_pos_per_query, - VF& x0_p0, VF& x0_p1, VF& x1_p0, VF& x1_p1, - VF& x2_p0, VF& x2_p1, VF& x3_p0, VF& x3_p1, - VF& x4_p0, VF& x4_p1, VF& x5_p0, VF& x5_p1, - VF& x6_p0, VF& x6_p1, VF& x7_p0, VF& x7_p1) { + const size_t* HWY_RESTRICT first_pos_per_query, + const size_t* HWY_RESTRICT last_pos_per_query, + VF& x0_p0, VF& x0_p1, VF& x1_p0, VF& x1_p1, + VF& x2_p0, VF& x2_p1, VF& x3_p0, VF& x3_p1, + VF& x4_p0, VF& x4_p1, VF& x5_p0, VF& x5_p1, + VF& x6_p0, VF& x6_p1, VF& x7_p0, VF& x7_p1) { VU lane_indices = hn::Iota(du, 0); HWY_LANES_CONSTEXPR size_t kTileSize = hn::Lanes(df); auto per_lane_pos_p0 = hn::Add(hn::Set(du, position), lane_indices); @@ -2017,6 +2017,268 @@ HWY_INLINE V SumReduceSegments(D d, V v) { } } +// Highway SIMD Vector Logit Masking Kernel (Requirement R5 Part 1). +// Applies vector logit mask to `logits`. +// Disallowed tokens are overwritten with `mask_value` (-infinity). +HWY_INLINE void ApplyLogitMaskKernel( + hwy::Span logits, const uint64_t* HWY_RESTRICT mask_words, + size_t vocab_size, + float mask_value = -std::numeric_limits::infinity()) { + if (HWY_UNLIKELY(vocab_size == 0 || logits.data() == nullptr || + mask_words == nullptr)) { + return; + } + if (HWY_UNLIKELY(logits.size() < vocab_size)) { + vocab_size = logits.size(); + } + + const size_t num_words = (vocab_size + 63) / 64; + + // Empty mask guardrail: check if any allowed bit exists within vocab_size. + // If entire mask is 0, fail-safe by preserving finite logits to avoid NaN in + // softmax. + bool has_any_allowed = false; + for (size_t w = 0; w < num_words; ++w) { + uint64_t w_val = mask_words[w]; + if (HWY_UNLIKELY(w == num_words - 1 && (vocab_size % 64 != 0))) { + const uint64_t valid_bits = (1ULL << (vocab_size % 64)) - 1; + w_val &= valid_bits; + } + if (w_val != 0ULL) { + has_any_allowed = true; + break; + } + } + if (HWY_UNLIKELY(!has_any_allowed)) { + return; + } + + const hn::ScalableTag df; + const size_t N = hn::Lanes(df); + const auto v_neg_inf = hn::Set(df, mask_value); + + for (size_t w = 0; w < num_words; ++w) { + const uint64_t word = mask_words[w]; + const size_t base_idx = w * 64; + + if (HWY_LIKELY(base_idx + 64 <= vocab_size)) { + // Fast Path 1: All 64 tokens are allowed. Skip writes. + if (word == ~0ULL) { + continue; + } + + // Fast Path 2: All 64 tokens are disallowed. Vector fill -inf. + if (word == 0ULL) { + for (size_t offset = 0; offset < 64; offset += N) { + hn::StoreU(v_neg_inf, df, &logits[base_idx + offset]); + } + continue; + } + + // Optimization for sparse disallowed tokens: + // If <= 4 tokens in the word are disallowed, directly overwrite those + // tokens. Eliminates all memory reads, vector operations, and bulk memory + // writes. + const uint64_t disallowed = ~word; + const size_t num_disallowed = hwy::PopCount(disallowed); + if (num_disallowed <= 4) { + uint64_t rem = disallowed; + while (rem != 0) { + const size_t bit = hwy::Num0BitsBelowLS1Bit_Nonzero64(rem); + logits[base_idx + bit] = mask_value; + rem &= rem - 1; + } + continue; + } + + // Optimization for sparse allowed tokens: + // If <= 4 tokens in the word are allowed, save them, vector fill -inf, + // and restore. Eliminates all memory reads of disallowed logits and + // vector mask blending. + const size_t num_allowed = 64 - num_disallowed; + if (num_allowed <= 4) { + float saved_logits[4]; + size_t saved_indices[4]; + size_t count = 0; + uint64_t rem = word; + while (rem != 0) { + const size_t bit = hwy::Num0BitsBelowLS1Bit_Nonzero64(rem); + saved_indices[count] = base_idx + bit; + saved_logits[count] = logits[saved_indices[count]]; + ++count; + rem &= rem - 1; + } + for (size_t offset = 0; offset < 64; offset += N) { + hn::StoreU(v_neg_inf, df, &logits[base_idx + offset]); + } + for (size_t i = 0; i < count; ++i) { + logits[saved_indices[i]] = saved_logits[i]; + } + continue; + } + + // General Path: Process in vector chunks of N lanes. + const uint8_t* word_bytes = + reinterpret_cast(&mask_words[w]); + for (size_t offset = 0; offset < 64; offset += N) { + const size_t idx = base_idx + offset; + const uint64_t shifted = word >> offset; + const uint64_t chunk_mask = + shifted & ((N >= 64) ? ~0ULL : ((1ULL << N) - 1)); + + // Sub-chunk is all allowed: skip memory writes. + if (chunk_mask == ((N >= 64) ? ~0ULL : ((1ULL << N) - 1))) { + continue; + } + + // Sub-chunk is all disallowed: store mask_value without loading logits. + if (chunk_mask == 0) { + hn::StoreU(v_neg_inf, df, &logits[idx]); + continue; + } + + // Sub-chunk has 1 or 2 disallowed tokens: write mask_value directly. + const uint64_t chunk_disallowed = + ((N >= 64) ? ~0ULL : ((1ULL << N) - 1)) & ~chunk_mask; + if (hwy::PopCount(chunk_disallowed) <= 2) { + uint64_t rem = chunk_disallowed; + while (rem != 0) { + const size_t bit = hwy::Num0BitsBelowLS1Bit_Nonzero64(rem); + logits[idx + bit] = mask_value; + rem &= rem - 1; + } + continue; + } + + // Mixed chunk: load mask bits, load logits, blend, and store. + const auto mask = + (N >= 8) ? hn::LoadMaskBits(df, word_bytes + (offset / 8)) + : hn::LoadMaskBits( + df, reinterpret_cast(&shifted)); + const auto v_logits = hn::LoadU(df, &logits[idx]); + const auto v_masked = hn::IfThenElse(mask, v_logits, v_neg_inf); + hn::StoreU(v_masked, df, &logits[idx]); + } + } else { + // Boundary word when vocab_size % 64 != 0. + const size_t valid_tokens = vocab_size - base_idx; + const uint64_t valid_mask = (1ULL << valid_tokens) - 1; + const uint64_t valid_word = word & valid_mask; + + // Fast Path 1: All remaining tokens are allowed. Skip writes. + if (valid_word == valid_mask) { + continue; + } + + // Fast Path 2: All remaining tokens are disallowed. Vector fill -inf. + if (valid_word == 0ULL) { + for (size_t offset = 0; offset < valid_tokens; offset += N) { + const size_t idx = base_idx + offset; + const size_t rem = vocab_size - idx; + if (rem >= N) { + hn::StoreU(v_neg_inf, df, &logits[idx]); + } else { + hn::StoreN(v_neg_inf, df, &logits[idx], rem); + } + } + continue; + } + + // Sparse disallowed tokens on boundary. + const uint64_t disallowed = valid_mask & ~valid_word; + const size_t num_disallowed = hwy::PopCount(disallowed); + if (num_disallowed <= 4) { + uint64_t rem = disallowed; + while (rem != 0) { + const size_t bit = hwy::Num0BitsBelowLS1Bit_Nonzero64(rem); + logits[base_idx + bit] = mask_value; + rem &= rem - 1; + } + continue; + } + + // Sparse allowed tokens on boundary. + const size_t num_allowed = valid_tokens - num_disallowed; + if (num_allowed <= 4) { + float saved_logits[4]; + size_t saved_indices[4]; + size_t count = 0; + uint64_t rem = valid_word; + while (rem != 0) { + const size_t bit = hwy::Num0BitsBelowLS1Bit_Nonzero64(rem); + saved_indices[count] = base_idx + bit; + saved_logits[count] = logits[saved_indices[count]]; + ++count; + rem &= rem - 1; + } + for (size_t offset = 0; offset < valid_tokens; offset += N) { + const size_t idx = base_idx + offset; + const size_t rem = vocab_size - idx; + if (rem >= N) { + hn::StoreU(v_neg_inf, df, &logits[idx]); + } else { + hn::StoreN(v_neg_inf, df, &logits[idx], rem); + } + } + for (size_t i = 0; i < count; ++i) { + logits[saved_indices[i]] = saved_logits[i]; + } + continue; + } + + // General Path for boundary word. + const uint8_t* word_bytes = + reinterpret_cast(&mask_words[w]); + for (size_t offset = 0; offset < valid_tokens; offset += N) { + const size_t idx = base_idx + offset; + const size_t rem = vocab_size - idx; + const uint64_t shifted = word >> offset; + const uint64_t chunk_valid_mask = + (rem >= N) ? ((N >= 64) ? ~0ULL : ((1ULL << N) - 1)) + : ((1ULL << rem) - 1); + const uint64_t chunk_mask = shifted & chunk_valid_mask; + + if (chunk_mask == chunk_valid_mask) { + continue; + } + if (chunk_mask == 0) { + if (rem >= N) { + hn::StoreU(v_neg_inf, df, &logits[idx]); + } else { + hn::StoreN(v_neg_inf, df, &logits[idx], rem); + } + continue; + } + + const uint64_t chunk_disallowed = chunk_valid_mask & ~chunk_mask; + if (hwy::PopCount(chunk_disallowed) <= 2) { + uint64_t r = chunk_disallowed; + while (r != 0) { + const size_t bit = hwy::Num0BitsBelowLS1Bit_Nonzero64(r); + logits[idx + bit] = mask_value; + r &= r - 1; + } + continue; + } + + const auto mask = + (N >= 8) ? hn::LoadMaskBits(df, word_bytes + (offset / 8)) + : hn::LoadMaskBits( + df, reinterpret_cast(&shifted)); + if (rem >= N) { + const auto v_logits = hn::LoadU(df, &logits[idx]); + const auto v_masked = hn::IfThenElse(mask, v_logits, v_neg_inf); + hn::StoreU(v_masked, df, &logits[idx]); + } else { + const auto v_logits = hn::LoadN(df, &logits[idx], rem); + const auto v_masked = hn::IfThenElse(mask, v_logits, v_neg_inf); + hn::StoreN(v_masked, df, &logits[idx], rem); + } + } + } + } +} + // NOLINTNEXTLINE(google-readability-namespace-comments) } // namespace HWY_NAMESPACE } // namespace gcpp diff --git a/ops/ops.cc b/ops/ops.cc new file mode 100644 index 00000000..d9386428 --- /dev/null +++ b/ops/ops.cc @@ -0,0 +1,42 @@ +// Copyright 2024 Google LLC +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "ops/ops.h" + +#include +#include + +// clang-format off +#undef HWY_TARGET_INCLUDE +#define HWY_TARGET_INCLUDE "ops/ops.cc" // NOLINT +#include "hwy/foreach_target.h" // IWYU pragma: keep +#include "hwy/highway.h" + +// After highway.h +#include "ops/ops-inl.h" +// clang-format on + +#if HWY_ONCE +namespace gcpp { +HWY_EXPORT(ApplyLogitMaskKernel); + +void ApplyLogitMaskKernel(hwy::Span logits, + const uint64_t* HWY_RESTRICT mask_words, + size_t vocab_size, float mask_value) { + HWY_DYNAMIC_DISPATCH(ApplyLogitMaskKernel)(logits, mask_words, vocab_size, + mask_value); +} +} // namespace gcpp +#endif // HWY_ONCE diff --git a/ops/ops.h b/ops/ops.h index 9dd3b917..1be39493 100644 --- a/ops/ops.h +++ b/ops/ops.h @@ -19,12 +19,30 @@ #include #include +#include +#include #include "util/mat.h" +#include "hwy/aligned_allocator.h" // Span #include "hwy/base.h" namespace gcpp { +// Applies vector logit mask to `logits`. +// `mask_words` points to an array of uint64_t words where bit i = 1 indicates +// token i is allowed, and bit i = 0 indicates token i is disallowed. +// Disallowed tokens are overwritten with `mask_value` (-infinity). +// +// Fast Path 1: word == ~0ULL -> all 64 tokens allowed; skip writes. +// Fast Path 2: word == 0ULL -> all 64 tokens disallowed; vector store +// mask_value (-inf). General Path: Highway SIMD LoadMaskBits, IfThenElse, and +// StoreU. Guardrail: Empty mask (0 set bits across entire vocabulary) preserves +// finite logits. +void ApplyLogitMaskKernel( + hwy::Span logits, const uint64_t* HWY_RESTRICT mask_words, + size_t vocab_size, + float mask_value = -std::numeric_limits::infinity()); + static inline HWY_MAYBE_UNUSED MatStorageT CreateInvTimescale( const Allocator& allocator, size_t qkv_dim, bool half_rope, double base_frequency, float partial_rotary_factor = 1.0f) { diff --git a/ops/ops_test.cc b/ops/ops_test.cc index b5bd2f67..9584b23a 100644 --- a/ops/ops_test.cc +++ b/ops/ops_test.cc @@ -20,9 +20,12 @@ #include #include +#include +#include #include #include +#include // NOLINT(build/c++11) #include #include #include @@ -47,8 +50,8 @@ #include "hwy/highway.h" // After highway.h #include "compression/test_util-inl.h" -#include "ops/ops-inl.h" #include "ops/fast_ops-inl.h" +#include "ops/ops-inl.h" #include "hwy/tests/test_util-inl.h" HWY_BEFORE_NAMESPACE(); @@ -925,6 +928,713 @@ void TestPackTokenAndProb() { EXPECT_LT(packed2, packed1); } +void TestApplyLogitMaskAllAllowed() { + const std::vector vocab_sizes = { + 1, 3, 7, 8, 16, 63, 64, 65, 127, 128, 129, + 255, 256, 257, 1000, 1001, 129280, 129281, 256000, 256001, 256063}; + + for (size_t vocab_size : vocab_sizes) { + const size_t kPad = 128; + std::vector buffer(vocab_size + 2 * kPad, 777.0f); + float* logits_ptr = buffer.data() + kPad; + + for (size_t i = 0; i < vocab_size; ++i) { + logits_ptr[i] = static_cast(i) * 0.05f - 10.0f; + } + const std::vector original(logits_ptr, logits_ptr + vocab_size); + + const size_t num_words = (vocab_size + 63) / 64; + std::vector mask(num_words, ~0ULL); + + ApplyLogitMaskKernel(hwy::Span(logits_ptr, vocab_size), mask.data(), + vocab_size); + + // Verify all allowed tokens remain unchanged (delta == 0.0f). + for (size_t i = 0; i < vocab_size; ++i) { + EXPECT_FLOAT_EQ(logits_ptr[i], original[i]); + } + // Verify padding canaries before and after are untouched. + for (size_t i = 0; i < kPad; ++i) { + EXPECT_FLOAT_EQ(buffer[i], 777.0f); + EXPECT_FLOAT_EQ(buffer[kPad + vocab_size + i], 777.0f); + } + } +} + +void TestApplyLogitMaskAllDisallowed() { + const std::vector vocab_sizes = {64, 65, 128, 129, 1000, + 1001, 129280, 256000, 256001}; + + for (size_t vocab_size : vocab_sizes) { + const size_t kPad = 128; + std::vector buffer(vocab_size + 2 * kPad, 777.0f); + float* logits_ptr = buffer.data() + kPad; + + for (size_t i = 0; i < vocab_size; ++i) { + logits_ptr[i] = static_cast(i) * 0.1f + 1.0f; + } + const std::vector original(logits_ptr, logits_ptr + vocab_size); + + const size_t num_words = (vocab_size + 63) / 64; + // Word 0 is all-disallowed (0ULL), word 1 has allowed token, remaining + // words 0ULL. + std::vector mask(num_words, 0ULL); + if (num_words > 1) { + mask[1] = 1ULL; // Token 64 allowed, words 0 and 2..end take Fast Path 2. + } else { + mask[0] = (1ULL << (vocab_size - 1)); // Only last token allowed. + } + + ApplyLogitMaskKernel(hwy::Span(logits_ptr, vocab_size), mask.data(), + vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + const bool is_allowed = (mask[i / 64] & (1ULL << (i % 64))) != 0; + if (is_allowed) { + EXPECT_FLOAT_EQ(logits_ptr[i], original[i]); + } else { + EXPECT_TRUE(std::isinf(logits_ptr[i]) && logits_ptr[i] < 0.0f); + } + } + // Verify canaries untouched. + for (size_t i = 0; i < kPad; ++i) { + EXPECT_FLOAT_EQ(buffer[i], 777.0f); + EXPECT_FLOAT_EQ(buffer[kPad + vocab_size + i], 777.0f); + } + } +} + +void TestApplyLogitMaskEmptyMaskGuardrail() { + const std::vector vocab_sizes = {1, 63, 64, 65, + 128, 1000, 129280, 256000}; + + for (size_t vocab_size : vocab_sizes) { + std::vector logits(vocab_size); + for (size_t i = 0; i < vocab_size; ++i) { + logits[i] = static_cast(i) * 0.2f - 3.0f; + } + const std::vector original = logits; + + const size_t num_words = (vocab_size + 63) / 64; + std::vector empty_mask(num_words, 0ULL); + + // Empty mask must NOT zero or -inf out logits; must preserve original + // finite values. + ApplyLogitMaskKernel(hwy::Span(logits.data(), vocab_size), + empty_mask.data(), vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + EXPECT_FLOAT_EQ(logits[i], original[i]); + } + } + + // Null/zero guardrails + std::vector dummy(10, 1.0f); + ApplyLogitMaskKernel(hwy::Span(), nullptr, 0); + ApplyLogitMaskKernel(hwy::Span(dummy.data(), dummy.size()), nullptr, + 10); + EXPECT_FLOAT_EQ(dummy[0], 1.0f); +} + +void TestApplyLogitMaskMixed() { + const std::vector vocab_sizes = {64, 65, 127, 128, 129, + 256, 1000, 1001, 129280, 256000}; + + for (size_t vocab_size : vocab_sizes) { + const size_t num_words = (vocab_size + 63) / 64; + std::vector logits(vocab_size); + for (size_t i = 0; i < vocab_size; ++i) { + logits[i] = static_cast(i) * 0.01f; + } + const std::vector original = logits; + + // Pattern 1: Alternating even/odd bits + std::vector alt_mask(num_words, 0xAAAAAAAAAAAAAAAAULL); + ApplyLogitMaskKernel(hwy::Span(logits.data(), vocab_size), + alt_mask.data(), vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + const bool is_allowed = (alt_mask[i / 64] & (1ULL << (i % 64))) != 0; + if (is_allowed) { + EXPECT_FLOAT_EQ(logits[i], original[i]); + } else { + EXPECT_TRUE(std::isinf(logits[i]) && logits[i] < 0.0f); + } + } + + // Pattern 2: Custom mask value (-1000.0f) with inverted alternating bits + std::vector logits_custom = original; + std::vector alt_mask2(num_words, 0x5555555555555555ULL); + ApplyLogitMaskKernel(hwy::Span(logits_custom.data(), vocab_size), + alt_mask2.data(), vocab_size, /*mask_value=*/-1000.0f); + + for (size_t i = 0; i < vocab_size; ++i) { + const bool is_allowed = (alt_mask2[i / 64] & (1ULL << (i % 64))) != 0; + if (is_allowed) { + EXPECT_FLOAT_EQ(logits_custom[i], original[i]); + } else { + EXPECT_FLOAT_EQ(logits_custom[i], -1000.0f); + } + } + } +} + +void TestApplyLogitMaskUnalignedVocabSizes() { + // Test irregular boundaries and odd non-multiples of 64 + const std::vector vocab_sizes = { + 1, 2, 3, 5, 7, 9, 15, 17, 31, + 33, 63, 65, 77, 99, 127, 129, 255, 257, + 500, 501, 1023, 1025, 129279, 129281, 255999, 256001, 256063}; + + std::mt19937_64 rng(42); + + for (size_t vocab_size : vocab_sizes) { + const size_t kPad = 64; + std::vector buffer(vocab_size + 2 * kPad, 9999.0f); + float* logits_ptr = buffer.data() + kPad; + + for (size_t i = 0; i < vocab_size; ++i) { + logits_ptr[i] = static_cast(i % 100); + } + const std::vector original(logits_ptr, logits_ptr + vocab_size); + + const size_t num_words = (vocab_size + 63) / 64; + std::vector mask(num_words); + for (size_t w = 0; w < num_words; ++w) { + mask[w] = rng(); + } + // Ensure at least one token is allowed so guardrail doesn't trigger + mask[0] |= 1ULL; + + ApplyLogitMaskKernel(hwy::Span(logits_ptr, vocab_size), mask.data(), + vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + const bool is_allowed = (mask[i / 64] & (1ULL << (i % 64))) != 0; + if (is_allowed) { + EXPECT_FLOAT_EQ(logits_ptr[i], original[i]); + } else { + EXPECT_TRUE(std::isinf(logits_ptr[i]) && logits_ptr[i] < 0.0f); + } + } + + // Verify bounds safety: no overrun before or after logits span + for (size_t i = 0; i < kPad; ++i) { + EXPECT_FLOAT_EQ(buffer[i], 9999.0f); + EXPECT_FLOAT_EQ(buffer[kPad + vocab_size + i], 9999.0f); + } + } +} + +void TestApplyLogitMaskUnalignedPointers() { + const size_t vocab_size = 500; + const size_t num_words = (vocab_size + 63) / 64; + std::mt19937_64 rng(12345); + + std::vector mask(num_words); + for (size_t w = 0; w < num_words; ++w) { + mask[w] = rng(); + } + mask[0] |= 1ULL; + + // Test various pointer offsets from unaligned base + for (size_t misalign = 0; misalign < 16; ++misalign) { + std::vector buffer(vocab_size + misalign + 32, 555.0f); + float* logits_ptr = buffer.data() + misalign; + + for (size_t i = 0; i < vocab_size; ++i) { + logits_ptr[i] = static_cast(i + 1); + } + const std::vector original(logits_ptr, logits_ptr + vocab_size); + + ApplyLogitMaskKernel(hwy::Span(logits_ptr, vocab_size), mask.data(), + vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + const bool is_allowed = (mask[i / 64] & (1ULL << (i % 64))) != 0; + if (is_allowed) { + EXPECT_FLOAT_EQ(logits_ptr[i], original[i]); + } else { + EXPECT_TRUE(std::isinf(logits_ptr[i]) && logits_ptr[i] < 0.0f); + } + } + } +} + +void TestApplyLogitMaskSingleBitSet() { + const std::vector vocab_sizes = {1, 63, 64, 65, 127, + 128, 129, 1000, 256000}; + const std::vector test_indices = { + 0, 1, 2, 31, 32, 63, 64, 65, 127, 128, 129, 999, 1000, 255998, 255999}; + + for (size_t vocab_size : vocab_sizes) { + const size_t num_words = (vocab_size + 63) / 64; + const size_t kPad = 128; + + for (size_t target_idx : test_indices) { + if (target_idx >= vocab_size) continue; + + std::vector buffer(vocab_size + 2 * kPad, 777.0f); + float* logits_ptr = buffer.data() + kPad; + for (size_t i = 0; i < vocab_size; ++i) { + logits_ptr[i] = static_cast(i) * 0.01f + 1.0f; + } + const std::vector original(logits_ptr, logits_ptr + vocab_size); + + std::vector mask(num_words, 0ULL); + mask[target_idx / 64] |= (1ULL << (target_idx % 64)); + + ApplyLogitMaskKernel(hwy::Span(logits_ptr, vocab_size), + mask.data(), vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + if (i == target_idx) { + EXPECT_FLOAT_EQ(logits_ptr[i], original[i]); + } else { + EXPECT_TRUE(std::isinf(logits_ptr[i]) && logits_ptr[i] < 0.0f); + } + } + + for (size_t i = 0; i < kPad; ++i) { + EXPECT_FLOAT_EQ(buffer[i], 777.0f); + EXPECT_FLOAT_EQ(buffer[kPad + vocab_size + i], 777.0f); + } + } + } +} + +void TestApplyLogitMaskSingleBitUnset() { + const std::vector vocab_sizes = {1, 63, 64, 65, 127, + 128, 129, 1000, 256000}; + const std::vector test_indices = { + 0, 1, 2, 31, 32, 63, 64, 65, 127, 128, 129, 999, 1000, 255998, 255999}; + + for (size_t vocab_size : vocab_sizes) { + const size_t num_words = (vocab_size + 63) / 64; + const size_t kPad = 128; + + for (size_t target_idx : test_indices) { + if (target_idx >= vocab_size) continue; + + std::vector buffer(vocab_size + 2 * kPad, 777.0f); + float* logits_ptr = buffer.data() + kPad; + for (size_t i = 0; i < vocab_size; ++i) { + logits_ptr[i] = static_cast(i) * 0.01f + 1.0f; + } + const std::vector original(logits_ptr, logits_ptr + vocab_size); + + std::vector mask(num_words, ~0ULL); + mask[target_idx / 64] &= ~(1ULL << (target_idx % 64)); + + ApplyLogitMaskKernel(hwy::Span(logits_ptr, vocab_size), + mask.data(), vocab_size); + + if (vocab_size == 1) { + EXPECT_FLOAT_EQ(logits_ptr[0], original[0]); + } else { + for (size_t i = 0; i < vocab_size; ++i) { + if (i == target_idx) { + EXPECT_TRUE(std::isinf(logits_ptr[i]) && logits_ptr[i] < 0.0f); + } else { + EXPECT_FLOAT_EQ(logits_ptr[i], original[i]); + } + } + } + + for (size_t i = 0; i < kPad; ++i) { + EXPECT_FLOAT_EQ(buffer[i], 777.0f); + EXPECT_FLOAT_EQ(buffer[kPad + vocab_size + i], 777.0f); + } + } + } +} + +void TestApplyLogitMaskAlternatingPatterns() { + const std::vector patterns = { + 0xAAAAAAAAAAAAAAAAULL, 0x5555555555555555ULL, 0x3333333333333333ULL, + 0x0F0F0F0F0F0F0F0FULL, 0x00FF00FF00FF00FFULL, + }; + const std::vector vocab_sizes = {63, 64, 65, 127, 128, 129, + 255, 256, 257, 1000, 256000}; + + for (size_t vocab_size : vocab_sizes) { + const size_t num_words = (vocab_size + 63) / 64; + const size_t kPad = 64; + + for (uint64_t pattern : patterns) { + std::vector buffer(vocab_size + 2 * kPad, 888.0f); + float* logits_ptr = buffer.data() + kPad; + for (size_t i = 0; i < vocab_size; ++i) { + logits_ptr[i] = static_cast(i % 256) - 128.0f; + } + const std::vector original(logits_ptr, logits_ptr + vocab_size); + + std::vector mask(num_words, pattern); + + ApplyLogitMaskKernel(hwy::Span(logits_ptr, vocab_size), + mask.data(), vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + const bool allowed = (mask[i / 64] & (1ULL << (i % 64))) != 0; + if (allowed) { + EXPECT_FLOAT_EQ(logits_ptr[i], original[i]); + } else { + EXPECT_TRUE(std::isinf(logits_ptr[i]) && logits_ptr[i] < 0.0f); + } + } + + for (size_t i = 0; i < kPad; ++i) { + EXPECT_FLOAT_EQ(buffer[i], 888.0f); + EXPECT_FLOAT_EQ(buffer[kPad + vocab_size + i], 888.0f); + } + } + } +} + +void TestApplyLogitMaskStressAndBenchmark() { + const size_t vocab_size = 256000; + const size_t num_words = (vocab_size + 63) / 64; + std::vector logits(vocab_size, 1.0f); + + struct BenchmarkCase { + const char* name; + std::vector mask; + }; + + std::vector cases; + cases.push_back( + {"All-Allowed (No-op)", std::vector(num_words, ~0ULL)}); + cases.push_back( + {"Empty Mask (Guardrail)", std::vector(num_words, 0ULL)}); + { + std::vector m(num_words, 0ULL); + m[0] = 1ULL; + cases.push_back({"All-Disallowed (Fast Path 2)", std::move(m)}); + } + { + std::vector m(num_words, 0ULL); + m[num_words - 1] = (1ULL << 63); + cases.push_back({"Single Bit Set (Tail)", std::move(m)}); + } + { + std::vector m(num_words, ~0ULL); + m[num_words - 1] &= ~(1ULL << 63); + cases.push_back({"Single Bit Unset (Tail)", std::move(m)}); + } + cases.push_back({"Alternating 0xAAAA (SIMD Worst-Case)", + std::vector(num_words, 0xAAAAAAAAAAAAAAAAULL)}); + cases.push_back({"Alternating 0x5555 (SIMD Worst-Case)", + std::vector(num_words, 0x5555555555555555ULL)}); + { + std::mt19937_64 rng(42); + std::vector m(num_words); + for (size_t w = 0; w < num_words; ++w) { + m[w] = rng(); + } + cases.push_back({"Random 50% Mask", std::move(m)}); + } + + const int kWarmupIterations = 10; + const int kBenchIterations = 100; + + for (const auto& bc : cases) { + for (int it = 0; it < kWarmupIterations; ++it) { + ApplyLogitMaskKernel(hwy::Span(logits.data(), vocab_size), + bc.mask.data(), vocab_size); + } + + const auto start = std::chrono::high_resolution_clock::now(); + for (int it = 0; it < kBenchIterations; ++it) { + ApplyLogitMaskKernel(hwy::Span(logits.data(), vocab_size), + bc.mask.data(), vocab_size); + } + const auto end = std::chrono::high_resolution_clock::now(); + + const double total_us = + std::chrono::duration(end - start).count(); + const double avg_us = total_us / kBenchIterations; + + printf("[BENCHMARK %s] Pattern '%s': %.2f us per call (vocab_size=%zu)\n", + hwy::TargetName(HWY_TARGET), bc.name, avg_us, vocab_size); + // Timing is logged for diagnostic purposes; wall-clock thresholds are + // omitted to prevent non-deterministic failures across varying hardware and + // emulation targets. + (void)avg_us; + } +} + +static void ScalarApplyLogitMaskOracle(std::vector& logits, + const uint64_t* mask_words, + size_t vocab_size, float mask_value) { + if (vocab_size == 0 || logits.empty() || mask_words == nullptr) return; + const size_t num_words = (vocab_size + 63) / 64; + bool has_any_allowed = false; + for (size_t w = 0; w < num_words; ++w) { + uint64_t w_val = mask_words[w]; + if (w == num_words - 1 && (vocab_size % 64 != 0)) { + w_val &= (1ULL << (vocab_size % 64)) - 1; + } + if (w_val != 0ULL) { + has_any_allowed = true; + break; + } + } + if (!has_any_allowed) return; + + for (size_t i = 0; i < vocab_size; ++i) { + const bool allowed = (mask_words[i / 64] & (1ULL << (i % 64))) != 0; + if (!allowed) { + logits[i] = mask_value; + } + } +} + +void TestApplyLogitMaskFuzzingDifferentialOracle() { + std::vector vocab_sizes; + // Sweep all small sizes 1 to 130 + for (size_t s = 1; s <= 130; ++s) { + vocab_sizes.push_back(s); + } + // Sweep around powers of 2 and model vocabulary limits up to 300,000 + const std::vector large_sizes = { + 255, 256, 257, 511, 512, 513, 1023, 1024, 1025, + 2047, 2048, 2049, 4095, 4096, 4097, 8191, 8192, 8193, + 16383, 16384, 16385, 32000, 32001, 65535, 65536, 65537, 129279, + 129280, 129281, 256000, 256063, 256064, 256128, 299993, 299999, 300000}; + vocab_sizes.insert(vocab_sizes.end(), large_sizes.begin(), large_sizes.end()); + + std::mt19937_64 rng(987654321ULL); + constexpr size_t kPad = 64; + + for (size_t vocab_size : vocab_sizes) { + const size_t num_words = (vocab_size + 63) / 64; + std::vector buffer(vocab_size + 2 * kPad, 12345.0f); + float* logits_ptr = buffer.data() + kPad; + + for (size_t i = 0; i < vocab_size; ++i) { + logits_ptr[i] = static_cast(i % 1000) * 0.1f - 50.0f; + } + const std::vector original(logits_ptr, logits_ptr + vocab_size); + + for (int pattern_type = 0; pattern_type < 6; ++pattern_type) { + std::vector mask(num_words, 0ULL); + switch (pattern_type) { + case 0: // Only first token allowed + mask[0] = 1ULL; + break; + case 1: // Only last valid token allowed + mask[(vocab_size - 1) / 64] = (1ULL << ((vocab_size - 1) % 64)); + break; + case 2: // Dense allowed (all 1s) + std::fill(mask.begin(), mask.end(), ~0ULL); + break; + case 3: // Alternating 0xAAAAAAAAAAAAAAAA + std::fill(mask.begin(), mask.end(), 0xAAAAAAAAAAAAAAAAULL); + break; + case 4: // Sparse random (~2% allowed) + for (size_t w = 0; w < num_words; ++w) { + mask[w] = (rng() & rng() & rng() & rng()); + } + mask[0] |= 1ULL; // Ensure at least one allowed + break; + case 5: // Uniform random + for (size_t w = 0; w < num_words; ++w) { + mask[w] = rng(); + } + mask[(vocab_size - 1) / 64] |= (1ULL << ((vocab_size - 1) % 64)); + break; + } + + std::vector expected = original; + ScalarApplyLogitMaskOracle(expected, mask.data(), vocab_size, + -std::numeric_limits::infinity()); + + std::copy(original.begin(), original.end(), logits_ptr); + + ApplyLogitMaskKernel(hwy::Span(logits_ptr, vocab_size), + mask.data(), vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + if (std::isinf(expected[i])) { + ASSERT_TRUE(std::isinf(logits_ptr[i]) && logits_ptr[i] < 0.0f) + << "Failed at vocab_size=" << vocab_size + << ", pattern=" << pattern_type << ", index=" << i; + } else { + ASSERT_FLOAT_EQ(logits_ptr[i], expected[i]) + << "Failed at vocab_size=" << vocab_size + << ", pattern=" << pattern_type << ", index=" << i; + } + } + + for (size_t i = 0; i < kPad; ++i) { + ASSERT_FLOAT_EQ(buffer[i], 12345.0f); + ASSERT_FLOAT_EQ(buffer[kPad + vocab_size + i], 12345.0f); + } + } + } +} + +void TestApplyLogitMaskNumericalStabilityAndBitExactness() { + const size_t vocab_size = 256; + const size_t num_words = (vocab_size + 63) / 64; + + std::vector logits(vocab_size); + logits[0] = 0.0f; + logits[1] = -0.0f; + logits[2] = std::numeric_limits::infinity(); + logits[3] = -std::numeric_limits::infinity(); + logits[4] = std::numeric_limits::denorm_min(); + logits[5] = -std::numeric_limits::denorm_min(); + logits[6] = std::numeric_limits::max(); + logits[7] = std::numeric_limits::lowest(); + uint32_t nan_bits = 0x7fc01234; + float custom_nan; + hwy::CopySameSize(&nan_bits, &custom_nan); + logits[8] = custom_nan; + + for (size_t i = 9; i < vocab_size; ++i) { + logits[i] = static_cast(i); + } + + std::vector mask(num_words, 0x5555555555555555ULL); + + std::vector test_logits = logits; + ApplyLogitMaskKernel(hwy::Span(test_logits.data(), vocab_size), + mask.data(), vocab_size); + + for (size_t i = 0; i < vocab_size; i += 2) { + uint32_t orig_bits, masked_bits; + hwy::CopySameSize(&logits[i], &orig_bits); + hwy::CopySameSize(&test_logits[i], &masked_bits); + EXPECT_EQ(orig_bits, masked_bits) + << "Bit mismatch for allowed token at index " << i; + } + + const uint32_t kNegInfBits = 0xFF800000; + for (size_t i = 1; i < vocab_size; i += 2) { + uint32_t masked_bits; + hwy::CopySameSize(&test_logits[i], &masked_bits); + EXPECT_EQ(masked_bits, kNegInfBits) + << "Exact -inf bit mismatch for disallowed token at index " << i; + } + + for (float custom_mask_val : {-1000.0f, -0.0f, 0.0f, 42.0f}) { + test_logits = logits; + ApplyLogitMaskKernel(hwy::Span(test_logits.data(), vocab_size), + mask.data(), vocab_size, custom_mask_val); + uint32_t expected_val_bits; + hwy::CopySameSize(&custom_mask_val, &expected_val_bits); + for (size_t i = 1; i < vocab_size; i += 2) { + uint32_t actual_val_bits; + hwy::CopySameSize(&test_logits[i], &actual_val_bits); + EXPECT_EQ(actual_val_bits, expected_val_bits); + } + } +} + +void TestApplyLogitMaskPageBoundaryProtection() { + const size_t page_size = sysconf(_SC_PAGESIZE); + const size_t total_alloc = 2 * page_size; + + void* addr = mmap(nullptr, total_alloc, PROT_READ | PROT_WRITE, + MAP_PRIVATE | MAP_ANONYMOUS, -1, 0); + ASSERT_NE(addr, MAP_FAILED); + + uint8_t* base = static_cast(addr); + ASSERT_EQ(mprotect(base + page_size, page_size, PROT_NONE), 0); + + for (size_t vocab_size : {1, 7, 15, 16, 31, 32, 63, 64, 65, 127, 128}) { + const size_t num_words = (vocab_size + 63) / 64; + const size_t bytes_needed = num_words * sizeof(uint64_t); + + uint64_t* mask_words = + reinterpret_cast(base + page_size - bytes_needed); + + for (size_t w = 0; w < num_words; ++w) { + mask_words[w] = 0xAAAAAAAAAAAAAAAAULL; + } + mask_words[0] |= 1ULL; + + std::vector logits(vocab_size, 1.0f); + + ApplyLogitMaskKernel(hwy::Span(logits.data(), vocab_size), + mask_words, vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + const bool allowed = (mask_words[i / 64] & (1ULL << (i % 64))) != 0; + if (allowed) { + EXPECT_FLOAT_EQ(logits[i], 1.0f); + } else { + EXPECT_TRUE(std::isinf(logits[i]) && logits[i] < 0.0f); + } + } + } + + for (size_t vocab_size : {1, 3, 7, 15, 16, 31, 32, 63, 64, 65, 127, 128}) { + const size_t bytes_needed = vocab_size * sizeof(float); + float* logits_at_edge = + reinterpret_cast(base + page_size - bytes_needed); + + for (size_t i = 0; i < vocab_size; ++i) { + logits_at_edge[i] = static_cast(i); + } + const size_t num_words = (vocab_size + 63) / 64; + std::vector mask(num_words, 0xAAAAAAAAAAAAAAAAULL); + mask[0] |= 1ULL; + + ApplyLogitMaskKernel(hwy::Span(logits_at_edge, vocab_size), + mask.data(), vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + const bool allowed = (mask[i / 64] & (1ULL << (i % 64))) != 0; + if (allowed) { + EXPECT_FLOAT_EQ(logits_at_edge[i], static_cast(i)); + } else { + EXPECT_TRUE(std::isinf(logits_at_edge[i]) && logits_at_edge[i] < 0.0f); + } + } + } + + munmap(addr, total_alloc); +} + +void TestApplyLogitMaskExtremeMisalignments() { + const std::vector vocab_sizes = {1, 7, 16, 63, 64, + 65, 127, 128, 500, 129280}; + std::mt19937_64 rng(54321); + + for (size_t vocab_size : vocab_sizes) { + const size_t num_words = (vocab_size + 63) / 64; + std::vector mask(num_words); + for (size_t w = 0; w < num_words; ++w) { + mask[w] = rng(); + } + mask[0] |= 1ULL; + + for (size_t misalign = 0; misalign < 16; ++misalign) { + std::vector buffer(vocab_size + misalign + 32, 333.0f); + float* logits_ptr = buffer.data() + misalign; + + for (size_t i = 0; i < vocab_size; ++i) { + logits_ptr[i] = static_cast(i + 1); + } + const std::vector original(logits_ptr, logits_ptr + vocab_size); + + ApplyLogitMaskKernel(hwy::Span(logits_ptr, vocab_size), + mask.data(), vocab_size); + + for (size_t i = 0; i < vocab_size; ++i) { + const bool is_allowed = (mask[i / 64] & (1ULL << (i % 64))) != 0; + if (is_allowed) { + EXPECT_FLOAT_EQ(logits_ptr[i], original[i]); + } else { + EXPECT_TRUE(std::isinf(logits_ptr[i]) && logits_ptr[i] < 0.0f); + } + } + } + } +} + // NOLINTNEXTLINE(google-readability-namespace-comments) } // namespace HWY_NAMESPACE } // namespace gcpp @@ -954,8 +1664,39 @@ HWY_EXPORT_AND_TEST_P(OpsTest, TestLayerNormSimple); HWY_EXPORT_AND_TEST_P(OpsTest, TestSampleTopK); HWY_EXPORT_AND_TEST_P(OpsTest, TestTopK); HWY_EXPORT_AND_TEST_P(OpsTest, TestPackTokenAndProb); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskAllAllowed); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskAllDisallowed); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskEmptyMaskGuardrail); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskMixed); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskUnalignedVocabSizes); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskUnalignedPointers); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskSingleBitSet); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskSingleBitUnset); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskAlternatingPatterns); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskStressAndBenchmark); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskFuzzingDifferentialOracle); +HWY_EXPORT_AND_TEST_P(OpsTest, + TestApplyLogitMaskNumericalStabilityAndBitExactness); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskPageBoundaryProtection); +HWY_EXPORT_AND_TEST_P(OpsTest, TestApplyLogitMaskExtremeMisalignments); HWY_AFTER_TEST(); +TEST(OpsTest, TestApplyLogitMaskDynamicDispatch) { + const size_t vocab_size = 100; + std::vector logits(vocab_size, 1.0f); + const size_t num_words = (vocab_size + 63) / 64; + std::vector mask(num_words, 0ULL); + mask[0] = 1ULL; // Only token 0 allowed + + gcpp::ApplyLogitMaskKernel(hwy::Span(logits.data(), vocab_size), + mask.data(), vocab_size); + + EXPECT_FLOAT_EQ(logits[0], 1.0f); + for (size_t i = 1; i < vocab_size; ++i) { + EXPECT_TRUE(std::isinf(logits[i]) && logits[i] < 0.0f); + } +} + } // namespace gcpp #endif