feat(inference): operator-tunable sampling knobs (top_k/min_p/repeat_penalty) for stage requests

GenerationConfig only carried temperature/top_p/max_tokens/stop/seed. Added nullable
topK/minP/repeatPenalty, serialized to the llama.cpp and OpenAI-compat request bodies
via @EncodeDefault(NEVER) so an unset knob is omitted (the model keeps its own default)
and behavior is unchanged unless the operator opts in.

Surfaced as a new [sampling] config section feeding the default stage GenerationConfig
(the main agentic loop) through TomlWorkflowLoader + ExecutionPlanCompiler; the former
hardcoded temperature=0.7/topP=1.0 stage defaults now come from config. Talkie
chat/narration keep their own generation settings.

Vikunja #46 (task 76) — sampling half.
This commit is contained in:
2026-07-12 17:54:25 +04:00
parent b2c7bbe401
commit d89d4e32a9
13 changed files with 125 additions and 16 deletions
@@ -3,6 +3,7 @@ package com.correx.infrastructure.inference.llama.cpp
import com.correx.core.inference.ChatMessage
import com.correx.core.inference.ToolCallRequest
import com.correx.core.inference.ToolDefinition
import kotlinx.serialization.EncodeDefault
import kotlinx.serialization.SerialName
import kotlinx.serialization.Serializable
@@ -15,6 +16,12 @@ data class ChatCompletionRequest(
@SerialName("max_tokens") val maxTokens: Int,
@SerialName("stop") val stopSequences: List<String> = emptyList(),
val seed: Long? = null,
// EncodeDefault.NEVER overrides the class-level encodeDefaults=true so an unset (null) sampling
// knob is omitted from the JSON entirely, letting the model keep its own default (rather than
// sending "top_k": null, which llama.cpp may reject or misread).
@EncodeDefault(EncodeDefault.Mode.NEVER) @SerialName("top_k") val topK: Int? = null,
@EncodeDefault(EncodeDefault.Mode.NEVER) @SerialName("min_p") val minP: Double? = null,
@EncodeDefault(EncodeDefault.Mode.NEVER) @SerialName("repeat_penalty") val repeatPenalty: Double? = null,
val stream: Boolean = false,
val grammar: String? = null,
val tools: List<ToolDefinition>? = null,
@@ -170,6 +170,9 @@ class LlamaCppInferenceProvider(
maxTokens = request.generationConfig.maxTokens,
stopSequences = request.generationConfig.stopSequences,
seed = request.generationConfig.seed,
topK = request.generationConfig.topK,
minP = request.generationConfig.minP,
repeatPenalty = request.generationConfig.repeatPenalty,
stream = false,
grammar = grammar,
tools = tools,
@@ -0,0 +1,37 @@
package com.correx.infrastructure.inference.llama.cpp
import kotlinx.serialization.encodeToString
import kotlinx.serialization.json.Json
import kotlin.test.Test
import kotlin.test.assertFalse
import kotlin.test.assertTrue
// Guards the EncodeDefault.NEVER behavior on the sampling knobs: with encodeDefaults=true (matching
// the provider's Json), an unset (null) top_k/min_p/repeat_penalty must be OMITTED from the body so
// the model keeps its own default; a set value must appear.
class SamplingRequestSerializationTest {
private val json = Json { encodeDefaults = true }
@Test
fun `unset sampling knobs are omitted from the request body`() {
val body = ChatCompletionRequest(
model = "m", messages = emptyList(), temperature = 0.7, topP = 1.0, maxTokens = 16,
)
val out = json.encodeToString(body)
assertFalse("top_k" in out, out)
assertFalse("min_p" in out, out)
assertFalse("repeat_penalty" in out, out)
}
@Test
fun `set sampling knobs are serialized`() {
val body = ChatCompletionRequest(
model = "m", messages = emptyList(), temperature = 0.7, topP = 1.0, maxTokens = 16,
topK = 40, minP = 0.05, repeatPenalty = 1.1,
)
val out = json.encodeToString(body)
assertTrue("\"top_k\":40" in out, out)
assertTrue("\"min_p\":0.05" in out, out)
assertTrue("\"repeat_penalty\":1.1" in out, out)
}
}
@@ -2,6 +2,7 @@ package com.correx.infrastructure.inference.openai
import com.correx.core.inference.ToolCallRequest
import com.correx.core.inference.ToolDefinition
import kotlinx.serialization.EncodeDefault
import kotlinx.serialization.SerialName
import kotlinx.serialization.Serializable
@@ -20,6 +21,11 @@ data class OpenAiChatCompletionRequest(
@SerialName("max_tokens") val maxTokens: Int,
@SerialName("stop") val stopSequences: List<String>? = null,
val seed: Long? = null,
// Non-standard OpenAI params accepted by local backends (vLLM, llama.cpp --api). EncodeDefault.NEVER
// omits them when unset so a strict endpoint never sees an unknown key unless the operator opts in.
@EncodeDefault(EncodeDefault.Mode.NEVER) @SerialName("top_k") val topK: Int? = null,
@EncodeDefault(EncodeDefault.Mode.NEVER) @SerialName("min_p") val minP: Double? = null,
@EncodeDefault(EncodeDefault.Mode.NEVER) @SerialName("repeat_penalty") val repeatPenalty: Double? = null,
val stream: Boolean = false,
val tools: List<ToolDefinition>? = null,
)
@@ -104,6 +104,9 @@ class OpenAiCompatInferenceProvider(
maxTokens = request.generationConfig.maxTokens,
stopSequences = request.generationConfig.stopSequences.ifEmpty { null },
seed = request.generationConfig.seed,
topK = request.generationConfig.topK,
minP = request.generationConfig.minP,
repeatPenalty = request.generationConfig.repeatPenalty,
stream = false,
tools = tools,
)
@@ -17,6 +17,7 @@ import com.correx.core.events.EventDispatcher
import com.correx.core.events.stores.EventStore
import com.correx.core.inference.CapabilityScore
import com.correx.core.inference.Embedder
import com.correx.core.inference.GenerationConfig
import com.correx.core.inference.InferenceProvider
import com.correx.core.inference.InferenceRouter
import com.correx.core.inference.ModelCapability
@@ -210,10 +211,13 @@ object InfrastructureModule {
),
)
fun createWorkflowLoader(extraKinds: List<ArtifactKind> = emptyList()): WorkflowLoader {
fun createWorkflowLoader(
extraKinds: List<ArtifactKind> = emptyList(),
samplingDefaults: GenerationConfig = GenerationConfig(temperature = 0.7, topP = 1.0, maxTokens = 0),
): WorkflowLoader {
val registry = DefaultArtifactKindRegistry()
extraKinds.forEach { registry.register(it) }
return TomlWorkflowLoader(registry)
return TomlWorkflowLoader(registry, samplingDefaults)
}
fun createPromptLoader(): PromptLoader = FileSystemPromptLoader()
@@ -59,12 +59,6 @@ private const val DEFAULT_STAGE_TOKEN_BUDGET = 16384
// default, or the model is truncated (finishReason=length) mid-artifact — and a degenerating local
// model burns the whole 2048 on garbage (e.g. a `<|channel>thought` repetition loop) before it can
// stop. Mirror the static TomlWorkflowLoader path: pin the completion cap to the stage token budget.
private val DEFAULT_STAGE_GENERATION = GenerationConfig(
temperature = 0.7,
topP = 1.0,
maxTokens = DEFAULT_STAGE_TOKEN_BUDGET,
)
class ExecutionPlanCompiler(
private val registry: ArtifactKindRegistry,
// Names of every registered tool. A stage that references a tool the runtime can't resolve
@@ -79,7 +73,11 @@ class ExecutionPlanCompiler(
// retries (Vikunja #41). Off by default so the compiler's own unit tests see only plan stages;
// the server's freestyle path (Main.kt) turns it on.
private val injectRecovery: Boolean = false,
// Operator sampling defaults for freestyle-compiled stages; maxTokens pinned to the stage budget.
// Default reproduces the former hardcoded temperature=0.7/topP=1.0.
private val samplingDefaults: GenerationConfig = GenerationConfig(temperature = 0.7, topP = 1.0, maxTokens = 0),
) {
private val defaultStageGeneration = samplingDefaults.copy(maxTokens = DEFAULT_STAGE_TOKEN_BUDGET)
private val mapper = JsonMapper.builder()
.addModule(kotlinModule())
.disable(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES)
@@ -150,7 +148,7 @@ class ExecutionPlanCompiler(
autoBuildGate = s.id == autoGateStageId,
semanticReview = s.semanticReview,
tokenBudget = DEFAULT_STAGE_TOKEN_BUDGET,
generationConfig = DEFAULT_STAGE_GENERATION,
generationConfig = defaultStageGeneration,
metadata = mapOf("promptInline" to s.prompt),
)
}
@@ -168,7 +166,7 @@ class ExecutionPlanCompiler(
StageId(RECOVERY_STAGE) to StageConfig(
allowedTools = knownTools.ifEmpty { setOf("file_write", "file_edit", "shell") },
tokenBudget = DEFAULT_STAGE_TOKEN_BUDGET,
generationConfig = DEFAULT_STAGE_GENERATION,
generationConfig = defaultStageGeneration,
metadata = mapOf("role" to "recovery", "promptInline" to RECOVERY_PROMPT),
)
}
@@ -76,6 +76,9 @@ private val mapper = TomlMapper.builder()
class TomlWorkflowLoader(
private val registry: ArtifactKindRegistry = DefaultArtifactKindRegistry(),
// Operator sampling defaults applied to every stage's inference request (maxTokens is still pinned
// per-stage to the token budget). Defaults reproduce the former hardcoded temperature=0.7/topP=1.0.
private val samplingDefaults: GenerationConfig = GenerationConfig(temperature = 0.7, topP = 1.0, maxTokens = 0),
) : WorkflowLoader {
override fun load(path: Path): WorkflowGraph {
val raw = path.readText()
@@ -116,11 +119,7 @@ class TomlWorkflowLoader(
// Propagate the declared token budget to the inference completion cap.
// Without this the StageConfig default (maxTokens=2048) is used, truncating
// larger artifacts (finishReason=length) → invalid JSON → validation failure.
generationConfig = GenerationConfig(
temperature = 0.7,
topP = 1.0,
maxTokens = s.tokenBudget,
),
generationConfig = samplingDefaults.copy(maxTokens = s.tokenBudget),
maxRetries = s.maxRetries,
metadata = s.toMetadata(workflowDir),
)