feat(inference): clamp max_tokens to context-window runway (#291)
The stage completion cap (e.g. 24_576) is a ceiling, not a promise. When the rendered prompt + tool schemas fill most of the model window, sending the raw cap makes llama.cpp truncate the prompt from the left — the traced Gemma4 32_767-token blow-up. Compute effectiveMaxTokens = min(cap, contextSize - promptTokens - toolSchemas - templateOverhead - reserve), counting prompt/tool tokens with the model's own tokenizer (char/4 fallback). Live-path only; deterministic replay never calls infer, so no event recording needed. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
+46
-1
@@ -13,6 +13,7 @@ import com.correx.core.inference.TokenUsage
|
|||||||
import com.correx.core.inference.Tokenizer
|
import com.correx.core.inference.Tokenizer
|
||||||
import com.correx.core.inference.ToolCallFunction
|
import com.correx.core.inference.ToolCallFunction
|
||||||
import com.correx.core.inference.ToolCallRequest
|
import com.correx.core.inference.ToolCallRequest
|
||||||
|
import com.correx.core.inference.ToolDefinition
|
||||||
import com.correx.infrastructure.inference.commons.ModelDescriptor
|
import com.correx.infrastructure.inference.commons.ModelDescriptor
|
||||||
import io.ktor.client.HttpClient
|
import io.ktor.client.HttpClient
|
||||||
import io.ktor.client.call.body
|
import io.ktor.client.call.body
|
||||||
@@ -148,6 +149,9 @@ class LlamaCppInferenceProvider(
|
|||||||
|
|
||||||
val tools = request.tools.takeIf { it.isNotEmpty() }
|
val tools = request.tools.takeIf { it.isNotEmpty() }
|
||||||
|
|
||||||
|
// #291: clamp the stage completion cap to the window's real runway (see clampMaxTokens).
|
||||||
|
val effectiveMaxTokens = clampMaxTokens(request.generationConfig.maxTokens, messages, tools)
|
||||||
|
|
||||||
// llama.cpp rejects requests that carry BOTH a custom grammar and tools
|
// llama.cpp rejects requests that carry BOTH a custom grammar and tools
|
||||||
// ("Cannot use custom grammar constraints with tools"). When a stage has tools,
|
// ("Cannot use custom grammar constraints with tools"). When a stage has tools,
|
||||||
// we drop the grammar and rely on post-hoc schema validation + retry to keep the
|
// we drop the grammar and rely on post-hoc schema validation + retry to keep the
|
||||||
@@ -167,7 +171,7 @@ class LlamaCppInferenceProvider(
|
|||||||
messages = messages,
|
messages = messages,
|
||||||
temperature = request.generationConfig.temperature,
|
temperature = request.generationConfig.temperature,
|
||||||
topP = request.generationConfig.topP,
|
topP = request.generationConfig.topP,
|
||||||
maxTokens = request.generationConfig.maxTokens,
|
maxTokens = effectiveMaxTokens,
|
||||||
stopSequences = request.generationConfig.stopSequences,
|
stopSequences = request.generationConfig.stopSequences,
|
||||||
seed = request.generationConfig.seed,
|
seed = request.generationConfig.seed,
|
||||||
topK = request.generationConfig.topK,
|
topK = request.generationConfig.topK,
|
||||||
@@ -222,6 +226,47 @@ class LlamaCppInferenceProvider(
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Per-message chat-template framing (role tags, delimiters) the client can't see but the server
|
||||||
|
// adds; a small fixed estimate is enough headroom for a safety clamp. ponytail: constant, not a
|
||||||
|
// per-family table — bump if a template proves heavier.
|
||||||
|
private val templateTokensPerMessage = 8
|
||||||
|
private val safetyReserveTokens = 512
|
||||||
|
private val minCompletionTokens = 256
|
||||||
|
|
||||||
|
/**
|
||||||
|
* #291: returns min(requestedCap, contextSize − promptTokens − toolSchema − templateOverhead −
|
||||||
|
* reserve), floored so a nearly-full window still asks for *something* rather than a negative or
|
||||||
|
* absurd allowance. Prompt/tool tokens are counted with the model's own tokenizer (one extra
|
||||||
|
* /tokenize round-trip); a char/4 estimate is the fallback if that call fails. Live-path only —
|
||||||
|
* deterministic replay never calls [infer], so nothing here needs to be an event.
|
||||||
|
*/
|
||||||
|
private suspend fun clampMaxTokens(
|
||||||
|
requestedCap: Int,
|
||||||
|
messages: List<LlamaCppChatMessage>,
|
||||||
|
tools: List<ToolDefinition>?,
|
||||||
|
): Int {
|
||||||
|
val countable = buildString {
|
||||||
|
messages.forEach { m ->
|
||||||
|
append(m.role).append('\n')
|
||||||
|
m.content?.let { append(it).append('\n') }
|
||||||
|
m.reasoningContent?.let { append(it).append('\n') }
|
||||||
|
m.toolCalls.forEach { append(it.function.name).append(it.function.arguments).append('\n') }
|
||||||
|
}
|
||||||
|
tools?.let { append(json.encodeToString(it)) }
|
||||||
|
}
|
||||||
|
val promptTokens = runCatching { tokenizer.countTokens(countable) }
|
||||||
|
.getOrElse { countable.length / 4 }
|
||||||
|
val overhead = messages.size * templateTokensPerMessage
|
||||||
|
val runway = descriptor.contextSize - promptTokens - overhead - safetyReserveTokens
|
||||||
|
if (runway < minCompletionTokens) {
|
||||||
|
log.warn(
|
||||||
|
"context runway low: ctx={} prompt~{} overhead={} reserve={} runway={} (cap was {})",
|
||||||
|
descriptor.contextSize, promptTokens, overhead, safetyReserveTokens, runway, requestedCap,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return requestedCap.coerceAtMost(runway).coerceAtLeast(0)
|
||||||
|
}
|
||||||
|
|
||||||
override suspend fun healthCheck(): ProviderHealth = try {
|
override suspend fun healthCheck(): ProviderHealth = try {
|
||||||
val response = httpClient.get("$baseUrl/health")
|
val response = httpClient.get("$baseUrl/health")
|
||||||
if (response.status.value in 200..299) {
|
if (response.status.value in 200..299) {
|
||||||
|
|||||||
+111
@@ -0,0 +1,111 @@
|
|||||||
|
package com.correx.infrastructure.inference.llama.cpp
|
||||||
|
|
||||||
|
import com.correx.core.context.model.ContextEntry
|
||||||
|
import com.correx.core.context.model.ContextLayer
|
||||||
|
import com.correx.core.context.model.ContextPack
|
||||||
|
import com.correx.core.context.model.EntryRole
|
||||||
|
import com.correx.core.events.types.ContextEntryId
|
||||||
|
import com.correx.core.events.types.ContextPackId
|
||||||
|
import com.correx.core.events.types.InferenceRequestId
|
||||||
|
import com.correx.core.events.types.SessionId
|
||||||
|
import com.correx.core.events.types.StageId
|
||||||
|
import com.correx.core.inference.GenerationConfig
|
||||||
|
import com.correx.core.inference.InferenceRequest
|
||||||
|
import com.correx.infrastructure.inference.commons.ModelDescriptor
|
||||||
|
import com.correx.infrastructure.inference.commons.ResidencyMode
|
||||||
|
import io.ktor.client.HttpClient
|
||||||
|
import io.ktor.client.engine.mock.MockEngine
|
||||||
|
import io.ktor.client.engine.mock.respond
|
||||||
|
import io.ktor.client.plugins.contentnegotiation.ContentNegotiation
|
||||||
|
import io.ktor.serialization.kotlinx.json.json
|
||||||
|
import io.ktor.http.ContentType
|
||||||
|
import io.ktor.http.HttpHeaders
|
||||||
|
import io.ktor.http.content.TextContent
|
||||||
|
import io.ktor.http.headersOf
|
||||||
|
import kotlinx.coroutines.runBlocking
|
||||||
|
import org.junit.jupiter.api.Test
|
||||||
|
import kotlin.test.assertEquals
|
||||||
|
|
||||||
|
// #291: the provider must clamp the requested completion cap against the model context window so a
|
||||||
|
// nearly-full prompt never asks for an impossible allowance (the traced Gemma4 truncation).
|
||||||
|
class LlamaCppMaxTokensClampTest {
|
||||||
|
|
||||||
|
private val promptTokenCount = 1000
|
||||||
|
private val reserve = 512
|
||||||
|
private val overheadPerMessage = 8
|
||||||
|
|
||||||
|
private fun request(cap: Int) = InferenceRequest(
|
||||||
|
requestId = InferenceRequestId("req-1"),
|
||||||
|
sessionId = SessionId("s-1"),
|
||||||
|
stageId = StageId("stage-1"),
|
||||||
|
contextPack = ContextPack(
|
||||||
|
id = ContextPackId("pack-1"),
|
||||||
|
sessionId = SessionId("s-1"),
|
||||||
|
stageId = StageId("stage-1"),
|
||||||
|
layers = mapOf(
|
||||||
|
ContextLayer.L1 to listOf(
|
||||||
|
ContextEntry(
|
||||||
|
id = ContextEntryId("e-1"),
|
||||||
|
layer = ContextLayer.L1,
|
||||||
|
content = "do the thing",
|
||||||
|
sourceType = "stagePrompt",
|
||||||
|
sourceId = "e-1",
|
||||||
|
tokenEstimate = 4,
|
||||||
|
role = EntryRole.USER,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
budgetUsed = 4,
|
||||||
|
budgetLimit = 24_576,
|
||||||
|
),
|
||||||
|
generationConfig = GenerationConfig(temperature = 1.0, topP = 1.0, maxTokens = cap),
|
||||||
|
)
|
||||||
|
|
||||||
|
private fun provider(contextSize: Int, capture: (Int) -> Unit): LlamaCppInferenceProvider {
|
||||||
|
val engine = MockEngine { req ->
|
||||||
|
when {
|
||||||
|
req.url.encodedPath.endsWith("/tokenize") -> respond(
|
||||||
|
content = """{"tokens":[${(1..promptTokenCount).joinToString(",")}]}""",
|
||||||
|
headers = headersOf(HttpHeaders.ContentType, ContentType.Application.Json.toString()),
|
||||||
|
)
|
||||||
|
else -> {
|
||||||
|
val body = (req.body as TextContent).text
|
||||||
|
val maxTokens = Regex("\"max_tokens\":(\\d+)").find(body)!!.groupValues[1].toInt()
|
||||||
|
capture(maxTokens)
|
||||||
|
val usage =
|
||||||
|
""""usage":{"prompt_tokens":$promptTokenCount,"completion_tokens":1,"total_tokens":1}"""
|
||||||
|
respond(
|
||||||
|
content = """{"id":"x","choices":[{"message":{"role":"assistant",""" +
|
||||||
|
""""content":"ok"},"finish_reason":"stop"}],$usage}""",
|
||||||
|
headers = headersOf(HttpHeaders.ContentType, ContentType.Application.Json.toString()),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
val descriptor = ModelDescriptor(
|
||||||
|
modelId = "test",
|
||||||
|
modelPath = "/dev/null",
|
||||||
|
residencyMode = ResidencyMode.EPHEMERAL,
|
||||||
|
contextSize = contextSize,
|
||||||
|
)
|
||||||
|
val client = HttpClient(engine) {
|
||||||
|
install(ContentNegotiation) { json(kotlinx.serialization.json.Json { ignoreUnknownKeys = true }) }
|
||||||
|
}
|
||||||
|
return LlamaCppInferenceProvider(descriptor, "http://localhost:10000", client)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test
|
||||||
|
fun `clamps completion cap to remaining context runway`(): Unit = runBlocking {
|
||||||
|
var sent = -1
|
||||||
|
// 1 message -> overhead 8. runway = 2000 - 1000 - 8 - 512 = 480, below the 24_576 cap.
|
||||||
|
provider(contextSize = 2000) { sent = it }.infer(request(cap = 24_576))
|
||||||
|
assertEquals(2000 - promptTokenCount - overheadPerMessage - reserve, sent)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test
|
||||||
|
fun `keeps the stage cap when the window has ample runway`(): Unit = runBlocking {
|
||||||
|
var sent = -1
|
||||||
|
provider(contextSize = 100_000) { sent = it }.infer(request(cap = 24_576))
|
||||||
|
assertEquals(24_576, sent)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user