diff --git a/benchmarks/DotLLM.Benchmarks/DirectKvWriteBenchmarks.cs b/benchmarks/DotLLM.Benchmarks/DirectKvWriteBenchmarks.cs new file mode 100644 index 00000000..2b39d5d2 --- /dev/null +++ b/benchmarks/DotLLM.Benchmarks/DirectKvWriteBenchmarks.cs @@ -0,0 +1,138 @@ +using BenchmarkDotNet.Attributes; +using DotLLM.Core.Attention; +using DotLLM.Core.Tensors; +using DotLLM.Engine.KvCache; +using DotLLM.HuggingFace; +using DotLLM.Models.Architectures; +using DotLLM.Models.Gguf; +using DotLLM.Tokenizers.Bpe; + +namespace DotLLM.Benchmarks; + +/// +/// Measures the per-decode-step saving from direct-to-cache K/V projection (#25 item 4): +/// the K and V GEMMs write straight into the cache slot, skipping the scratch + Update +/// memcpy. forces the pre-#278 baseline by returning +/// false from . +/// +/// Per-decode-step saving on a 7B GQA-2 (kvStride = 1024 floats / 4 KiB) over 32 layers: +/// 2 × 4 KiB × 32 = 256 KiB of copy per token eliminated. The exact wall-time impact is +/// model and cache-size dependent — this benchmark prints the delta on the bundled +/// SmolLM-135M Q8_0 model (small but real). +/// +[SimpleJob(warmupCount: 2, iterationCount: 5)] +public class DirectKvWriteBenchmarks +{ + private GgufFile _gguf = null!; + private TransformerModel _model = null!; + private BpeTokenizer _tokenizer = null!; + private int[] _promptIds = null!; + private int[] _positions = null!; + + // Two pre-prefilled caches, reset to the same prefill state between iterations. + private SimpleKvCache _baseCache = null!; + + [GlobalSetup] + public void Setup() + { + const string Repo = "QuantFactory/SmolLM-135M-GGUF"; + const string FileName = "SmolLM-135M.Q8_0.gguf"; + string cacheDir = Path.Combine( + Environment.GetFolderPath(Environment.SpecialFolder.UserProfile), + ".dotllm", "test-cache"); + string cachedPath = Path.Combine(cacheDir, Repo.Replace('/', Path.DirectorySeparatorChar), FileName); + string path; + if (File.Exists(cachedPath)) + { + path = cachedPath; + } + else + { + using var downloader = new HuggingFaceDownloader(); + path = downloader.DownloadFileAsync(Repo, FileName, cacheDir).GetAwaiter().GetResult(); + } + + _gguf = GgufFile.Open(path); + var cfg = GgufModelConfigExtractor.Extract(_gguf.Metadata); + _model = TransformerModel.LoadFromGguf(_gguf, cfg); + _tokenizer = GgufBpeTokenizerFactory.Load(_gguf.Metadata); + + _promptIds = _tokenizer.Encode("The capital of France is"); + _positions = new int[_promptIds.Length + 64]; + for (int i = 0; i < _positions.Length; i++) _positions[i] = i; + + _baseCache = NewPrefilledCache(); + } + + [GlobalCleanup] + public void Cleanup() + { + _baseCache.Dispose(); + _model.Dispose(); + _gguf.Dispose(); + } + + private SimpleKvCache NewPrefilledCache() + { + var cache = new SimpleKvCache( + _model.Config.NumLayers, _model.Config.NumKvHeads, _model.Config.HeadDim, + _positions.Length); + using var _ = _model.Forward(_promptIds, _positions.AsSpan(0, _promptIds.Length), -1, cache); + return cache; + } + + /// One decode step with the direct-to-cache path enabled (the new default). + [Benchmark(Baseline = false)] + public int Decode_DirectToCache() + { + using var cache = NewPrefilledCache(); + int pos = _promptIds.Length; + using var logits = _model.Forward([_promptIds[^1]], _positions.AsSpan(pos, 1), -1, cache); + return cache.CurrentLength; + } + + /// + /// One decode step forced onto the legacy path + /// via . The delta to + /// is the K/V scratch→cache memcpy saved. + /// + [Benchmark(Baseline = true)] + public int Decode_LegacyUpdate() + { + using var inner = NewPrefilledCache(); + using var legacy = new LegacyUpdateOnlyCache(inner); + int pos = _promptIds.Length; + using var logits = _model.Forward([_promptIds[^1]], _positions.AsSpan(pos, 1), -1, legacy); + return legacy.CurrentLength; + } + + /// + /// IKvCache decorator that forces the legacy path by + /// returning false from . Used to A/B + /// the direct-to-cache optimisation against the pre-#278 baseline behaviour. + /// + private sealed class LegacyUpdateOnlyCache : IKvCache + { + private readonly IKvCache _inner; + public LegacyUpdateOnlyCache(IKvCache inner) => _inner = inner; + public int CurrentLength => _inner.CurrentLength; + public int MaxLength => _inner.MaxLength; + public void Update(ITensor keys, ITensor values, ReadOnlySpan positions, int layerIndex) => + _inner.Update(keys, values, positions, layerIndex); + public void Update(TensorRef keys, TensorRef values, ReadOnlySpan positions, int layerIndex) => + _inner.Update(keys, values, positions, layerIndex); + public ITensor GetKeys(int layerIndex) => _inner.GetKeys(layerIndex); + public ITensor GetValues(int layerIndex) => _inner.GetValues(layerIndex); + public TensorRef GetKeysRef(int layerIndex) => _inner.GetKeysRef(layerIndex); + public TensorRef GetValuesRef(int layerIndex) => _inner.GetValuesRef(layerIndex); + public void Rollback(int length) => _inner.Rollback(length); + public bool TryReserveSlot(int layerIndex, ReadOnlySpan positions, out Span kDst, out Span vDst) + { + kDst = default; + vDst = default; + return false; + } + public void CommitSlot(int layerIndex, ReadOnlySpan positions) { } + public void Dispose() { /* outer benchmark owns inner */ } + } +} diff --git a/src/DotLLM.Core/Attention/IKvCache.cs b/src/DotLLM.Core/Attention/IKvCache.cs index 6052835c..f181ad24 100644 --- a/src/DotLLM.Core/Attention/IKvCache.cs +++ b/src/DotLLM.Core/Attention/IKvCache.cs @@ -56,4 +56,63 @@ public interface IKvCache : IDisposable /// /// The new current length (must be <= ). void Rollback(int length); + + /// + /// Attempts to reserve in-place write slots for the K and V projections at the given + /// . When successful, callers can target + /// and as the K/V projection output buffers, and run the + /// post-projection in-place pipeline (AddBias, LoRA delta, QK-norm, RoPE) directly on + /// those spans — avoiding the scratch buffer and the subsequent Update + /// memcpy. Length advancement is deferred to ; the caller must + /// invoke after writing to keep + /// consistent. + /// + /// + /// + /// Returns false when the cache cannot expose an in-place slot for the given + /// positions — most commonly because positions are non-contiguous, exceed + /// , would span a paged-block boundary, or the underlying storage + /// is quantized / device-resident. The caller must then fall back to the existing + /// scratch + Update path. + /// + /// + /// The default implementation returns false, preserving backward compatibility + /// for every implementation that has not opted in. + /// + /// + /// Transformer layer index. + /// Position indices for the new entries. Must be contiguous for + /// the slot to be reservable. + /// On success, span covering the K cache slot for these positions + /// (positions.Length * kvStride FP32 elements). Undefined on failure. + /// On success, span covering the V cache slot for these positions. + /// Undefined on failure. + /// true when a slot was reserved and / + /// are valid in-place targets; false otherwise. + bool TryReserveSlot( + int layerIndex, + ReadOnlySpan positions, + out Span kDst, + out Span vDst) + { + kDst = default; + vDst = default; + return false; + } + + /// + /// Commits a prior successful call by advancing + /// based on . Idempotent across + /// layers within the same forward pass — the maximum-position computation matches + /// Update's semantics. + /// + /// + /// The default implementation is a no-op. Callers must only invoke this after a + /// successful on the same cache for the same positions. + /// + /// Transformer layer index. + /// Position indices for the entries written during the slot. + void CommitSlot(int layerIndex, ReadOnlySpan positions) + { + } } diff --git a/src/DotLLM.Engine/KvCache/PagedKvCache.cs b/src/DotLLM.Engine/KvCache/PagedKvCache.cs index 9864fc99..4c1ba63a 100644 --- a/src/DotLLM.Engine/KvCache/PagedKvCache.cs +++ b/src/DotLLM.Engine/KvCache/PagedKvCache.cs @@ -174,6 +174,66 @@ public void Rollback(int length) _blockTable.SetCurrentLength(length); } + /// + public bool TryReserveSlot( + int layerIndex, + ReadOnlySpan positions, + out Span kDst, + out Span vDst) + { + kDst = default; + vDst = default; + + int seqLen = positions.Length; + if (seqLen == 0) return false; + + int start = positions[0]; + + // Contiguous run required (GEMM output is contiguous). + for (int i = 1; i < seqLen; i++) + { + if (positions[i] != start + i) return false; + } + + // Bounds: entire run must fit within MaxLength. + if ((uint)start >= (uint)_maxSeqLen) return false; + if (start + seqLen > _maxSeqLen) return false; + + // Single-block run only: the run must not cross a block boundary, otherwise the + // in-place slot wouldn't be physically contiguous. Decode (seqLen=1) always + // satisfies this; multi-token runs only when they fit inside one block. + int blockSize = _pool.BlockSize; + int offset = start % blockSize; + if (offset + seqLen > blockSize) return false; + + // Ensure a block exists (with refcount-1 fast-path) for the start position. + _blockTable.EnsureCapacity(start + seqLen); + _blockTable.EnsureWritable(start); + var (blockId, offsetInBlock) = _blockTable.Resolve(start); + + int totalFloats = seqLen * _kvStride; + kDst = new Span(_pool.GetKeyPtr(blockId, layerIndex) + offsetInBlock * _kvStride, totalFloats); + vDst = new Span(_pool.GetValuePtr(blockId, layerIndex) + offsetInBlock * _kvStride, totalFloats); + return true; + } + + /// + public void CommitSlot(int layerIndex, ReadOnlySpan positions) + { + int seqLen = positions.Length; + if (seqLen == 0) return; + + int maxPos = positions[0]; + for (int i = 1; i < seqLen; i++) + { + if (positions[i] > maxPos) maxPos = positions[i]; + } + + int newLength = maxPos + 1; + if (newLength > _blockTable.CurrentLength) + _blockTable.Advance(newLength); + } + /// /// Gathers block data into a contiguous staging buffer for attention kernel consumption. /// Copies block-by-block in logical order. diff --git a/src/DotLLM.Engine/KvCache/SimpleKvCache.cs b/src/DotLLM.Engine/KvCache/SimpleKvCache.cs index 56336499..4a19d0e4 100644 --- a/src/DotLLM.Engine/KvCache/SimpleKvCache.cs +++ b/src/DotLLM.Engine/KvCache/SimpleKvCache.cs @@ -151,6 +151,58 @@ public void Rollback(int length) _currentLength = length; } + /// + public bool TryReserveSlot( + int layerIndex, + ReadOnlySpan positions, + out Span kDst, + out Span vDst) + { + kDst = default; + vDst = default; + + int seqLen = positions.Length; + if (seqLen == 0) return false; + + int start = positions[0]; + + // Contiguous run required: the GEMM writes [seqLen, kvStride] as a single + // contiguous block, which only maps to a contiguous cache region. + for (int i = 1; i < seqLen; i++) + { + if (positions[i] != start + i) return false; + } + + // Bounds: the entire run must fit within the cache. + if ((uint)start >= (uint)_maxSeqLen) return false; + if (start + seqLen > _maxSeqLen) return false; + + if ((uint)layerIndex >= (uint)_numLayers) + throw new ArgumentOutOfRangeException(nameof(layerIndex)); + + int totalFloats = seqLen * _kvStride; + kDst = new Span((float*)_keys[layerIndex] + (long)start * _kvStride, totalFloats); + vDst = new Span((float*)_values[layerIndex] + (long)start * _kvStride, totalFloats); + return true; + } + + /// + public void CommitSlot(int layerIndex, ReadOnlySpan positions) + { + int seqLen = positions.Length; + if (seqLen == 0) return; + + int maxPos = positions[0]; + for (int i = 1; i < seqLen; i++) + { + if (positions[i] > maxPos) maxPos = positions[i]; + } + + int newLength = maxPos + 1; + if (newLength > _currentLength) + _currentLength = newLength; + } + /// public void Dispose() { diff --git a/src/DotLLM.Models/Architectures/TransformerModel.cs b/src/DotLLM.Models/Architectures/TransformerModel.cs index 689a1334..1c94e83c 100644 --- a/src/DotLLM.Models/Architectures/TransformerModel.cs +++ b/src/DotLLM.Models/Architectures/TransformerModel.cs @@ -193,6 +193,27 @@ public ITensor Forward(ReadOnlySpan tokenIds, ReadOnlySpan positions, // b. RMSNorm + Pre-quantize + Q/K/V projections byte* inputQ8Scratch = (byte*)_state.InputQ8Scratch; + // Direct-to-cache K/V opt (#25 item 4): when the cache can expose an + // in-place slot for these positions, point the K and V projection + // outputs at the slot so the GEMM, bias, QK-norm, and RoPE all run + // directly on the cache buffer — skipping the scratch + `Update` + // memcpy. The slot is committed (length advance) after the in-place + // pipeline completes. Caches that can't expose a slot (quantized, + // CUDA, hybrid, or paged spanning a block boundary) return false from + // TryReserveSlot and the legacy scratch + Update path runs unchanged. + // + // Q always stays in scratch — it isn't cached. + float* kTarget = k; + float* vTarget = v; + bool kvSlotReserved = false; + if (kvCache is not null && kvCache.TryReserveSlot(layer, positions, + out Span kSlot, out Span vSlot)) + { + kTarget = (float*)Unsafe.AsPointer(ref MemoryMarshal.GetReference(kSlot)); + vTarget = (float*)Unsafe.AsPointer(ref MemoryMarshal.GetReference(vSlot)); + kvSlotReserved = true; + } + if (seqLen == 1 && _threadPool != null) { // Decode path: try fused RmsNorm+Quantize (skips normOut intermediate) @@ -214,7 +235,7 @@ public ITensor Forward(ReadOnlySpan tokenIds, ReadOnlySpan positions, preQuantNorm = QuantizeInput(normOut, inputQ8Scratch, hiddenSize, 1, lw.QQuantType); } - FusedQkvDecode(in lw, normOut, preQuantNorm, q, k, v); + FusedQkvDecode(in lw, normOut, preQuantNorm, q, kTarget, vTarget); } else { @@ -234,27 +255,29 @@ public ITensor Forward(ReadOnlySpan tokenIds, ReadOnlySpan positions, var rwV = rl?.V ?? default; GemmInterleaved(lw.QWeight, lw.QQuantType, normOut, q, lw.QOutputDim, lw.QInputDim, seqLen, preQuantNorm, in rwQ); - GemmInterleaved(lw.KWeight, lw.KQuantType, normOut, k, lw.KOutputDim, lw.KInputDim, seqLen, + GemmInterleaved(lw.KWeight, lw.KQuantType, normOut, kTarget, lw.KOutputDim, lw.KInputDim, seqLen, IsCompatiblePreQuant(lw.QQuantType, lw.KQuantType) ? preQuantNorm : null, in rwK); - GemmInterleaved(lw.VWeight, lw.VQuantType, normOut, v, lw.VOutputDim, lw.VInputDim, seqLen, + GemmInterleaved(lw.VWeight, lw.VQuantType, normOut, vTarget, lw.VOutputDim, lw.VInputDim, seqLen, IsCompatiblePreQuant(lw.QQuantType, lw.VQuantType) ? preQuantNorm : null, in rwV); } - // Optional bias: y = Wx + b (no-op when null) + // Optional bias: y = Wx + b (no-op when null). All run in place on + // kTarget / vTarget, which is either scratch or the cache slot. AddBias(lw.QBias, q, lw.QOutputDim, seqLen); - AddBias(lw.KBias, k, lw.KOutputDim, seqLen); - AddBias(lw.VBias, v, lw.VOutputDim, seqLen); + AddBias(lw.KBias, kTarget, lw.KOutputDim, seqLen); + AddBias(lw.VBias, vTarget, lw.VOutputDim, seqLen); // Optional QK-norms (Qwen3-style): per-head RMSNorm on Q/K after projection, before RoPE if (lw.QNormWeight is not null) ApplyPerHeadNorm(lw.QNormWeight, q, numHeads, headDim, seqLen, eps); if (lw.KNormWeight is not null) - ApplyPerHeadNorm(lw.KNormWeight, k, numKvHeads, headDim, seqLen, eps); + ApplyPerHeadNorm(lw.KNormWeight, kTarget, numKvHeads, headDim, seqLen, eps); - // d. RoPE (in-place on Q and K for all tokens) + // d. RoPE (in-place on Q and K for all tokens) — K rotates inside + // the cache slot when direct-to-cache is active. RoPE.Execute( new Span(q, seqLen * numHeads * headDim), - new Span(k, seqLen * kvStride), + new Span(kTarget, seqLen * kvStride), positions, numHeads, numKvHeads, headDim, _ropeDim, _state.CosTable, _state.SinTable, _ropeType); @@ -262,11 +285,18 @@ public ITensor Forward(ReadOnlySpan tokenIds, ReadOnlySpan positions, // e. Attention — with or without KV-cache if (kvCache is not null) { - // Store new K/V in cache, then attend over full cached context (zero allocations) - var kRef = new TensorRef(seqLen, kvStride, DType.Float32, -1, (nint)k); - var vRef = new TensorRef(seqLen, kvStride, DType.Float32, -1, (nint)v); - - kvCache.Update(kRef, vRef, positions, layer); + if (kvSlotReserved) + { + // K/V are already in the cache slot — just advance length. + kvCache.CommitSlot(layer, positions); + } + else + { + // Legacy path: K/V live in scratch; copy into the cache. + var kRef = new TensorRef(seqLen, kvStride, DType.Float32, -1, (nint)kTarget); + var vRef = new TensorRef(seqLen, kvStride, DType.Float32, -1, (nint)vTarget); + kvCache.Update(kRef, vRef, positions, layer); + } int seqKv = kvCache.CurrentLength; @@ -289,7 +319,9 @@ public ITensor Forward(ReadOnlySpan tokenIds, ReadOnlySpan positions, } else { - Attention.Execute(q, k, v, attnOut, + // kvCache==null path: TryReserveSlot was never attempted; K/V are + // in scratch (kTarget == k, vTarget == v). + Attention.Execute(q, kTarget, vTarget, attnOut, seqLen, seqLen, numHeads, numKvHeads, headDim, 0, _threadPool, _slidingWindowSize); } diff --git a/tests/DotLLM.Tests.Integration/Models/Architectures/DirectKvWriteParityTests.cs b/tests/DotLLM.Tests.Integration/Models/Architectures/DirectKvWriteParityTests.cs new file mode 100644 index 00000000..b6178e0d --- /dev/null +++ b/tests/DotLLM.Tests.Integration/Models/Architectures/DirectKvWriteParityTests.cs @@ -0,0 +1,304 @@ +using DotLLM.Core.Attention; +using DotLLM.Core.Tensors; +using DotLLM.Engine.KvCache; +using DotLLM.Models.Architectures; +using DotLLM.Models.Gguf; +using DotLLM.Tests.Integration.Fixtures; +using DotLLM.Tokenizers.Bpe; +using Xunit; + +namespace DotLLM.Tests.Integration.Models.Architectures; + +/// +/// Parity tests for the direct-to-cache K/V write path (issue #25 item 4). +/// +/// +/// The optimisation lets the K and V projection GEMMs (and the subsequent in-place +/// AddBias / QK-norm / RoPE pipeline) write straight into the KV-cache slot via +/// / , skipping +/// the scratch buffer and the Update memcpy. This test exercises both paths +/// against the SmolLM-135M Q8_0 model and asserts byte-identical logits and KV-cache +/// state — proving the optimisation is a pure copy elimination with no math change. +/// +/// +/// +/// The legacy path is forced via , a decorator +/// that intercepts and returns false, +/// pushing the caller back onto the Update branch in +/// . +/// +/// +[Collection("SmallModel")] +public class DirectKvWriteParityTests +{ + private readonly SmallModelFixture _fixture; + + public DirectKvWriteParityTests(SmallModelFixture fixture) + { + _fixture = fixture; + } + + private (TransformerModel model, GgufFile gguf, BpeTokenizer tokenizer) LoadModel() + { + var gguf = GgufFile.Open(_fixture.FilePath); + var config = GgufModelConfigExtractor.Extract(gguf.Metadata); + var model = TransformerModel.LoadFromGguf(gguf, config); + var tokenizer = GgufBpeTokenizerFactory.Load(gguf.Metadata); + return (model, gguf, tokenizer); + } + + /// + /// Prefill parity: a single forward over an N-token prompt with both caches must + /// produce byte-identical logits and byte-identical KV-cache buffers. + /// + [Fact] + public void Prefill_DirectToCache_MatchesLegacyUpdate_BitExact() + { + var (model, gguf, tokenizer) = LoadModel(); + using var _ = gguf; + using var __ = model; + + int[] tokenIds = tokenizer.Encode("The capital of France is"); + int[] positions = new int[tokenIds.Length]; + for (int i = 0; i < positions.Length; i++) positions[i] = i; + + int cacheSize = tokenIds.Length + 8; + + using var directCache = new SimpleKvCache( + model.Config.NumLayers, model.Config.NumKvHeads, model.Config.HeadDim, cacheSize); + using var legacyInner = new SimpleKvCache( + model.Config.NumLayers, model.Config.NumKvHeads, model.Config.HeadDim, cacheSize); + using var legacyCache = new LegacyUpdateOnlyCache(legacyInner); + + using ITensor directLogits = model.Forward(tokenIds, positions, -1, directCache); + using ITensor legacyLogits = model.Forward(tokenIds, positions, -1, legacyCache); + + AssertLogitsByteEqual(directLogits, legacyLogits); + AssertKvCacheByteEqual(directCache, legacyInner, model.Config.NumLayers); + } + + /// + /// Decode parity: prefill + several single-token decode steps under both caches + /// must produce byte-identical decode-step logits and byte-identical KV state. + /// This is the case the optimisation primarily targets — every decode step would + /// otherwise pay a kvStride * 4-byte memcpy per layer per token. + /// + [Fact] + public void Decode_DirectToCache_MatchesLegacyUpdate_BitExact() + { + var (model, gguf, tokenizer) = LoadModel(); + using var _ = gguf; + using var __ = model; + + int[] promptIds = tokenizer.Encode("The capital of France is"); + int numDecodeSteps = 4; + int cacheSize = promptIds.Length + numDecodeSteps; + + int[] positions = new int[cacheSize]; + for (int i = 0; i < cacheSize; i++) positions[i] = i; + + using var directCache = new SimpleKvCache( + model.Config.NumLayers, model.Config.NumKvHeads, model.Config.HeadDim, cacheSize); + using var legacyInner = new SimpleKvCache( + model.Config.NumLayers, model.Config.NumKvHeads, model.Config.HeadDim, cacheSize); + using var legacyCache = new LegacyUpdateOnlyCache(legacyInner); + + int vocabSize = model.Config.VocabSize; + + // Prefill both caches. + int firstDirect, firstLegacy; + using (ITensor d = model.Forward(promptIds, positions.AsSpan(0, promptIds.Length), -1, directCache)) + using (ITensor l = model.Forward(promptIds, positions.AsSpan(0, promptIds.Length), -1, legacyCache)) + { + AssertLogitsByteEqual(d, l); + firstDirect = ArgMaxLast(d, promptIds.Length, vocabSize); + firstLegacy = ArgMaxLast(l, promptIds.Length, vocabSize); + Assert.Equal(firstLegacy, firstDirect); + } + AssertKvCacheByteEqual(directCache, legacyInner, model.Config.NumLayers); + + int nextDirect = firstDirect; + int nextLegacy = firstLegacy; + + // Decode steps: each step is a single-token forward at position prompt + step. + // After every step both caches must be byte-identical and both logits buffers + // must match exactly. + for (int step = 0; step < numDecodeSteps - 1; step++) + { + int pos = promptIds.Length + step; + using ITensor d = model.Forward([nextDirect], positions.AsSpan(pos, 1), -1, directCache); + using ITensor l = model.Forward([nextLegacy], positions.AsSpan(pos, 1), -1, legacyCache); + + AssertLogitsByteEqual(d, l); + AssertKvCacheByteEqual(directCache, legacyInner, model.Config.NumLayers); + + unsafe + { + nextDirect = ArgMax(new ReadOnlySpan((void*)d.DataPointer, vocabSize)); + nextLegacy = ArgMax(new ReadOnlySpan((void*)l.DataPointer, vocabSize)); + } + Assert.Equal(nextLegacy, nextDirect); + } + } + + /// + /// Confirms TryReserveSlot is actually exercised end-to-end. If wiring regresses + /// and the model never calls TryReserveSlot, the parity test would still pass + /// trivially (Update would run on both paths). This counter ensures we actually + /// took the direct-to-cache branch. + /// + [Fact] + public void TryReserveSlot_IsActuallyCalled_FromTransformerModel() + { + var (model, gguf, tokenizer) = LoadModel(); + using var _ = gguf; + using var __ = model; + + int[] promptIds = tokenizer.Encode("Hello"); + int[] positions = new int[promptIds.Length]; + for (int i = 0; i < positions.Length; i++) positions[i] = i; + + using var inner = new SimpleKvCache( + model.Config.NumLayers, model.Config.NumKvHeads, model.Config.HeadDim, promptIds.Length + 1); + using var counting = new CountingCache(inner); + + using var _logits = model.Forward(promptIds, positions, -1, counting); + + // SimpleKvCache reserves contiguous positions starting at 0 — must succeed + // for every layer of the prefill. + Assert.Equal(model.Config.NumLayers, counting.TryReserveSucceededCount); + Assert.Equal(0, counting.UpdateCallCount); + Assert.Equal(model.Config.NumLayers, counting.CommitSlotCount); + } + + // ── Helpers ──────────────────────────────────────────────────────── + + private static unsafe void AssertLogitsByteEqual(ITensor a, ITensor b) + { + Assert.Equal(a.ElementCount, b.ElementCount); + int bytes = (int)a.ElementCount * sizeof(float); + var sa = new ReadOnlySpan((void*)a.DataPointer, bytes); + var sb = new ReadOnlySpan((void*)b.DataPointer, bytes); + Assert.True(sa.SequenceEqual(sb), "Logits must be byte-identical between direct-to-cache and legacy paths."); + } + + private static unsafe void AssertKvCacheByteEqual(SimpleKvCache a, SimpleKvCache b, int numLayers) + { + Assert.Equal(a.CurrentLength, b.CurrentLength); + // GetKeysRef returns a TensorRef of shape [CurrentLength, kvStride] — use Dim1 + // for the per-row width rather than referencing internal fields. + var probe = a.GetKeysRef(0); + int floatsPerLayer = probe.Dim0 * probe.Dim1; + int bytesPerLayer = floatsPerLayer * sizeof(float); + for (int layer = 0; layer < numLayers; layer++) + { + var refA = a.GetKeysRef(layer); + var refB = b.GetKeysRef(layer); + var refAv = a.GetValuesRef(layer); + var refBv = b.GetValuesRef(layer); + + var ka = new ReadOnlySpan((void*)refA.DataPointer, bytesPerLayer); + var kb = new ReadOnlySpan((void*)refB.DataPointer, bytesPerLayer); + var va = new ReadOnlySpan((void*)refAv.DataPointer, bytesPerLayer); + var vb = new ReadOnlySpan((void*)refBv.DataPointer, bytesPerLayer); + + Assert.True(ka.SequenceEqual(kb), $"Layer {layer} K buffer must be byte-identical."); + Assert.True(va.SequenceEqual(vb), $"Layer {layer} V buffer must be byte-identical."); + } + } + + private static unsafe int ArgMaxLast(ITensor logits, int seqLen, int vocabSize) + { + float* ptr = (float*)(logits.DataPointer + (long)(seqLen - 1) * vocabSize * sizeof(float)); + return ArgMax(new ReadOnlySpan(ptr, vocabSize)); + } + + private static int ArgMax(ReadOnlySpan values) + { + int best = 0; + float bestVal = values[0]; + for (int i = 1; i < values.Length; i++) + { + if (values[i] > bestVal) + { + bestVal = values[i]; + best = i; + } + } + return best; + } + + /// + /// IKvCache decorator that forces the legacy path by + /// short-circuiting to false. All other + /// operations delegate to the wrapped cache. Used to compare the direct-to-cache + /// optimisation against the pre-#278 baseline behaviour on identical state. + /// + private sealed class LegacyUpdateOnlyCache : IKvCache + { + private readonly IKvCache _inner; + public LegacyUpdateOnlyCache(IKvCache inner) => _inner = inner; + public int CurrentLength => _inner.CurrentLength; + public int MaxLength => _inner.MaxLength; + public void Update(ITensor keys, ITensor values, ReadOnlySpan positions, int layerIndex) => + _inner.Update(keys, values, positions, layerIndex); + public void Update(TensorRef keys, TensorRef values, ReadOnlySpan positions, int layerIndex) => + _inner.Update(keys, values, positions, layerIndex); + public ITensor GetKeys(int layerIndex) => _inner.GetKeys(layerIndex); + public ITensor GetValues(int layerIndex) => _inner.GetValues(layerIndex); + public TensorRef GetKeysRef(int layerIndex) => _inner.GetKeysRef(layerIndex); + public TensorRef GetValuesRef(int layerIndex) => _inner.GetValuesRef(layerIndex); + public void Rollback(int length) => _inner.Rollback(length); + public bool TryReserveSlot(int layerIndex, ReadOnlySpan positions, out Span kDst, out Span vDst) + { + kDst = default; + vDst = default; + return false; // force the legacy Update branch + } + public void CommitSlot(int layerIndex, ReadOnlySpan positions) { /* never called when TryReserveSlot returns false */ } + public void Dispose() { /* outer test owns inner */ } + } + + /// + /// IKvCache decorator that counts TryReserveSlot success vs Update fallback, used + /// to assert that the direct-to-cache branch is actually being exercised end-to-end + /// (so a parity test passing trivially can't mask a wiring regression). + /// + private sealed class CountingCache : IKvCache + { + private readonly IKvCache _inner; + public int TryReserveSucceededCount { get; private set; } + public int UpdateCallCount { get; private set; } + public int CommitSlotCount { get; private set; } + public CountingCache(IKvCache inner) => _inner = inner; + public int CurrentLength => _inner.CurrentLength; + public int MaxLength => _inner.MaxLength; + public void Update(ITensor keys, ITensor values, ReadOnlySpan positions, int layerIndex) + { + UpdateCallCount++; + _inner.Update(keys, values, positions, layerIndex); + } + public void Update(TensorRef keys, TensorRef values, ReadOnlySpan positions, int layerIndex) + { + UpdateCallCount++; + _inner.Update(keys, values, positions, layerIndex); + } + public ITensor GetKeys(int layerIndex) => _inner.GetKeys(layerIndex); + public ITensor GetValues(int layerIndex) => _inner.GetValues(layerIndex); + public TensorRef GetKeysRef(int layerIndex) => _inner.GetKeysRef(layerIndex); + public TensorRef GetValuesRef(int layerIndex) => _inner.GetValuesRef(layerIndex); + public void Rollback(int length) => _inner.Rollback(length); + public bool TryReserveSlot(int layerIndex, ReadOnlySpan positions, out Span kDst, out Span vDst) + { + bool ok = _inner.TryReserveSlot(layerIndex, positions, out kDst, out vDst); + if (ok) TryReserveSucceededCount++; + return ok; + } + public void CommitSlot(int layerIndex, ReadOnlySpan positions) + { + CommitSlotCount++; + _inner.CommitSlot(layerIndex, positions); + } + public void Dispose() { /* outer test owns inner */ } + } +} diff --git a/tests/DotLLM.Tests.Unit/Engine/KvCache/ReserveSlotTests.cs b/tests/DotLLM.Tests.Unit/Engine/KvCache/ReserveSlotTests.cs new file mode 100644 index 00000000..6cf33d51 --- /dev/null +++ b/tests/DotLLM.Tests.Unit/Engine/KvCache/ReserveSlotTests.cs @@ -0,0 +1,437 @@ +using System.Runtime.InteropServices; +using DotLLM.Core.Attention; +using DotLLM.Core.Configuration; +using DotLLM.Core.Tensors; +using DotLLM.Engine.KvCache; +using Xunit; + +namespace DotLLM.Tests.Unit.Engine.KvCache; + +/// +/// Coverage for + . +/// The primitive lets transformer K/V projections write directly into the cache slot, +/// skipping the scratch + Update memcpy. The contract is that the resulting +/// cache state must be byte-identical to the legacy Update path. +/// +public sealed unsafe class ReserveSlotTests +{ + private const int NumLayers = 2; + private const int NumKvHeads = 4; + private const int HeadDim = 8; + private const int KvStride = NumKvHeads * HeadDim; // 32 + + // ── SimpleKvCache ─────────────────────────────────────────────────── + + [Fact] + public void Simple_TryReserveSlot_Contiguous_ReturnsTrueAndExposesInPlaceSlot() + { + const int MaxSeqLen = 16; + using var cache = new SimpleKvCache(NumLayers, NumKvHeads, HeadDim, MaxSeqLen); + + Span positions = stackalloc int[] { 0, 1, 2 }; + bool ok = cache.TryReserveSlot(layerIndex: 0, positions, out var kDst, out var vDst); + + Assert.True(ok); + Assert.Equal(3 * KvStride, kDst.Length); + Assert.Equal(3 * KvStride, vDst.Length); + } + + [Fact] + public void Simple_TryReserveSlot_NonContiguous_ReturnsFalse() + { + const int MaxSeqLen = 16; + using var cache = new SimpleKvCache(NumLayers, NumKvHeads, HeadDim, MaxSeqLen); + + Span positions = stackalloc int[] { 0, 2, 3 }; + bool ok = cache.TryReserveSlot(layerIndex: 0, positions, out var kDst, out var vDst); + + Assert.False(ok); + Assert.True(kDst.IsEmpty); + Assert.True(vDst.IsEmpty); + } + + [Fact] + public void Simple_TryReserveSlot_OutOfRange_ReturnsFalse() + { + const int MaxSeqLen = 16; + using var cache = new SimpleKvCache(NumLayers, NumKvHeads, HeadDim, MaxSeqLen); + + // Run [15..17) exceeds maxSeqLen=16. + Span positions = stackalloc int[] { 15, 16, 17 }; + bool ok = cache.TryReserveSlot(layerIndex: 0, positions, out _, out _); + + Assert.False(ok); + } + + [Fact] + public void Simple_TryReserveSlot_EmptyPositions_ReturnsFalse() + { + const int MaxSeqLen = 16; + using var cache = new SimpleKvCache(NumLayers, NumKvHeads, HeadDim, MaxSeqLen); + + bool ok = cache.TryReserveSlot(layerIndex: 0, ReadOnlySpan.Empty, out _, out _); + Assert.False(ok); + } + + [Fact] + public void Simple_CommitSlot_AdvancesCurrentLength() + { + const int MaxSeqLen = 16; + using var cache = new SimpleKvCache(NumLayers, NumKvHeads, HeadDim, MaxSeqLen); + + Span positions = stackalloc int[] { 0, 1, 2 }; + Assert.True(cache.TryReserveSlot(0, positions, out _, out _)); + cache.CommitSlot(0, positions); + + Assert.Equal(3, cache.CurrentLength); + } + + /// + /// Bit-exact: building the cache via TryReserveSlot+write+CommitSlot produces + /// byte-identical buffers to the legacy scratch+Update path. + /// + [Fact] + public void Simple_ReserveSlot_BitExactWithUpdate_Prefill() + { + const int MaxSeqLen = 16; + const int SeqLen = 6; + + using var cacheUpdate = new SimpleKvCache(NumLayers, NumKvHeads, HeadDim, MaxSeqLen); + using var cacheSlot = new SimpleKvCache(NumLayers, NumKvHeads, HeadDim, MaxSeqLen); + + // Deterministic synthetic K/V. + nint kSrc = (nint)NativeMemory.AlignedAlloc((nuint)(SeqLen * KvStride * sizeof(float)), 64); + nint vSrc = (nint)NativeMemory.AlignedAlloc((nuint)(SeqLen * KvStride * sizeof(float)), 64); + try + { + for (int t = 0; t < SeqLen; t++) + for (int d = 0; d < KvStride; d++) + { + ((float*)kSrc)[t * KvStride + d] = MathF.Sin(t * 0.37f + d * 0.013f); + ((float*)vSrc)[t * KvStride + d] = MathF.Cos(t * 0.41f + d * 0.017f); + } + + int[] positions = [0, 1, 2, 3, 4, 5]; + + // Path A: legacy Update. + for (int layer = 0; layer < NumLayers; layer++) + { + var kRef = new TensorRef(SeqLen, KvStride, DType.Float32, -1, kSrc); + var vRef = new TensorRef(SeqLen, KvStride, DType.Float32, -1, vSrc); + cacheUpdate.Update(kRef, vRef, positions, layer); + } + + // Path B: TryReserveSlot + write + CommitSlot. + for (int layer = 0; layer < NumLayers; layer++) + { + Assert.True(cacheSlot.TryReserveSlot(layer, positions, out var kDst, out var vDst)); + new ReadOnlySpan((void*)kSrc, SeqLen * KvStride).CopyTo(kDst); + new ReadOnlySpan((void*)vSrc, SeqLen * KvStride).CopyTo(vDst); + cacheSlot.CommitSlot(layer, positions); + } + + Assert.Equal(cacheUpdate.CurrentLength, cacheSlot.CurrentLength); + + for (int layer = 0; layer < NumLayers; layer++) + { + var kA = cacheUpdate.GetKeysRef(layer); + var kB = cacheSlot.GetKeysRef(layer); + var vA = cacheUpdate.GetValuesRef(layer); + var vB = cacheSlot.GetValuesRef(layer); + + int floats = SeqLen * KvStride; + AssertBytesEqual(kA.DataPointer, kB.DataPointer, floats); + AssertBytesEqual(vA.DataPointer, vB.DataPointer, floats); + } + } + finally + { + NativeMemory.AlignedFree((void*)kSrc); + NativeMemory.AlignedFree((void*)vSrc); + } + } + + /// + /// Decode pattern: per-step single-token writes via TryReserveSlot must produce + /// byte-identical state to the legacy Update path. + /// + [Fact] + public void Simple_ReserveSlot_BitExactWithUpdate_DecodeSequence() + { + const int MaxSeqLen = 16; + const int Steps = 8; + + using var cacheUpdate = new SimpleKvCache(NumLayers, NumKvHeads, HeadDim, MaxSeqLen); + using var cacheSlot = new SimpleKvCache(NumLayers, NumKvHeads, HeadDim, MaxSeqLen); + + nint kStep = (nint)NativeMemory.AlignedAlloc((nuint)(KvStride * sizeof(float)), 64); + nint vStep = (nint)NativeMemory.AlignedAlloc((nuint)(KvStride * sizeof(float)), 64); + try + { + for (int step = 0; step < Steps; step++) + { + for (int d = 0; d < KvStride; d++) + { + ((float*)kStep)[d] = MathF.Tan((step + 1) * 0.07f + d * 0.003f); + ((float*)vStep)[d] = MathF.Sinh((step + 1) * 0.05f + d * 0.011f); + } + + int[] positions = [step]; + + for (int layer = 0; layer < NumLayers; layer++) + { + var kRef = new TensorRef(1, KvStride, DType.Float32, -1, kStep); + var vRef = new TensorRef(1, KvStride, DType.Float32, -1, vStep); + cacheUpdate.Update(kRef, vRef, positions, layer); + + Assert.True(cacheSlot.TryReserveSlot(layer, positions, out var kDst, out var vDst)); + new ReadOnlySpan((void*)kStep, KvStride).CopyTo(kDst); + new ReadOnlySpan((void*)vStep, KvStride).CopyTo(vDst); + cacheSlot.CommitSlot(layer, positions); + } + } + + Assert.Equal(cacheUpdate.CurrentLength, cacheSlot.CurrentLength); + + for (int layer = 0; layer < NumLayers; layer++) + { + var kA = cacheUpdate.GetKeysRef(layer); + var kB = cacheSlot.GetKeysRef(layer); + var vA = cacheUpdate.GetValuesRef(layer); + var vB = cacheSlot.GetValuesRef(layer); + AssertBytesEqual(kA.DataPointer, kB.DataPointer, Steps * KvStride); + AssertBytesEqual(vA.DataPointer, vB.DataPointer, Steps * KvStride); + } + } + finally + { + NativeMemory.AlignedFree((void*)kStep); + NativeMemory.AlignedFree((void*)vStep); + } + } + + // ── PagedKvCache ──────────────────────────────────────────────────── + + [Fact] + public void Paged_TryReserveSlot_SingleBlock_ReturnsTrue() + { + const int BlockSize = 4; + const int TotalBlocks = 8; + const int MaxSeqLen = 16; + using var pool = new KvBlockPool(NumLayers, NumKvHeads, HeadDim, BlockSize, TotalBlocks); + using var cache = new PagedKvCache(pool, NumLayers, KvStride, MaxSeqLen); + + // Run fits entirely within block 0 (positions 0..2 of blockSize=4). + Span positions = stackalloc int[] { 0, 1, 2 }; + bool ok = cache.TryReserveSlot(0, positions, out var kDst, out var vDst); + + Assert.True(ok); + Assert.Equal(3 * KvStride, kDst.Length); + Assert.Equal(3 * KvStride, vDst.Length); + } + + [Fact] + public void Paged_TryReserveSlot_BlockBoundary_ReturnsFalse() + { + const int BlockSize = 4; + const int TotalBlocks = 8; + const int MaxSeqLen = 16; + using var pool = new KvBlockPool(NumLayers, NumKvHeads, HeadDim, BlockSize, TotalBlocks); + using var cache = new PagedKvCache(pool, NumLayers, KvStride, MaxSeqLen); + + // Run [3,4,5] crosses block 0 → block 1. + Span positions = stackalloc int[] { 3, 4, 5 }; + bool ok = cache.TryReserveSlot(0, positions, out var kDst, out var vDst); + + Assert.False(ok); + Assert.True(kDst.IsEmpty); + Assert.True(vDst.IsEmpty); + } + + [Fact] + public void Paged_TryReserveSlot_SingleTokenDecode_AlwaysFits() + { + const int BlockSize = 4; + const int TotalBlocks = 8; + const int MaxSeqLen = 16; + using var pool = new KvBlockPool(NumLayers, NumKvHeads, HeadDim, BlockSize, TotalBlocks); + using var cache = new PagedKvCache(pool, NumLayers, KvStride, MaxSeqLen); + + // seqLen=1 always fits in any block — every decode position is reservable. + Span positionBuf = stackalloc int[1]; + for (int p = 0; p < MaxSeqLen; p++) + { + positionBuf[0] = p; + Assert.True(cache.TryReserveSlot(0, positionBuf, out var kDst, out var vDst), + $"position {p} should be reservable as a single-token slot"); + Assert.Equal(KvStride, kDst.Length); + Assert.Equal(KvStride, vDst.Length); + } + } + + [Fact] + public void Paged_TryReserveSlot_NonContiguous_ReturnsFalse() + { + const int BlockSize = 4; + const int TotalBlocks = 8; + const int MaxSeqLen = 16; + using var pool = new KvBlockPool(NumLayers, NumKvHeads, HeadDim, BlockSize, TotalBlocks); + using var cache = new PagedKvCache(pool, NumLayers, KvStride, MaxSeqLen); + + Span positions = stackalloc int[] { 0, 2 }; + bool ok = cache.TryReserveSlot(0, positions, out _, out _); + Assert.False(ok); + } + + /// + /// Bit-exact: paged decode sequence built via TryReserveSlot must match the legacy + /// Update path on the data the attention kernel reads through GetKeysRef/GetValuesRef + /// (the staging buffer). + /// + [Fact] + public void Paged_ReserveSlot_BitExactWithUpdate_DecodeSequence() + { + const int BlockSize = 4; + const int TotalBlocks = 8; + const int MaxSeqLen = 16; + const int Steps = 10; + + using var poolA = new KvBlockPool(NumLayers, NumKvHeads, HeadDim, BlockSize, TotalBlocks); + using var poolB = new KvBlockPool(NumLayers, NumKvHeads, HeadDim, BlockSize, TotalBlocks); + using var cacheUpdate = new PagedKvCache(poolA, NumLayers, KvStride, MaxSeqLen); + using var cacheSlot = new PagedKvCache(poolB, NumLayers, KvStride, MaxSeqLen); + + nint kStep = (nint)NativeMemory.AlignedAlloc((nuint)(KvStride * sizeof(float)), 64); + nint vStep = (nint)NativeMemory.AlignedAlloc((nuint)(KvStride * sizeof(float)), 64); + try + { + for (int step = 0; step < Steps; step++) + { + for (int d = 0; d < KvStride; d++) + { + ((float*)kStep)[d] = MathF.Sin((step + 1) * 0.13f + d * 0.007f); + ((float*)vStep)[d] = MathF.Cos((step + 1) * 0.11f + d * 0.005f); + } + int[] positions = [step]; + + for (int layer = 0; layer < NumLayers; layer++) + { + var kRef = new TensorRef(1, KvStride, DType.Float32, -1, kStep); + var vRef = new TensorRef(1, KvStride, DType.Float32, -1, vStep); + cacheUpdate.Update(kRef, vRef, positions, layer); + + Assert.True(cacheSlot.TryReserveSlot(layer, positions, out var kDst, out var vDst)); + new ReadOnlySpan((void*)kStep, KvStride).CopyTo(kDst); + new ReadOnlySpan((void*)vStep, KvStride).CopyTo(vDst); + cacheSlot.CommitSlot(layer, positions); + } + } + + Assert.Equal(cacheUpdate.CurrentLength, cacheSlot.CurrentLength); + + // Compare via the staging-gathered contiguous view (what attention sees). + for (int layer = 0; layer < NumLayers; layer++) + { + var kA = cacheUpdate.GetKeysRef(layer); + var kB = cacheSlot.GetKeysRef(layer); + var vA = cacheUpdate.GetValuesRef(layer); + var vB = cacheSlot.GetValuesRef(layer); + AssertBytesEqual(kA.DataPointer, kB.DataPointer, Steps * KvStride); + AssertBytesEqual(vA.DataPointer, vB.DataPointer, Steps * KvStride); + } + } + finally + { + NativeMemory.AlignedFree((void*)kStep); + NativeMemory.AlignedFree((void*)vStep); + } + } + + /// + /// Prefill: a single multi-token reservation that fits in one block produces + /// byte-identical state to Update. + /// + [Fact] + public void Paged_ReserveSlot_BitExactWithUpdate_SingleBlockPrefill() + { + const int BlockSize = 8; + const int TotalBlocks = 4; + const int MaxSeqLen = 16; + const int SeqLen = 5; // fits in block 0 (size 8) + + using var poolA = new KvBlockPool(NumLayers, NumKvHeads, HeadDim, BlockSize, TotalBlocks); + using var poolB = new KvBlockPool(NumLayers, NumKvHeads, HeadDim, BlockSize, TotalBlocks); + using var cacheUpdate = new PagedKvCache(poolA, NumLayers, KvStride, MaxSeqLen); + using var cacheSlot = new PagedKvCache(poolB, NumLayers, KvStride, MaxSeqLen); + + nint kSrc = (nint)NativeMemory.AlignedAlloc((nuint)(SeqLen * KvStride * sizeof(float)), 64); + nint vSrc = (nint)NativeMemory.AlignedAlloc((nuint)(SeqLen * KvStride * sizeof(float)), 64); + try + { + for (int t = 0; t < SeqLen; t++) + for (int d = 0; d < KvStride; d++) + { + ((float*)kSrc)[t * KvStride + d] = MathF.Sin(t * 0.37f + d * 0.013f); + ((float*)vSrc)[t * KvStride + d] = MathF.Cos(t * 0.41f + d * 0.017f); + } + + int[] positions = [0, 1, 2, 3, 4]; + for (int layer = 0; layer < NumLayers; layer++) + { + var kRef = new TensorRef(SeqLen, KvStride, DType.Float32, -1, kSrc); + var vRef = new TensorRef(SeqLen, KvStride, DType.Float32, -1, vSrc); + cacheUpdate.Update(kRef, vRef, positions, layer); + + Assert.True(cacheSlot.TryReserveSlot(layer, positions, out var kDst, out var vDst)); + new ReadOnlySpan((void*)kSrc, SeqLen * KvStride).CopyTo(kDst); + new ReadOnlySpan((void*)vSrc, SeqLen * KvStride).CopyTo(vDst); + cacheSlot.CommitSlot(layer, positions); + } + + Assert.Equal(cacheUpdate.CurrentLength, cacheSlot.CurrentLength); + for (int layer = 0; layer < NumLayers; layer++) + { + var kA = cacheUpdate.GetKeysRef(layer); + var kB = cacheSlot.GetKeysRef(layer); + var vA = cacheUpdate.GetValuesRef(layer); + var vB = cacheSlot.GetValuesRef(layer); + AssertBytesEqual(kA.DataPointer, kB.DataPointer, SeqLen * KvStride); + AssertBytesEqual(vA.DataPointer, vB.DataPointer, SeqLen * KvStride); + } + } + finally + { + NativeMemory.AlignedFree((void*)kSrc); + NativeMemory.AlignedFree((void*)vSrc); + } + } + + // ── Caches that opt out (default IKvCache fallback) ──────────────── + + [Fact] + public void Quantized_TryReserveSlot_ReturnsFalse_NoSlotExposed() + { + // Quantized caches store quantized rows, not F32 — no in-place slot. + // Default IKvCache implementation returns false. + using var cache = new QuantizedKvCache( + NumLayers, NumKvHeads, HeadDim, maxSeqLen: 16, + keyDType: KvCacheDType.Q8_0, valueDType: KvCacheDType.Q8_0, windowSize: 0); + + IKvCache ikv = cache; + Span positions = stackalloc int[] { 0, 1, 2 }; + bool ok = ikv.TryReserveSlot(0, positions, out var kDst, out var vDst); + + Assert.False(ok); + Assert.True(kDst.IsEmpty); + Assert.True(vDst.IsEmpty); + } + + // ── Helpers ──────────────────────────────────────────────────────── + + private static void AssertBytesEqual(nint a, nint b, int floatCount) + { + var sa = new ReadOnlySpan((void*)a, floatCount * sizeof(float)); + var sb = new ReadOnlySpan((void*)b, floatCount * sizeof(float)); + Assert.True(sa.SequenceEqual(sb), "KV buffers must be byte-identical between Update and ReserveSlot paths."); + } +}