diff --git a/Sources/SwiftLM/Server.swift b/Sources/SwiftLM/Server.swift index 9823593..8c81887 100644 --- a/Sources/SwiftLM/Server.swift +++ b/Sources/SwiftLM/Server.swift @@ -1948,10 +1948,17 @@ actor PromptCache { /// The generation prompt after it (`assistant\n…`) is not always a prefix of /// its own re-rendered form. Returns nil when the model is not hybrid, the template is /// not ChatML, or there is no history to cache. +/// +/// The attention layers may be KVCacheSimple or, under `--ctx-size`, RotatingKVCache. +/// A ring is safe here even once it has wrapped: this path never trims, it restores the +/// exact snapshot taken at the boundary, and the ring's state + metaState round-trip it. func hybridCacheBoundary(promptTokens: [Int], imStartId: Int?, cache: [KVCache]) -> Int? { guard let imStartId, cache.contains(where: { $0 is MambaCache }), - cache.allSatisfy({ $0 is MambaCache || type(of: $0) == KVCacheSimple.self }), + cache.allSatisfy({ + $0 is MambaCache || type(of: $0) == KVCacheSimple.self + || type(of: $0) == RotatingKVCache.self + }), let boundary = promptTokens.lastIndex(of: imStartId), boundary > 0 else { return nil } return boundary @@ -2294,7 +2301,9 @@ func handleChatCompletion( // ── Hybrid (recurrent + attention) prompt cache ── // Qwen3.5/3.6-style models pair MambaCache (linear attention) with KVCacheSimple - // layers, and the generic path below refuses them: recurrent state cannot be + // layers (RotatingKVCache under --ctx-size; this path must run before the + // sliding-window one, which misses on any MambaCache), and the generic path below + // refuses them: recurrent state cannot be // trimmed, and the onPrefillDone save runs after the first decode token has been // fed. Without this, every agent turn re-prefills the whole conversation. Instead: // resume from an exact cached prefix, prefill to the turn boundary, snapshot there diff --git a/tests/SwiftLMTests/HybridPromptCacheTests.swift b/tests/SwiftLMTests/HybridPromptCacheTests.swift index 6ebda96..b32031e 100644 --- a/tests/SwiftLMTests/HybridPromptCacheTests.swift +++ b/tests/SwiftLMTests/HybridPromptCacheTests.swift @@ -77,4 +77,57 @@ final class HybridPromptCacheTests: XCTestCase { into: [KVCacheSimple(), MambaCache()]) XCTAssertNil(n, "only the boundary snapshot may persist recurrent state") } + + // MARK: - --ctx-size: RotatingKVCache attention layers + + /// A ring fed `n` single tokens; values encode the token index. + private func makeRing(maxSize: Int, tokens n: Int) -> RotatingKVCache { + let ring = RotatingKVCache(maxSize: maxSize, keep: 0, step: 4) + for t in 0 ..< n { + let k = MLXArray([Float(t)]).reshaped([1, 1, 1, 1]) + _ = ring.update(keys: k, values: k) + } + return ring + } + + private func makeMamba() -> MambaCache { + let mamba = MambaCache() + mamba.state = [MLXArray.ones([1, 3, 8]), MLXArray.ones([1, 2, 4, 4])] + return mamba + } + + func testBoundaryAcceptsRotatingAttentionLayers() { + let tokens = [imStart, 1, imStart, 2] + XCTAssertEqual(hybridCacheBoundary(promptTokens: tokens, imStartId: imStart, + cache: [RotatingKVCache(maxSize: 16), MambaCache()]), 2, + "--ctx-size must not turn the hybrid cache off") + } + + func testRestoresWrappedRingExactly() async { + let ring = makeRing(maxSize: 16, tokens: 40) // wrapped + let pc = PromptCache() + await pc.save(tokens: Array(0 ..< 40), cache: [ring, makeMamba()], allowRecurrent: true) + + // Decode on the live ring after saving; the snapshot must not see it. + for t in 40 ..< 50 { + let k = MLXArray([Float(t)]).reshaped([1, 1, 1, 1]) + _ = ring.update(keys: k, values: k) + } + let fresh: [any KVCache] = [RotatingKVCache(maxSize: 16), MambaCache()] + let n = await pc.restoreExactPrefix(newTokens: Array(0 ..< 40) + [999], limit: 40, into: fresh) + + XCTAssertEqual(n, 40) + XCTAssertEqual(fresh[0].offset, 40) + XCTAssertEqual(fresh[0].state[0].asArray(Float.self).sorted(), (24 ..< 40).map { Float($0) }, + "the restored window is the snapshot's, not the live ring's") + } + + func testRotatingHybridMissesWhenPrefixDiverges() async { + let pc = PromptCache() + await pc.save(tokens: Array(0 ..< 8), cache: [makeRing(maxSize: 16, tokens: 8), makeMamba()], + allowRecurrent: true) + let n = await pc.restoreExactPrefix(newTokens: [0, 1, 9, 3, 4, 5, 6, 7, 8], limit: 9, + into: [RotatingKVCache(maxSize: 16), MambaCache()]) + XCTAssertNil(n) + } }