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:
+7
@@ -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,
|
||||
|
||||
+3
@@ -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,
|
||||
|
||||
+37
@@ -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)
|
||||
}
|
||||
}
|
||||
+6
@@ -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,
|
||||
)
|
||||
|
||||
+3
@@ -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()
|
||||
|
||||
+6
-8
@@ -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),
|
||||
)
|
||||
}
|
||||
|
||||
+4
-5
@@ -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),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user