From 44d2170f3b45f8888cd1de1ebcfb060911713f4f Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Fri, 25 Sep 2026 19:04:58 +0200 Subject: [PATCH 01/10] feat(`LlamaChatSession`): provide an optional document in `decide` --- docs/guide/structured-decisions.md | 5 +- src/evaluator/LlamaChat/LlamaChat.ts | 15 ++- .../utils/prepareDecisionContextWindow.ts | 23 +++-- .../LlamaChatSession/LlamaChatSession.ts | 12 +++ .../generateCriteriaChoiceOptionTokens.ts | 2 +- .../gemma4-e2b/structuredDecisions.test.ts | 93 ++++++++++++++----- 6 files changed, 119 insertions(+), 31 deletions(-) diff --git a/docs/guide/structured-decisions.md b/docs/guide/structured-decisions.md index 1458eb91..cb20bfb9 100644 --- a/docs/guide/structured-decisions.md +++ b/docs/guide/structured-decisions.md @@ -92,7 +92,8 @@ There are two places where you can use the structured decisions API: ### On a Decision Context {#decision-context} -When using a [`LlamaDecisionContext`](../api/classes/LlamaDecisionContext.md), the document you provide as context is only evaluated once, +When using a [`LlamaDecisionContext`](../api/classes/LlamaDecisionContext.md) (via [`.decide()`](../api/classes/LlamaDecisionContext.md#decide)), +the document you provide as context is only evaluated once, and then all questions are evaluated in parallel (up to the configured parallelism limit). It's recommended to configure the [`contextSize`](../api/type-aliases/LlamaDecisionContextOptions.md#contextsize) to limit its size if you only expect short documents. @@ -122,7 +123,7 @@ const context = await model.createDecisionContext({ await context.warmup(); // optional, makes timing the next decision more accurate const startTime = Date.now(); -const ticket = "I still can't sign in after resetting my password. My whole team is locked out."; +const ticket = "I can't sign in after resetting my password. My whole team is locked out."; const answers = await context.decide(ticket, { troubleshootingAttempted: { type: "noul", diff --git a/src/evaluator/LlamaChat/LlamaChat.ts b/src/evaluator/LlamaChat/LlamaChat.ts index d792c72c..acad1546 100644 --- a/src/evaluator/LlamaChat/LlamaChat.ts +++ b/src/evaluator/LlamaChat/LlamaChat.ts @@ -473,6 +473,17 @@ export type LLamaChatLoadAndCompleteUserMessageOptions> { const { + document, evaluationPriority = defaultEvaluationPriority, contextShift = defaultContextShiftOptions, functions, @@ -1202,7 +1214,8 @@ export class LlamaChat { chatWrapper: this._chatWrapper, sequence: this.sequence, functions, - documentFunctionParams + documentFunctionParams, + injectedDocument: document }); const answers: {[key: string]: DecisionAnswer} = {} as DecisionAnswers; diff --git a/src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts b/src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts index 522b62ee..99b22e93 100644 --- a/src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts +++ b/src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts @@ -11,12 +11,14 @@ export async function prepareDecisionContextWindow({ fullHistory, lastEvaluationContextWindowHistory, resolvedContextShift, fallbackToDefaultContextShiftStrategy = true, fitInContextSize, chatWrapper, sequence, - functions, documentFunctionParams, minFreeContextTokens = 1 + functions, documentFunctionParams, minFreeContextTokens = 1, + injectedDocument }: { fullHistory: ChatHistoryItem[], lastEvaluationContextWindowHistory?: ChatHistoryItem[], resolvedContextShift: false | Required, fallbackToDefaultContextShiftStrategy?: boolean, fitInContextSize: number, chatWrapper: ChatWrapper, sequence: LlamaContextSequence, - functions?: ChatModelFunctions, documentFunctionParams?: boolean, minFreeContextTokens?: number + functions?: ChatModelFunctions, documentFunctionParams?: boolean, minFreeContextTokens?: number, + injectedDocument?: string }): Promise<{ prefix: LlamaText, afterQuestion: LlamaText, @@ -27,7 +29,7 @@ export async function prepareDecisionContextWindow({ const context = sequence.context; function generateResponseForChatHistory(contextWindowChatHistory: ChatHistoryItem[], compressionMetadata: object | null | undefined) { - const questionContext = addQuestionMarkerToContextWindow(chatWrapper, contextWindowChatHistory); + const questionContext = addQuestionMarkerToContextWindow(chatWrapper, contextWindowChatHistory, injectedDocument); const questionContextState = chatWrapper.generateContextState({ chatHistory: questionContext.contextWindow, availableFunctions: functions, @@ -174,7 +176,7 @@ function getEvaluationTextParts(contextText: LlamaText, questionMarker: string, }; } -function addQuestionMarkerToContextWindow(chatWrapper: ChatWrapper, contextWindow: ChatHistoryItem[]): { +function addQuestionMarkerToContextWindow(chatWrapper: ChatWrapper, contextWindow: ChatHistoryItem[], injectedDocument?: string): { contextWindow: ChatHistoryItem[], questionMarker: string, decisionMarker: string @@ -218,11 +220,11 @@ function addQuestionMarkerToContextWindow(chatWrapper: ChatWrapper, contextWindo contextWindow: lastItem?.type === "user" ? [...contextWindow.slice(0, -1), { ...lastItem, - text: lastItem.text + "\n\n" + questionMarker + text: addDocumentToUserMessageText(injectedDocument, lastItem.text + "\n\n" + questionMarker) }, modelMessage] : [...contextWindow, { type: "user", - text: questionMarker + text: addDocumentToUserMessageText(injectedDocument, questionMarker) }, modelMessage] }; } @@ -254,3 +256,12 @@ function splitLlamaTextByLastRegularTextMatch(llamaText: LlamaText, textToMatch: return null; } + +function addDocumentToUserMessageText(document: string | undefined, userMessage: string) { + if (document == null || document.trim() === "") + return userMessage; + else if (userMessage.trim() === "") + return document ?? userMessage; + + return document + "\n\n" + userMessage; +} diff --git a/src/evaluator/LlamaChatSession/LlamaChatSession.ts b/src/evaluator/LlamaChatSession/LlamaChatSession.ts index 68019178..0c918920 100644 --- a/src/evaluator/LlamaChatSession/LlamaChatSession.ts +++ b/src/evaluator/LlamaChatSession/LlamaChatSession.ts @@ -415,6 +415,16 @@ export type LLamaChatPreloadPromptOptions = { }; export type LlamaChatSessionDecideOptions = { + /** + * An optional document to add to the context window before a question. + * Only added for generating decisions in this call; won't be appended to the actual chat history. + * + * Will be put in the same user message as the question, with `"\n\n"` in between the document and the question. + * + * Note that a long document can incur a context shift, so make sure to not use a too big document. + */ + document?: string, + signal?: LLamaChatCompletePromptOptions["signal"], evaluationPriority?: LLamaChatCompletePromptOptions["evaluationPriority"], @@ -1276,6 +1286,7 @@ export class LlamaChatSession { } }> { const { + document, signal, evaluationPriority, functions, @@ -1306,6 +1317,7 @@ export class LlamaChatSession { answers, tokenUsage } = await this._chat.generateDecisions(this._chatHistory, questions, { + document, signal, evaluationPriority, functions, diff --git a/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts b/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts index 53a7d1d1..1768e512 100644 --- a/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts +++ b/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts @@ -12,8 +12,8 @@ export function generateCriteriaChoiceOptionTokens(count: number, model: LlamaMo const res = new Set(); const ranges = - "AZ" + "09" + + "AZ" + "\u03b1\u03c1\u03c3\u03c9" + // greek symbols "\u0531\u0556" + // hy "\u05d0\u05d9\u05db\u05dc\u05de\u05de\u05e0\u05e2\u05e4\u05e4\u05e6\u05ea" + // he diff --git a/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts b/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts index c71e9b78..6b579eaa 100644 --- a/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts +++ b/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts @@ -98,8 +98,8 @@ describe("gemma4 e2b", () => { "confidence": 1, "probabilities": { "API": 1, - "codebase": 0.0000237, - "database": 0.00000558, + "codebase": 0.000028, + "database": 9.01e-7, }, "type": "choice", }, @@ -108,10 +108,10 @@ describe("gemma4 e2b", () => { "value": 1, }, "level": { - "confidence": 0.999, + "confidence": 1, "probabilities": [ - 0.0000978, - 0.0002, + 7.93e-8, + 1.66e-7, 1, ], "score": 2, @@ -122,8 +122,8 @@ describe("gemma4 e2b", () => { "confidence": 1, "probabilities": { "engineering": 1, - "hr": 0.000214, - "sales": 0.000036, + "hr": 0.0000956, + "sales": 0.0000154, }, "type": "choice", }, @@ -255,13 +255,12 @@ describe("gemma4 e2b", () => { contextSequence: context.getSequence() }); - const chatResponse = await chat.prompt("Tell me about llamas and where they are from, and a related animal", { + await chat.prompt("Tell me about llamas and where they are from, and a related animal", { maxTokens: 100, budgets: { thoughtTokens: 40 } }); - console.log(chatResponse); const res = await chat.decide({ animal: { @@ -306,7 +305,7 @@ describe("gemma4 e2b", () => { { "animal": { "type": "noul", - "value": 1, + "value": 0.999, }, "animalOrigins": { "type": "noul", @@ -322,7 +321,7 @@ describe("gemma4 e2b", () => { }, "relatedAnimals": { "type": "noul", - "value": 1, + "value": 0.998, }, } `); @@ -348,13 +347,12 @@ describe("gemma4 e2b", () => { contextSequence: sequence }); - const chatResponse = await chat.prompt("Where is wood coming from?", { + await chat.prompt("Where is wood coming from?", { maxTokens: 100, budgets: { thoughtTokens: 40 } }); - console.log(chatResponse); const res = await chat.decide({ animal: { @@ -413,32 +411,32 @@ describe("gemma4 e2b", () => { { "animal": { "type": "noul", - "value": 0.00238, + "value": 0.0000739, }, "cookingRecipe": { "type": "noul", - "value": 0.0000225, + "value": 0.000116, }, "fictionalStory": { "type": "noul", - "value": 0.0000138, + "value": 0.00000194, }, "mineralOrigin": { "type": "noul", - "value": 0.0000348, + "value": 0.0000033, }, "spaceTravel": { "type": "noul", - "value": 0.0000123, + "value": 0.0000856, }, "subject": { "choice": "materials", "confidence": 1, "probabilities": { - "brushing": 0.00000313, - "food": 0.00000662, + "brushing": 0.0000072, + "food": 0.000124, "materials": 1, - "other": 0.0000181, + "other": 0.0000823, }, "type": "choice", }, @@ -449,6 +447,59 @@ describe("gemma4 e2b", () => { } `); }); + + test("with document", {timeout: 1000 * 60 * 60 * 2}, async () => { + const modelPath = await getModelFile("gemma-4-E2B-it-Q4_K_M.gguf"); + const llama = await getTestLlama(); + + const model = await llama.loadModel({ + modelPath + }); + const context = await model.createContext({ + contextSize: 4096 + }); + const chat = new LlamaChatSession({ + contextSequence: context.getSequence() + }); + + await chat.prompt("Tell me about llamas and where they are from, and a related animal", { + maxTokens: 100, + budgets: { + thoughtTokens: 40 + } + }); + + const res1 = await chat.decide({ + locks: { + type: "choice", + instruction: "What are locks?", + criteria: { + useful: "They are useful", + notDoors: "Not doors", + cats: "Not cats" + } + } + }, { + document: "Locks are not doors" + }); + const res2 = await chat.decide({ + locks: { + type: "choice", + instruction: "What are locks?", + criteria: { + useful: "They are useful", + notDoors: "Not doors", + cats: "Not cats" + } + } + }); + + expect(res1.locks.confidence).to.be.greaterThan(0.8); + expect(res1.locks.choice).to.equal("notDoors"); + + expect(res2.locks.confidence).to.be.greaterThan(0.8); + expect(res2.locks.choice).to.equal("useful"); + }); }); }); }); From d8c66de2c6726916a15538eba914e59c996adb01 Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Fri, 25 Sep 2026 21:03:49 +0200 Subject: [PATCH 02/10] fix: prefer uncased characters for options --- src/evaluator/LlamaChat/LlamaChat.ts | 2 +- .../LlamaDecisionContext.ts | 2 +- .../utils/createDecisionAnswer.ts | 18 +++++--- .../utils/createQuestionInputs.ts | 17 +++---- .../generateCriteriaChoiceOptionTokens.ts | 44 ++++++++++++++----- 5 files changed, 54 insertions(+), 29 deletions(-) diff --git a/src/evaluator/LlamaChat/LlamaChat.ts b/src/evaluator/LlamaChat/LlamaChat.ts index acad1546..982caa90 100644 --- a/src/evaluator/LlamaChat/LlamaChat.ts +++ b/src/evaluator/LlamaChat/LlamaChat.ts @@ -1183,7 +1183,7 @@ export class LlamaChat { } = options; return await withLock([this._chatLock, "evaluate"], signal, async (): Promise> => { - const inputs = createQuestionInputs(questions, this.model); + const inputs = createQuestionInputs(questions, this.model.tokenizer); const maxInputLength = Object.values(inputs).reduce((max, item) => Math.max(max, item.input.length), 0); if (maxInputLength > this.sequence.contextSize) throw new Error( diff --git a/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts b/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts index 34474923..13b9ed01 100644 --- a/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts +++ b/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts @@ -229,7 +229,7 @@ export class LlamaDecisionContext { disposeAggregator.add(() => signal.removeEventListener("abort", disposeAggregator.dispose)); } - const inputs = createQuestionInputs(questions, this.model); + const inputs = createQuestionInputs(questions, this.model.tokenizer); const maxInputLength = Object.values(inputs).reduce((max, item) => Math.max(max, item.input.length), 0); if (maxInputLength > this.contextSize) throw new Error( diff --git a/src/evaluator/LlamaDecisionContext/utils/createDecisionAnswer.ts b/src/evaluator/LlamaDecisionContext/utils/createDecisionAnswer.ts index 1e9b69a2..70becd6e 100644 --- a/src/evaluator/LlamaDecisionContext/utils/createDecisionAnswer.ts +++ b/src/evaluator/LlamaDecisionContext/utils/createDecisionAnswer.ts @@ -102,12 +102,12 @@ export function createDecisionAnswer(input: QuestionInput, rawLogits: Map highestProbability) highestProbability = prob; - score += probabilities[i]! * i; + score += prob * i; } const noneDiff = (logits.get(input.tokens[levels]!) ?? 0) - (maxLogit ?? 0) - Math.log(totalWeight); @@ -174,9 +174,17 @@ function getNormalizedInputTokenLogits(input: QuestionInput, logits: Map logit) - res.set(token, destinationLogit); + let alignedLogit = textToLogit.get(text); + const lowercaseText = text.toLowerCase(); + + if (lowercaseText !== text) { + const alignedLogitFromAlignedText = textToLogit.get(lowercaseText); + if (alignedLogitFromAlignedText != null && (alignedLogit == null || alignedLogitFromAlignedText > alignedLogit)) + alignedLogit = alignedLogitFromAlignedText; + } + + if (alignedLogit != null && alignedLogit > logit) + res.set(token, alignedLogit); } return res; diff --git a/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts b/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts index 27c06c96..338fcaba 100644 --- a/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts +++ b/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts @@ -1,21 +1,18 @@ import {LlamaText} from "../../../utils/LlamaText.js"; import {pushAll} from "../../../utils/pushAll.js"; import {generateCriteriaChoiceOptionTokens} from "./generateCriteriaChoiceOptionTokens.js"; -import type {LlamaModel} from "../../../index.js"; -import type {Token} from "../../../types.js"; +import type {Token, Tokenizer} from "../../../types.js"; import type {DecisionQuestions} from "../types.js"; -export function createQuestionInputs(questions: DecisionQuestions, model: LlamaModel) { +export function createQuestionInputs(questions: DecisionQuestions, tokenizer: Tokenizer) { return Object.fromEntries( Object.entries(questions) - .map(([key, question]) => [key, createQuestionInput(key, question, model)]) + .map(([key, question]) => [key, createQuestionInput(key, question, tokenizer)]) ); } export type QuestionInput = ReturnType; -function createQuestionInput(keyName: string, question: DecisionQuestions[number], model: LlamaModel) { - const tokenizer = model.tokenizer; - +function createQuestionInput(keyName: string, question: DecisionQuestions[number], tokenizer: Tokenizer) { if ( (LlamaText.isLlamaText(question.instruction) && question.instruction.values.length === 0) || (typeof question.instruction === "string" && question.instruction.length === 0) || @@ -24,7 +21,7 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number throw new Error(`Question instruction for key "${keyName}" is empty`); if (question.type === "noul") { - const [yesToken, noToken] = generateCriteriaChoiceOptionTokens(2, model); + const [yesToken, noToken] = generateCriteriaChoiceOptionTokens(2, tokenizer); if (yesToken == null || noToken == null) throw new Error("Failed to generate yes/no choice option tokens"); @@ -68,7 +65,7 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number }; } else if (question.type === "choice") { const keys = Object.keys(question.criteria); - const choiceOptions = generateCriteriaChoiceOptionTokens(keys.length, model); + const choiceOptions = generateCriteriaChoiceOptionTokens(keys.length, tokenizer); if (choiceOptions.length < 2) throw new Error('Question with type "choice" must have at least 2 criteria'); @@ -101,7 +98,7 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number }; } else if (question.type === "score") { const additionalChoices = 1; - const scoreOptions = generateCriteriaChoiceOptionTokens(question.criteria.length + additionalChoices, model); + const scoreOptions = generateCriteriaChoiceOptionTokens(question.criteria.length + additionalChoices, tokenizer); if (scoreOptions.length - additionalChoices < 2) throw new Error('Question with type "score" must have at least 2 criteria'); diff --git a/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts b/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts index 1768e512..c2937e6f 100644 --- a/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts +++ b/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts @@ -1,23 +1,23 @@ -import type {Token} from "../../../types.js"; -import type {LlamaModel} from "../../LlamaModel/LlamaModel.js"; +import type {Token, Tokenizer} from "../../../types.js"; const charCode0 = "0".charCodeAt(0); const charCodeA = "A".charCodeAt(0); const charCodeZ = "Z".charCodeAt(0); -export function generateCriteriaChoiceOptionTokens(count: number, model: LlamaModel) { +export function generateCriteriaChoiceOptionTokens(count: number, tokenizer: Tokenizer) { if (count <= 0) return []; const res = new Set(); + const cased = new Set(); const ranges = "09" + "AZ" + - "\u03b1\u03c1\u03c3\u03c9" + // greek symbols + "\u0391\u03a1\u03a3\u03a9" + // greek symbols "\u0531\u0556" + // hy "\u05d0\u05d9\u05db\u05dc\u05de\u05de\u05e0\u05e2\u05e4\u05e4\u05e6\u05ea" + // he - "\u10d0\u10f0" + // ka + "\u1c90\u1cb0" + // ka "\u0915\u0928\u092a\u0930\u0932\u0932\u0935\u0939" + // hi "\u0a95\u0aa8\u0aaa\u0ab0\u0ab2\u0ab2\u0ab5\u0ab9\u0ab3\u0ab3" + // gu "\u0e01\u0e02\u0e04\u0e04\u0e07\u0e23\u0e25\u0e25\u0e27\u0e2e" + // th @@ -65,9 +65,13 @@ export function generateCriteriaChoiceOptionTokens(count: number, model: LlamaMo if (res.size >= count) return true; - const token = findSingleToken(char, model); - if (token != null) - res.add(token); + const token = findSingleToken(char, tokenizer); + if (token != null) { + if (isCasedCharacter(char)) + cased.add(token); + else + res.add(token); + } return res.size >= count; } @@ -130,13 +134,25 @@ export function generateCriteriaChoiceOptionTokens(count: number, model: LlamaMo return false; } + function addPendingCharacters() { + for (const token of cased) { + if (res.size >= count) + break; + + res.add(token); + } + + return res.size >= count; + } + const hasEnoughCharacters = addRanges(ranges) || addDeltaSequence("\u3042", japaneseDeltas) || // Japanese Hiragana addDeltaSequence("\u30a2", japaneseDeltas) || // Japanese Katakana addDeltaSequence("\u3131", koreanConsonantDeltas) || // Korean consonant ordering addDeltaSequence("\uac00", koreanSyllableDeltas, hangulInitialStride) || // Korean syllable ordering addCharacters(cjkOrdinalCharacters) || // CJK Heavenly Stems and Earthly Branches - addRanges(languageDigitRanges); + addRanges(languageDigitRanges) || + addPendingCharacters(); if (!hasEnoughCharacters) throw new RangeError( @@ -148,12 +164,16 @@ export function generateCriteriaChoiceOptionTokens(count: number, model: LlamaMo return [...res]; } -function findSingleToken(text: string, model: LlamaModel) { - const tokens = model.tokenize(text, false, "trimLeadingSpace"); +function findSingleToken(text: string, tokenizer: Tokenizer) { + const tokens = tokenizer(text, false, "trimLeadingSpace"); for (const token of tokens) { - if (model.detokenize([token], false).trim() === text && model._model.getTokenString(token) === text) + if (tokenizer.detokenize([token], false).trim() === text) return token; } return undefined; } + +function isCasedCharacter(text: string): boolean { + return text.toUpperCase() !== text || text.toLowerCase() !== text; +} From 96cdf186b3c34efcfefe7eb285d0c37268b0b74d Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Fri, 25 Sep 2026 22:29:50 +0200 Subject: [PATCH 03/10] perf: sampling --- llama/addon/AddonContext.cpp | 19 +++++++++++++------ llama/addon/AddonContext.h | 7 ++++++- llama/addon/AddonSampler.cpp | 9 +++++++-- llama/addon/AddonSampler.h | 3 ++- 4 files changed, 28 insertions(+), 10 deletions(-) diff --git a/llama/addon/AddonContext.cpp b/llama/addon/AddonContext.cpp index 0338d297..83013ded 100644 --- a/llama/addon/AddonContext.cpp +++ b/llama/addon/AddonContext.cpp @@ -56,6 +56,9 @@ class AddonContextDecodeBatchWorker : public Napi::AsyncWorker { void Execute() { try { + ctx->sharedSamplerData.gotLogit = false; + ctx->sharedSamplerData.hasLogits = false; + // Perform the evaluation using llama_decode. int r = llama_decode(ctx->ctx, ctx->batch); @@ -315,17 +318,21 @@ class AddonContextSampleTokenWorker : public Napi::AsyncWorker { } void SampleToken() { - std::unique_lock samplingLock(ctx->samplingMutex); - if (llama_get_logits(ctx->ctx) == nullptr) { + sampler->rebuildChainIfNeeded(); + + std::unique_lock samplingLock(ctx->sharedSamplerData.mutex); + if (!ctx->sharedSamplerData.gotLogit) { + ctx->sharedSamplerData.hasLogits = llama_get_logits(ctx->ctx) != nullptr; + ctx->sharedSamplerData.gotLogit = true; + } + + if (!ctx->sharedSamplerData.hasLogits) { SetError("This model does not support token generation"); return; } - sampler->rebuildChainIfNeeded(); - llama_token_data_array cur_p; - sampler->sample(ctx->ctx, batchLogitIndex, cur_p, returnProbabilities || returnConfidence || returnLogits.enabled != ReturnLogits::Disabled); - samplingLock.unlock(); + sampler->sampleAndReleaseLock(ctx->sharedSamplerData.mutex, samplingLock, ctx->ctx, batchLogitIndex, cur_p, returnProbabilities || returnConfidence || returnLogits.enabled != ReturnLogits::Disabled); if (cur_p.size == 0 || !(cur_p.selected >= 0 && cur_p.selected < (int32_t)cur_p.size)) { no_output = true; diff --git a/llama/addon/AddonContext.h b/llama/addon/AddonContext.h index f5acd4bb..58dbd84c 100644 --- a/llama/addon/AddonContext.h +++ b/llama/addon/AddonContext.h @@ -22,7 +22,12 @@ class AddonContext : public Napi::ObjectWrap { uint64_t loadedContextMemorySize = 0; bool contextLoaded = false; std::mutex disposeMutex; - std::mutex samplingMutex; + + struct SharedSamplerData { + std::mutex mutex; + bool gotLogit = false; + bool hasLogits = false; + } sharedSamplerData; bool disposed = false; bool memoryDisposed = false; diff --git a/llama/addon/AddonSampler.cpp b/llama/addon/AddonSampler.cpp index bb7c3d90..2286095a 100644 --- a/llama/addon/AddonSampler.cpp +++ b/llama/addon/AddonSampler.cpp @@ -1,4 +1,5 @@ #include +#include #include "common/common.h" #include "globals/addonLog.h" #include "ggml.h" @@ -170,8 +171,9 @@ void AddonSampler::acceptToken(llama_token token) { } } -void AddonSampler::sample(struct llama_context* llamaContext, int32_t batchLogitIndex, llama_token_data_array& curP, bool forceGrammar) { +void AddonSampler::sampleAndReleaseLock(std::mutex & samplingMutex, std::unique_lock & samplingLock, struct llama_context* llamaContext, int32_t batchLogitIndex, llama_token_data_array& curP, bool forceGrammar) { setTokenCandidates(llamaContext, batchLogitIndex, curP); + samplingLock.unlock(); if (curP.size == 0) { return; @@ -206,7 +208,10 @@ void AddonSampler::sample(struct llama_context* llamaContext, int32_t batchLogit return; } - setTokenCandidates(llamaContext, batchLogitIndex, curP); + { + std::unique_lock samplingLock(samplingMutex); + setTokenCandidates(llamaContext, batchLogitIndex, curP); + } llama_sampler_apply(grammarEvaluationState->sampler, &curP); llama_sampler_apply(chain, &curP); diff --git a/llama/addon/AddonSampler.h b/llama/addon/AddonSampler.h index 99cac374..a5fdfbe1 100644 --- a/llama/addon/AddonSampler.h +++ b/llama/addon/AddonSampler.h @@ -1,4 +1,5 @@ #pragma once +#include #include "llama.h" #include "napi.h" #include "RingBuffer.h" @@ -64,7 +65,7 @@ class AddonSampler : public Napi::ObjectWrap { void freeChain(); void rebuildChainIfNeeded(); void acceptToken(llama_token token); - void sample(struct llama_context* llamaContext, int32_t batchLogitIndex, llama_token_data_array& curP, bool forceGrammar); + void sampleAndReleaseLock(std::mutex & samplingMutex, std::unique_lock & samplingLock, struct llama_context* llamaContext, int32_t batchLogitIndex, llama_token_data_array& curP, bool forceGrammar); void setTokenCandidates(struct llama_context* llamaContext, int32_t batchLogitIndex, llama_token_data_array& curP); Napi::Value Dispose(const Napi::CallbackInfo& info); From 79ec6b62ff3eccd44e1f6f68734dfbe87d2948f9 Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Sat, 26 Sep 2026 11:00:50 +0200 Subject: [PATCH 04/10] fix: use a safer char range --- src/evaluator/LlamaChat/LlamaChat.ts | 52 ++- src/evaluator/LlamaContext/LlamaContext.ts | 22 +- .../LlamaDecisionContext.ts | 111 ++++-- src/evaluator/LlamaDecisionContext/types.ts | 4 + .../utils/createDecisionAnswer.ts | 94 +---- .../utils/createQuestionInputs.ts | 149 +++++-- .../utils/evaluateChoiceDecision.ts | 377 ++++++++++++++++++ .../generateCriteriaChoiceOptionTokens.ts | 138 +------ .../gemma4-e2b/structuredDecisions.test.ts | 3 + 9 files changed, 649 insertions(+), 301 deletions(-) create mode 100644 src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts diff --git a/src/evaluator/LlamaChat/LlamaChat.ts b/src/evaluator/LlamaChat/LlamaChat.ts index 982caa90..8059f03b 100644 --- a/src/evaluator/LlamaChat/LlamaChat.ts +++ b/src/evaluator/LlamaChat/LlamaChat.ts @@ -1,4 +1,4 @@ -import {DisposeAggregator, DisposedError, EventRelay, withLock} from "lifecycle-utils"; +import {AsyncQueue, DisposeAggregator, DisposedError, EventRelay, withLock} from "lifecycle-utils"; import {ChatWrapper} from "../../ChatWrapper.js"; import {internalCheckpoints, LlamaContextSequence} from "../LlamaContext/LlamaContext.js"; import { @@ -30,9 +30,10 @@ import {LlamaLogLevel} from "../../bindings/types.js"; import {replaceRegularTextInLlamaText} from "../../chatWrappers/utils/replaceRegularTextInLlamaText.js"; import {DecisionAnswer, DecisionAnswers, DecisionQuestions} from "../LlamaDecisionContext/types.js"; import {createQuestionInputs} from "../LlamaDecisionContext/utils/createQuestionInputs.js"; -import {createEmptyInvalidDecisionAnswer, createDecisionAnswer, decisionAnswerMinimumTopLogits} from "../LlamaDecisionContext/utils/createDecisionAnswer.js"; +import {createEmptyInvalidDecisionAnswer, createDecisionAnswer} from "../LlamaDecisionContext/utils/createDecisionAnswer.js"; import {trimCommonLlamaTextPrefix} from "../../utils/llamaTextUtils.js"; import {TokenMeter} from "../TokenMeter.js"; +import {evaluateChoiceDecision} from "../LlamaDecisionContext/utils/evaluateChoiceDecision.js"; import {FunctionCallNameGrammar} from "./utils/FunctionCallNameGrammar.js"; import {FunctionCallParamsGrammar} from "./utils/FunctionCallParamsGrammar.js"; import {compressHistoryToFitContextSize} from "./utils/compressHistoryToFitContextSize.js"; @@ -1228,7 +1229,7 @@ export class LlamaChat { for (const [questionId, input] of entries) { signal?.throwIfAborted(); - if (input.tokens.length === 0) { + if ((input.type === "choice" && input.keys.length === 0) || (input.type !== "choice" && input.tokens.length === 0)) { answers[questionId] = createEmptyInvalidDecisionAnswer(input); continue; } @@ -1255,33 +1256,44 @@ export class LlamaChat { throw new Error("Not enough tokens to generate a response"); await this.sequence.adaptStateToTokens(fullInput, false); - await this.sequence.evaluateWithoutGeneratingNewTokens(fullInput.slice(this.sequence.nextTokenIndex)); - controlledEvaluateInput = [lastToken]; + controlledEvaluateInput = [...fullInput.slice(this.sequence.nextTokenIndex), lastToken]; } signal?.throwIfAborted(); if (controlledEvaluateInput.length === 0) throw new Error("Evaluate input is empty"); - controlledEvaluateInput[controlledEvaluateInput.length - 1] = [ - controlledEvaluateInput[controlledEvaluateInput.length - 1] as Token, - { - generateNext: { - logits: { - filter: { - tokens: input.tokens, - includeTop: Math.max(input.tokens.length * 2, decisionAnswerMinimumTopLogits) + if (input.type === "choice") { + const choiceSeqQueue = new AsyncQueue([]); + choiceSeqQueue.push(this.sequence); + const answerRes = await evaluateChoiceDecision({ + question: input, + seqQueue: choiceSeqQueue, + signal, + evaluationPriority, + baseSeqTokens: [...prefixTokens, ...input.input, ...afterQuestionTokens, ...input.answerPrefix] + }); + answers[questionId] = answerRes.answer; + } else { + controlledEvaluateInput[controlledEvaluateInput.length - 1] = [ + controlledEvaluateInput[controlledEvaluateInput.length - 1] as Token, + { + generateNext: { + logits: { + filter: { + tokens: input.tokens + } } } } - } - ]; - const res = await this.sequence.controlledEvaluate(controlledEvaluateInput, {evaluationPriority}); - const lastTokenResult = res[res.length - 1]; - if (lastTokenResult == null || lastTokenResult.next?.logits == null) - throw new Error("Failed to generate decisions"); + ]; + const res = await this.sequence.controlledEvaluate(controlledEvaluateInput, {evaluationPriority}); + const lastTokenResult = res[res.length - 1]; + if (lastTokenResult == null || lastTokenResult.next?.logits == null) + throw new Error("Failed to generate decisions"); - answers[questionId] = createDecisionAnswer(input, lastTokenResult.next.logits, this.sequence.model.tokenizer); + answers[questionId] = createDecisionAnswer(input, lastTokenResult.next.logits, this.sequence.model.tokenizer); + } isFirstEvaluation = false; } diff --git a/src/evaluator/LlamaContext/LlamaContext.ts b/src/evaluator/LlamaContext/LlamaContext.ts index f89e9f05..28cd9061 100644 --- a/src/evaluator/LlamaContext/LlamaContext.ts +++ b/src/evaluator/LlamaContext/LlamaContext.ts @@ -60,6 +60,10 @@ export const internalCheckpoints = { decisions: { name: "decisions", maxCheckpoints: 1 + }, + choiceDecision: { + name: "choiceDecision", + maxCheckpoints: 1 } }; @@ -1317,7 +1321,9 @@ export class LlamaContextSequence { await this._eraseContextTokenRanges([{ start: firstDifferentIndex, end: this._nextTokenIndex - }]); + }], { + avoidEvaluation: true + }); return; } @@ -1352,7 +1358,9 @@ export class LlamaContextSequence { }); if (eraseRanges.length > 0) - await this._eraseContextTokenRanges(eraseRanges); + await this._eraseContextTokenRanges(eraseRanges, { + avoidEvaluation: true + }); } /** @@ -1379,11 +1387,13 @@ export class LlamaContextSequence { { canResetTokenPredictor = true, canRemovePredictionTokens = true, - skipLock = false + skipLock = false, + avoidEvaluation = false }: { canResetTokenPredictor?: boolean, canRemovePredictionTokens?: boolean, - skipLock?: boolean + skipLock?: boolean, + avoidEvaluation?: boolean } = {} ) { this._ensureNotDisposed(); @@ -1509,7 +1519,7 @@ export class LlamaContextSequence { this._nextTokenIndex = restoreCheckpointIndex + 1; // wait for the evaluation outside the "context" lock to avoid deadlocks - if (tokensToEvaluate.length > 0) + if (!avoidEvaluation && tokensToEvaluate.length > 0) awaitEvaluationPromise = this.evaluateWithoutGeneratingNewTokens(tokensToEvaluate, {_skipLock: skipLock}); return; } @@ -1521,7 +1531,7 @@ export class LlamaContextSequence { this._contextTokens = []; // wait for the evaluation outside the "context" lock to avoid deadlocks - if (newSequenceTokens.length > 0) + if (!avoidEvaluation && newSequenceTokens.length > 0) awaitEvaluationPromise = this.evaluateWithoutGeneratingNewTokens(newSequenceTokens, {_skipLock: skipLock}); }); diff --git a/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts b/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts index 13b9ed01..1fc2a77f 100644 --- a/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts +++ b/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts @@ -8,7 +8,8 @@ import {ChatWrapper} from "../../ChatWrapper.js"; import {resolveChatWrapper} from "../../chatWrappers/utils/resolveChatWrapper.js"; import {TokenMeter} from "../TokenMeter.js"; import {createQuestionInputs} from "./utils/createQuestionInputs.js"; -import {createDecisionAnswer, decisionAnswerMinimumTopLogits} from "./utils/createDecisionAnswer.js"; +import {createDecisionAnswer} from "./utils/createDecisionAnswer.js"; +import {evaluateChoiceDecision} from "./utils/evaluateChoiceDecision.js"; import type {DecisionAnswer, DecisionAnswers, DecisionQuestions} from "./types.js"; import type {LlamaModel} from "../LlamaModel/LlamaModel.js"; import type {ChatHistoryItem, Token, Tokenizer} from "../../types.js"; @@ -358,28 +359,54 @@ export class LlamaDecisionContext { await mainSeqLease.item.evaluateWithoutGeneratingNewTokens(fullInput.slice(mainSeqLease.item.nextTokenIndex)); signal?.throwIfAborted(); - const controlledEvaluateInput: ControlledEvaluateInputItem[] = [[lastToken, { - generateNext: { - logits: { - filter: { - tokens: input.tokens, - includeTop: Math.max(input.tokens.length * 2, decisionAnswerMinimumTopLogits) + let usedInputTokens = 0; + let usedOutputTokens = 0; + if (input.type === "choice") { + const choiceSeqQueue = new AsyncQueue([], {parent: localQueue}); + choiceSeqQueue.push(mainSeqLease.item); + + mainSeqLease.move(); + const tokenUsageDiff = TokenMeter.diff(mainSeqLease.item.tokenMeter.getState(), mainSeqMeterInitialSnapshot); + usedInputTokens = tokenUsageDiff.usedInputTokens; + usedOutputTokens = tokenUsageDiff.usedOutputTokens; + + const answerRes = await evaluateChoiceDecision({ + question: input, + seqQueue: choiceSeqQueue, + signal, + evaluationPriority, + baseSeqTokens: [...prefixTokens, ...input.input, ...afterQuestionTokens, ...input.answerPrefix] + }); + answers[questionId] = answerRes.answer; + usedInputTokens += answerRes.tokenUsage.input; + usedOutputTokens += answerRes.tokenUsage.output; + } else { + const controlledEvaluateInput: ControlledEvaluateInputItem[] = [[lastToken, { + generateNext: { + logits: { + filter: { + tokens: input.tokens + } } } - } - }]]; - const res = await mainSeqLease.item.controlledEvaluate(controlledEvaluateInput, {evaluationPriority}); - const lastTokenResult = res[res.length - 1]; - if (lastTokenResult == null || lastTokenResult.next?.logits == null) - throw new Error("Failed to generate decisions"); - - answers[questionId] = createDecisionAnswer(input, lastTokenResult.next.logits, mainSeqLease.item.model.tokenizer); - const tokenUsageDiff = TokenMeter.diff(mainSeqLease.item.tokenMeter.getState(), mainSeqMeterInitialSnapshot); + }]]; + const res = await mainSeqLease.item.controlledEvaluate(controlledEvaluateInput, {evaluationPriority}); + const lastTokenResult = res[res.length - 1]; + if (lastTokenResult == null || lastTokenResult.next?.logits == null) + throw new Error("Failed to generate decisions"); + + answers[questionId] = createDecisionAnswer(input, lastTokenResult.next.logits, mainSeqLease.item.model.tokenizer); + + const tokenUsageDiff = TokenMeter.diff(mainSeqLease.item.tokenMeter.getState(), mainSeqMeterInitialSnapshot); + usedInputTokens = tokenUsageDiff.usedInputTokens; + usedOutputTokens = tokenUsageDiff.usedOutputTokens; + } + return { answers: answers as DecisionAnswers, tokenUsage: { - input: tokenUsageDiff.usedInputTokens, - output: tokenUsageDiff.usedOutputTokens + input: usedInputTokens, + output: usedOutputTokens } }; } @@ -548,26 +575,42 @@ export class LlamaDecisionContext { signal?.throwIfAborted(); } - const controlledEvaluateInput: ControlledEvaluateInputItem[] = evaluateInput; - controlledEvaluateInput[controlledEvaluateInput.length - 1] = [ - controlledEvaluateInput[controlledEvaluateInput.length - 1] as Token, - { - generateNext: { - logits: { - filter: { - tokens: input.tokens, - includeTop: Math.max(input.tokens.length * 2, decisionAnswerMinimumTopLogits) + if (input.type === "choice") { + const choiceSeqQueue = new AsyncQueue([], {parent: localQueue}); + choiceSeqQueue.push(seq); + seqLease.move(); + updateTokenUsageExitHandle.call(); + const answerRes = await evaluateChoiceDecision({ + question: input, + seqQueue: choiceSeqQueue, + signal, + evaluationPriority, + baseSeqTokens: [...prefixTokens, ...input.input, ...afterQuestionTokens, ...input.answerPrefix] + }); + answers[questionId] = answerRes.answer; + inputTokens += answerRes.tokenUsage.input; + outputTokens += answerRes.tokenUsage.output; + } else { + const controlledEvaluateInput: ControlledEvaluateInputItem[] = [...evaluateInput]; + controlledEvaluateInput[controlledEvaluateInput.length - 1] = [ + controlledEvaluateInput[controlledEvaluateInput.length - 1] as Token, + { + generateNext: { + logits: { + filter: { + tokens: input.tokens + } } } } - } - ]; - const res = await seq.controlledEvaluate(controlledEvaluateInput, {evaluationPriority}); - const lastTokenResult = res[res.length - 1]; - if (lastTokenResult == null || lastTokenResult.next?.logits == null) - throw new Error("Failed to generate decisions"); + ]; + const res = await seq.controlledEvaluate(controlledEvaluateInput, {evaluationPriority}); + const lastTokenResult = res[res.length - 1]; + if (lastTokenResult == null || lastTokenResult.next?.logits == null) + throw new Error("Failed to generate decisions"); - answers[questionId] = createDecisionAnswer(input, lastTokenResult.next.logits, seq.model.tokenizer); + answers[questionId] = createDecisionAnswer(input, lastTokenResult.next.logits, seq.model.tokenizer); + } }) ); diff --git a/src/evaluator/LlamaDecisionContext/types.ts b/src/evaluator/LlamaDecisionContext/types.ts index cc82704a..7db9fb48 100644 --- a/src/evaluator/LlamaDecisionContext/types.ts +++ b/src/evaluator/LlamaDecisionContext/types.ts @@ -59,6 +59,8 @@ export type DecisionChoiceQuestion = { * * Each key represents a choice, and the value is a description of that choice. * You can set `null` for a choice to reuse its key as the description. + * + * You can set up to 256 choices. * @example * ```ts * { @@ -90,6 +92,8 @@ export type DecisionScoreQuestion = { /** * The criteria describing the possible levels for the score. + * + * You can set up to 10 criteria levels. * @example * ```ts * [ diff --git a/src/evaluator/LlamaDecisionContext/utils/createDecisionAnswer.ts b/src/evaluator/LlamaDecisionContext/utils/createDecisionAnswer.ts index 70becd6e..190c698e 100644 --- a/src/evaluator/LlamaDecisionContext/utils/createDecisionAnswer.ts +++ b/src/evaluator/LlamaDecisionContext/utils/createDecisionAnswer.ts @@ -1,11 +1,13 @@ import type {Token, Tokenizer} from "../../../types.js"; import type {DecisionAnswer, DecisionChoiceAnswer, DecisionNoulAnswer, DecisionScoreAnswer} from "../types.js"; -import type {QuestionInput} from "./createQuestionInputs.js"; +import type {NoulQuestionInput, QuestionInput, ScoreQuestionInput} from "./createQuestionInputs.js"; export const decisionAnswerMinimumTopLogits = 20; -export function createDecisionAnswer(input: QuestionInput, rawLogits: Map, tokenizer: Tokenizer): DecisionAnswer { - const logits = getNormalizedInputTokenLogits(input, rawLogits, tokenizer); - +export function createDecisionAnswer( + input: NoulQuestionInput | ScoreQuestionInput, + logits: Map, + tokenizer: Tokenizer +): DecisionAnswer { if (input.type === "noul") { const [yesToken, noToken] = input.tokens; const yesLogit = logits.get(yesToken); @@ -26,54 +28,6 @@ export function createDecisionAnswer(input: QuestionInput, rawLogits: Map maxLogit) { - secondMaxLogit = maxLogit; - maxLogit = logit; - maxToken = token; - } else if (secondMaxLogit === null || logit > secondMaxLogit) - secondMaxLogit = logit; - } - - const probabilities: Record = {}; - for (const key of input.keys) - probabilities[key] = 0; - - let totalWeight = 0; - let choice: string | null = null; - for (let i = 0; i < input.tokens.length; i++) { - const token = input.tokens[i]!; - const key = input.keys[i]!; - - const logit = logits.get(token) ?? 0; - const weight = Math.exp(logit - (maxLogit ?? 0)); - - probabilities[key] = weight; - totalWeight += weight; - - if (token === maxToken) - choice = key; - } - - for (const key of input.keys) - probabilities[key]! /= totalWeight; - - if (choice == null) - throw new Error("Unable to determine choice"); - - return { - type: "choice", - choice, - confidence: -Math.expm1((secondMaxLogit ?? 0) - (maxLogit ?? 0)) / totalWeight, - probabilities - } satisfies DecisionChoiceAnswer; } else if (input.type === "score") { const additionalChoices = 1; const levels = input.tokens.length - additionalChoices; @@ -153,39 +107,3 @@ export function createEmptyInvalidDecisionAnswer(input: QuestionInput): Decision throw new Error(`Unsupported input type: ${(input as any).type}`); } - -function getNormalizedInputTokenLogits(input: QuestionInput, logits: Map, tokenizer: Tokenizer) { - const res = new Map(); - const textToLogit = new Map(); - - for (const [token, logit] of logits.entries()) { - const text = tokenizer.detokenize([token], false).trim(); - if (text.length !== 1) - continue; - - textToLogit.set(text, logit); - } - - for (const token of input.tokens) { - const logit = logits.get(token) ?? 0; - res.set(token, logit); - - const text = tokenizer.detokenize([token], false).trim(); - if (text.length !== 1) - continue; - - let alignedLogit = textToLogit.get(text); - const lowercaseText = text.toLowerCase(); - - if (lowercaseText !== text) { - const alignedLogitFromAlignedText = textToLogit.get(lowercaseText); - if (alignedLogitFromAlignedText != null && (alignedLogit == null || alignedLogitFromAlignedText > alignedLogit)) - alignedLogit = alignedLogitFromAlignedText; - } - - if (alignedLogit != null && alignedLogit > logit) - res.set(token, alignedLogit); - } - - return res; -} diff --git a/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts b/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts index 338fcaba..1a250a42 100644 --- a/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts +++ b/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts @@ -4,15 +4,17 @@ import {generateCriteriaChoiceOptionTokens} from "./generateCriteriaChoiceOption import type {Token, Tokenizer} from "../../../types.js"; import type {DecisionQuestions} from "../types.js"; +const maxScoreLevels = 10; +const maxChoiceOptions = 256; + export function createQuestionInputs(questions: DecisionQuestions, tokenizer: Tokenizer) { return Object.fromEntries( Object.entries(questions) .map(([key, question]) => [key, createQuestionInput(key, question, tokenizer)]) ); } -export type QuestionInput = ReturnType; -function createQuestionInput(keyName: string, question: DecisionQuestions[number], tokenizer: Tokenizer) { +function createQuestionInput(keyName: string, question: DecisionQuestions[number], tokenizer: Tokenizer): QuestionInput { if ( (LlamaText.isLlamaText(question.instruction) && question.instruction.values.length === 0) || (typeof question.instruction === "string" && question.instruction.length === 0) || @@ -21,9 +23,9 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number throw new Error(`Question instruction for key "${keyName}" is empty`); if (question.type === "noul") { - const [yesToken, noToken] = generateCriteriaChoiceOptionTokens(2, tokenizer); + const [noToken, yesToken] = generateCriteriaChoiceOptionTokens(2, 2, tokenizer); - if (yesToken == null || noToken == null) + if (noToken == null || yesToken == null) throw new Error("Failed to generate yes/no choice option tokens"); let yesLabel: string = (question.criteria == null || question.criteria.true == null || question.criteria.true.trim() === "") @@ -53,19 +55,23 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number "" ]).tokenize(tokenizer, "trimLeadingSpace"), yesToken, - ...LlamaText([": ", yesLabel, "\n"]).tokenize(tokenizer, "trimLeadingSpace"), + ...LlamaText([". ", yesLabel, "\n"]).tokenize(tokenizer, "trimLeadingSpace"), noToken, - ...LlamaText([": ", noLabel]).tokenize(tokenizer, "trimLeadingSpace") + ...LlamaText([". ", noLabel]).tokenize(tokenizer, "trimLeadingSpace") ]; return { - type: "noul" as const, + type: "noul", tokens: [yesToken, noToken] as const, input - }; + } satisfies NoulQuestionInput; } else if (question.type === "choice") { const keys = Object.keys(question.criteria); - const choiceOptions = generateCriteriaChoiceOptionTokens(keys.length, tokenizer); + + if (keys.length > maxChoiceOptions) + throw new Error(`Question with type "choice" cannot have more than ${maxChoiceOptions} criteria`); + + const choiceOptions = generateCriteriaChoiceOptionTokens(2, keys.length, tokenizer); if (choiceOptions.length < 2) throw new Error('Question with type "choice" must have at least 2 criteria'); @@ -73,32 +79,78 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number ...LlamaText.joinValues("\n", [ ["Question: ", question.instruction], "", - "Reply the letter of the best answer:", + "Reply the letters of the best answer:", "" ]).tokenize(tokenizer, "trimLeadingSpace") ]; - for (let i = 0; i < choiceOptions.length; i++) { - input.push(choiceOptions[i]!); - const criteriaKey = keys[i]!; + let nesting = 1; + while (Math.pow(choiceOptions.length, nesting) < keys.length) + nesting++; + + const answerPrefix = (nesting === 1 || choiceOptions[0] == null) + ? [] + : [choiceOptions[0]]; + + let nestedInputs: number = 0; + let keyIndex: number = 0; + const addAnswerTokens = (answerTokens: Map, nestingLeft: number, tokenTrail: Token[] = []) => { + for (let i = 0; i < choiceOptions.length && keyIndex < keys.length; i++) { + const token = choiceOptions[i]!; + + if (nestingLeft !== 1) { + const nestedAnswerTokens = new Map(); + tokenTrail.push(token); + addAnswerTokens(nestedAnswerTokens, nestingLeft - 1, tokenTrail); + tokenTrail.pop(); + + if (nestedAnswerTokens.size > 0) { + answerTokens.set(token, { + type: "input", + answerTokens: nestedAnswerTokens + }); + nestedInputs++; + } + } else { + const criteriaKey = keys[keyIndex]!; + + pushAll(input, tokenTrail); + input.push(token); + pushAll(input, LlamaText([ + ". ", question.criteria[criteriaKey] ?? criteriaKey, + keyIndex === choiceOptions.length - 1 + ? "" + : "\n" + ]).tokenize(tokenizer, "trimLeadingSpace")); + + answerTokens.set(choiceOptions[i]!, { + type: "result", + value: criteriaKey + }); + keyIndex++; + } + } + }; - pushAll(input, LlamaText([ - ": ", question.criteria[criteriaKey] ?? criteriaKey, - i === choiceOptions.length - 1 - ? "" - : "\n" - ]).tokenize(tokenizer, "trimLeadingSpace")); - } + const answerTokens = new Map(); + addAnswerTokens(answerTokens, nesting, [...answerPrefix]); return { - type: "choice" as const, - tokens: choiceOptions, + type: "choice", + answerPrefix, keys, - input - }; + nesting, + answerTokens, + input, + nestedInputs + } satisfies ChoiceQuestionInput; } else if (question.type === "score") { + if (question.criteria.length > maxScoreLevels) + throw new Error(`Question with type "score" cannot have more than ${maxScoreLevels} criteria`); + const additionalChoices = 1; - const scoreOptions = generateCriteriaChoiceOptionTokens(question.criteria.length + additionalChoices, tokenizer); + const requiredOptions = question.criteria.length + additionalChoices; + const scoreOptions = generateCriteriaChoiceOptionTokens(requiredOptions, requiredOptions, tokenizer, "?-!@"); if (scoreOptions.length - additionalChoices < 2) throw new Error('Question with type "score" must have at least 2 criteria'); @@ -118,7 +170,7 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number ? question.criteria[i]! : getAdditionalChoiceOption(i - (scoreOptions.length - additionalChoices)); pushAll(input, LlamaText([ - ": ", text, + ". ", text, i === scoreOptions.length - 1 ? "" : "\n" @@ -129,7 +181,7 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number type: "score" as const, tokens: scoreOptions, input - }; + } satisfies ScoreQuestionInput; } else void (question satisfies never); @@ -142,3 +194,46 @@ function getAdditionalChoiceOption(index: number) { throw new Error(`Unsupported additional choice option index: ${index}`); } + +export function getQuestionInputMaxTokenLength(input: QuestionInput): number { + if (input.type === "noul" || input.type === "score") + return input.input.length; + else if (input.type === "choice") + return input.input.length + input.answerPrefix.length + input.nesting; + else + void (input satisfies never); + + throw new Error(`Unsupported question input type: ${(input as any).type}`); +} + +export type QuestionInput = NoulQuestionInput | ScoreQuestionInput | ChoiceQuestionInput; + +export type NoulQuestionInput = { + type: "noul", + tokens: [yes: Token, no: Token], + input: Token[] +}; + +export type ScoreQuestionInput = { + type: "score", + tokens: Token[], + input: Token[] +}; + +export type ChoiceQuestionInput = { + type: "choice", + answerPrefix: Token[], + keys: string[], + input: Token[], + nesting: number, + answerTokens: Map, + nestedInputs: number +}; + +export type FunnelInput = { + type: "result", + value: string +} | { + type: "input", + answerTokens: Map +}; diff --git a/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts b/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts new file mode 100644 index 00000000..4ff2edf5 --- /dev/null +++ b/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts @@ -0,0 +1,377 @@ +import {AbortablePromise, AsyncQueue, scopeExit} from "lifecycle-utils"; +import {internalCheckpoints, LlamaContextSequence} from "../../LlamaContext/LlamaContext.js"; +import {Token} from "../../../types.js"; +import {ControlledEvaluateInputItem, EvaluationPriority} from "../../LlamaContext/types.js"; +import {TokenMeter} from "../../TokenMeter.js"; +import {pushAll} from "../../../utils/pushAll.js"; +import type {ChoiceQuestionInput, FunnelInput} from "./createQuestionInputs.js"; +import type {DecisionChoiceAnswer} from "../types.js"; + +export async function evaluateChoiceDecision({ + question, + seqQueue: localQueue, + signal, + evaluationPriority, + + baseSeqTokens +}: { + question: ChoiceQuestionInput, + seqQueue: AsyncQueue, + signal?: AbortSignal, + evaluationPriority?: EvaluationPriority, + + baseSeqTokens: Token[] +}) { + const localSeqs = new Set(); + let evaluationsLeft = question.nestedInputs; + let inputTokens: number = 0; + let outputTokens: number = 0; + using localQueueScopeHandle = scopeExit(() => { + if (localQueue.parent == null) + return; + + localQueue.forwardPushesToParent = true; + localQueue.drainToParent(); + }); + + const scores: Record = {}; + const probabilities: Record = {}; + let maxScore: number | null = null; + let maxScoreKey: string | null = null; + let secondMaxScore: number | null = null; + for (const key of question.keys) + probabilities[key] = 0; + + using mainSeqLease = await localQueue.acquire(signal); + const mainSeqMeterInitialSnapshot = mainSeqLease.item.tokenMeter.getState(); + + const needPrefixSeqs = new Map) => void, reject: (reason?: any) => void]>(); + using fixSeqsQueueHandle = scopeExit(() => { + for (const [, [, reject]] of needPrefixSeqs) { + reject(new Error("Disposed")); + } + needPrefixSeqs.clear(); + }); + + async function fixPendingSeqs(seq: LlamaContextSequence) { + if (needPrefixSeqs.size === 0) + return; + + const entriesNeedFixing = [...needPrefixSeqs.entries()]; + needPrefixSeqs.clear(); + await Promise.all( + entriesNeedFixing + .map(async ([otherSeq, [accept, reject]]) => { + try { + const copied = await otherSeq._copyStateFromOtherSequence(seq, baseSeqTokens.length); + if (!copied) + accept( + otherSeq.adaptStateToTokens(baseSeqTokens, false) + .then(() => ( + otherSeq.evaluateWithoutGeneratingNewTokens(baseSeqTokens.slice(otherSeq.nextTokenIndex)) + )) + .then(() => void 0) + ); + else + accept(); + } catch (err) { + reject(err); + } + }) + ); + } + + async function preloadSeqs(seqToCopyFrom: LlamaContextSequence, maxPreloadCount: number) { + const seqsToPreload: LlamaContextSequence[] = []; + for (let i = 0; i < maxPreloadCount; i++) { + const seq = localQueue.tryShift(); + + if (seq == null) + break; + + seqsToPreload.push(seq); + } + using putBackInQueueHandle = scopeExit(() => { + for (const seq of seqsToPreload) + localQueue.push(seq); + + seqsToPreload.length = 0; + }); + + const preloadResults = await Promise.allSettled( + seqsToPreload + .filter((preloadSeq) => !localSeqs.has(preloadSeq)) + .map(async (preloadSeq) => { + const initialMeterSnapshot = preloadSeq.tokenMeter.getState(); + using updateTokenUsageExitHandle = scopeExit(() => { + const diff = TokenMeter.diff(preloadSeq.tokenMeter.getState(), initialMeterSnapshot); + inputTokens += diff.usedInputTokens; + outputTokens += diff.usedOutputTokens; + }); + + const copied = await preloadSeq._copyStateFromOtherSequence(seqToCopyFrom, baseSeqTokens.length); + if (!copied) { + signal?.throwIfAborted(); + await preloadSeq.adaptStateToTokens(baseSeqTokens, false); + await preloadSeq.evaluateWithoutGeneratingNewTokens(baseSeqTokens.slice(preloadSeq.nextTokenIndex)); + } + + localSeqs.add(preloadSeq); + }) + ); + signal?.throwIfAborted(); + + for (const result of preloadResults) { + if (result.status === "rejected") + throw result.reason; + } + } + + async function evaluateFunnel(answerTokens: Map, evaluateInput: Token[], trailScore: number) { + let logits: Map; + { + using seqLease = await localQueue.acquire(signal); + const seq = seqLease.item; + + evaluationsLeft--; + using drainToParentOnFinishHandle = scopeExit(() => { + if (evaluationsLeft > 0 || localQueue.parent == null) + return; + + localQueue.forwardPushesToParent = true; + localQueue.drainToParent(); + }); + + const initialMeterSnapshot = seq.tokenMeter.getState(); + using updateTokenUsageExitHandle = scopeExit(() => { + const diff = TokenMeter.diff(seq.tokenMeter.getState(), initialMeterSnapshot); + inputTokens += diff.usedInputTokens; + outputTokens += diff.usedOutputTokens; + }); + + if (!localSeqs.has(seq)) { + await new AbortablePromise(signal, (accept, reject) => { + needPrefixSeqs.set(seq, [accept, reject]); + + return () => { + needPrefixSeqs.delete(seq); + }; + }); + localSeqs.add(seq); + + await using exitHandle = scopeExit(() => fixPendingSeqs(seq)); + if (needPrefixSeqs.size != 0) { + await fixPendingSeqs(seq); + signal?.throwIfAborted(); + } + } + + if (evaluateInput.length === 0) + throw new Error("Evaluate input is empty"); + + if (seq.nextTokenIndex > baseSeqTokens.length) { + let firstDifferentIndex = baseSeqTokens.length; + for (let i = 0; i < evaluateInput.length - 1; i++) { + if (seq.contextTokens[baseSeqTokens.length + i] !== evaluateInput[i]) + break; + + firstDifferentIndex = baseSeqTokens.length + i + 1; + } + + if (firstDifferentIndex == seq.nextTokenIndex) { + evaluateInput = evaluateInput.slice(firstDifferentIndex - baseSeqTokens.length); + } else { + evaluateInput = [...baseSeqTokens, ...evaluateInput]; + const lastToken = evaluateInput.pop(); + + await seq.adaptStateToTokens(evaluateInput, false); + evaluateInput = evaluateInput.slice(seq.nextTokenIndex); + + if (lastToken != null) + evaluateInput.push(lastToken); + + signal?.throwIfAborted(); + } + } + + const controlledEvaluateInput: ControlledEvaluateInputItem[] = [...evaluateInput]; + controlledEvaluateInput[controlledEvaluateInput.length - 1] = [ + controlledEvaluateInput[controlledEvaluateInput.length - 1] as Token, + { + generateNext: { + logits: { + filter: { + tokens: [...answerTokens.keys()] + } + } + } + } + ]; + + const res = await seq.controlledEvaluate(controlledEvaluateInput, {evaluationPriority}); + const lastTokenResult = res[res.length - 1]; + if (lastTokenResult == null || lastTokenResult.next?.logits == null) + throw new Error("Failed to generate decisions"); + + logits = lastTokenResult.next.logits; + + let totalInputFunnels = 0; + for (const funnel of answerTokens.values()) { + if (funnel.type === "input") + totalInputFunnels++; + } + if (totalInputFunnels !== 0) + await preloadSeqs(seq, totalInputFunnels); + } + + await handleEvaluationResult(answerTokens, evaluateInput, trailScore, logits); + } + + async function handleEvaluationResult( + answerTokens: Map, + evaluateInput: Token[], + trailScore: number, + logits: Map + ) { + let maxLogit: number | null = null; + for (const token of answerTokens.keys()) { + const logit = logits.get(token); + + if (logit == null) + throw new Error("Failed to get logit for token: " + token); + + if (maxLogit === null || logit > maxLogit) + maxLogit = logit; + } + + let totalLogitWeight = 0; + for (const token of answerTokens.keys()) { + const logit = logits.get(token); + + if (logit == null) + throw new Error("Failed to get logit for token: " + token); + + totalLogitWeight += Math.exp(logit - (maxLogit ?? 0)); + } + + const logSumExp = (maxLogit ?? 0) + Math.log(totalLogitWeight); + + const allSettledResults = await Promise.allSettled( + [...answerTokens.entries()].map(async ([token, funnel]) => { + const logit = logits.get(token); + + if (logit == null) + throw new Error("Failed to get logit for token: " + token); + + const tokenScore = logit - logSumExp; + const score = trailScore + tokenScore; + + if (funnel.type === "result") { + const existingValue = scores[funnel.value]; + if (existingValue == null) + scores[funnel.value] = score; + else + scores[funnel.value] = Math.max(score, existingValue); + + if (maxScore == null || score > maxScore) { + secondMaxScore = maxScore; + maxScore = score; + maxScoreKey = funnel.value; + } else if (secondMaxScore == null || score > secondMaxScore) + secondMaxScore = score; + } else if (funnel.type === "input") + await evaluateFunnel(funnel.answerTokens, [...evaluateInput, token], score); + else + void (funnel satisfies never); + }) + ); + + for (const result of allSettledResults) { + if (result.status === "rejected") + throw result.reason; + } + } + + + await mainSeqLease.item.adaptStateToTokens(baseSeqTokens, false); + const controlledEvaluateInput: ControlledEvaluateInputItem[] = baseSeqTokens.slice(mainSeqLease.item.nextTokenIndex); + if (controlledEvaluateInput.length === 0) { + await mainSeqLease.item.eraseContextTokenRanges([{start: baseSeqTokens.length - 1, end: mainSeqLease.item.nextTokenIndex}]); + pushAll(controlledEvaluateInput, baseSeqTokens.slice(mainSeqLease.item.nextTokenIndex)); + } + + if (controlledEvaluateInput.length === 0) + throw new Error("Evaluate input is empty"); + + controlledEvaluateInput[controlledEvaluateInput.length - 1] = [ + controlledEvaluateInput[controlledEvaluateInput.length - 1] as Token, + { + generateNext: { + logits: { + filter: { + tokens: [...question.answerTokens.keys()] + } + } + } + } + ]; + + const res = await mainSeqLease.item.controlledEvaluate(controlledEvaluateInput, {evaluationPriority}); + const lastTokenResult = res[res.length - 1]; + if (lastTokenResult == null || lastTokenResult.next?.logits == null) + throw new Error("Failed to generate decisions"); + + signal?.throwIfAborted(); + + if (question.nesting !== 1) + await mainSeqLease.item._takeNamedCheckpoint( + internalCheckpoints.choiceDecision.name, + internalCheckpoints.choiceDecision.maxCheckpoints + ); + + let totalInputFunnels = 0; + for (const funnel of question.answerTokens.values()) { + if (funnel.type === "input") + totalInputFunnels++; + } + if (totalInputFunnels !== 0) + await preloadSeqs(mainSeqLease.item, totalInputFunnels); + + signal?.throwIfAborted(); + + const mainSeqLeaseTokenUsageDiff = TokenMeter.diff(mainSeqLease.item.tokenMeter.getState(), mainSeqMeterInitialSnapshot); + inputTokens += mainSeqLeaseTokenUsageDiff.usedInputTokens; + outputTokens += mainSeqLeaseTokenUsageDiff.usedOutputTokens; + + localSeqs.add(mainSeqLease.item); + mainSeqLease.dispose(); + + await handleEvaluationResult(question.answerTokens, [], 0, lastTokenResult.next.logits); + signal?.throwIfAborted(); + + let totalScoreWeight = 0; + for (const [key, score] of Object.entries(scores)) { + const weight = Math.exp(score - (maxScore ?? 0)); + totalScoreWeight += weight; + probabilities[key] = weight; + } + + for (const key of Object.keys(scores)) + probabilities[key]! /= totalScoreWeight; + + if (maxScoreKey == null) + throw new Error("Failed to determine the choice decision"); + + return { + answer: { + type: "choice", + choice: maxScoreKey, + confidence: -Math.expm1((secondMaxScore ?? 0) - (maxScore ?? 0)) / totalScoreWeight, + probabilities + } satisfies DecisionChoiceAnswer, + tokenUsage: { + input: inputTokens, + output: outputTokens + } + }; +} diff --git a/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts b/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts index c2937e6f..714c4d07 100644 --- a/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts +++ b/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts @@ -1,79 +1,20 @@ import type {Token, Tokenizer} from "../../../types.js"; -const charCode0 = "0".charCodeAt(0); -const charCodeA = "A".charCodeAt(0); -const charCodeZ = "Z".charCodeAt(0); - -export function generateCriteriaChoiceOptionTokens(count: number, tokenizer: Tokenizer) { - if (count <= 0) +export function generateCriteriaChoiceOptionTokens(min: number, max: number, tokenizer: Tokenizer, additionalChars: string = ""): Token[] { + if (max <= 0) return []; const res = new Set(); - const cased = new Set(); - - const ranges = - "09" + - "AZ" + - "\u0391\u03a1\u03a3\u03a9" + // greek symbols - "\u0531\u0556" + // hy - "\u05d0\u05d9\u05db\u05dc\u05de\u05de\u05e0\u05e2\u05e4\u05e4\u05e6\u05ea" + // he - "\u1c90\u1cb0" + // ka - "\u0915\u0928\u092a\u0930\u0932\u0932\u0935\u0939" + // hi - "\u0a95\u0aa8\u0aaa\u0ab0\u0ab2\u0ab2\u0ab5\u0ab9\u0ab3\u0ab3" + // gu - "\u0e01\u0e02\u0e04\u0e04\u0e07\u0e23\u0e25\u0e25\u0e27\u0e2e" + // th - "\u3105\u3129"; - - const languageDigitRanges = - "\u0966\u096f" + // hi - "\u0a66\u0a6f" + // pa - "\u0ae6\u0aef" + // gu - "\u0b66\u0b6f" + // or - "\u0be6\u0bef" + // ta - "\u0c66\u0c6f" + // te - "\u0ce6\u0cef" + // kn - "\u0d66\u0d6f" + // ml - "\u0e50\u0e59" + // th - "\u0ed0\u0ed9" + // lo - "\u0f20\u0f29" + // bo - "\u1040\u1049" + // my - "\u17e0\u17e9" + // km - "\u1810\u1819"; // mn - - /** - * Delta sequences: - * - [0-9] as the Unicode delta. - * - [A-Z] to repeat the previous delta an additional 1-26 times. - * - * For example, `"2C"` expands to [2, 2, 2, 2] - */ - const japaneseDeltas = "2C12J32B1D3D1C2B1D21C"; - const koreanConsonantDeltas = "3A28132121C"; - const koreanSyllableDeltas = "2121A2A121C"; - - // Hangul syllables are arranged so changing the initial consonant while - // keeping the vowel and final consonant fixed has a constant stride - const hangulInitialStride = - "\uae4c".charCodeAt(0) - // next Hangul initial - "\uac00".charCodeAt(0); // first Hangul initial - - // CJK Heavenly Stems and Earthly Branches - const cjkOrdinalCharacters = - "\u7532\u4e59\u4e19\u4e01\u620a\u5df1\u5e9a\u8f9b\u58ec\u7678" + - "\u5b50\u4e11\u5bc5\u536f\u8fb0\u5df3\u5348\u672a\u7533\u9149\u620c\u4ea5"; function pushCharacter(char: string) { - if (res.size >= count) + if (res.size >= max) return true; const token = findSingleToken(char, tokenizer); - if (token != null) { - if (isCasedCharacter(char)) - cased.add(token); - else - res.add(token); - } + if (token != null) + res.add(token); - return res.size >= count; + return res.size >= max; } function pushCode(code: number) { @@ -94,70 +35,19 @@ export function generateCriteriaChoiceOptionTokens(count: number, tokenizer: Tok return false; } - function addDeltaSequence(start: string, compressedDeltas: string, stride: number = 1) { - let code = start.charCodeAt(0); - let lastDelta = 0; - - if (pushCode(code)) - return false; - - for (let i = 0; i < compressedDeltas.length; i++) { - const encodedDelta = compressedDeltas.charCodeAt(i); - - if (encodedDelta >= charCodeA && encodedDelta <= charCodeZ) { - const repetitions = encodedDelta - charCodeA + 1; - - for (let repetition = 0; repetition < repetitions; repetition++) { - code += lastDelta * stride; - - if (pushCode(code)) - return false; - } - } else { - lastDelta = encodedDelta - charCode0; - code += lastDelta * stride; - - if (pushCode(code)) - return true; - } - } - - return false; - } - - function addCharacters(characters: string) { - for (let i = 0; i < characters.length; i++) { - if (pushCharacter(characters[i]!)) - return true; - } - - return false; - } - - function addPendingCharacters() { - for (const token of cased) { - if (res.size >= count) + if (additionalChars !== "") { + for (const char of additionalChars) { + if (pushCharacter(char)) break; - - res.add(token); } - - return res.size >= count; } - const hasEnoughCharacters = addRanges(ranges) || - addDeltaSequence("\u3042", japaneseDeltas) || // Japanese Hiragana - addDeltaSequence("\u30a2", japaneseDeltas) || // Japanese Katakana - addDeltaSequence("\u3131", koreanConsonantDeltas) || // Korean consonant ordering - addDeltaSequence("\uac00", koreanSyllableDeltas, hangulInitialStride) || // Korean syllable ordering - addCharacters(cjkOrdinalCharacters) || // CJK Heavenly Stems and Earthly Branches - addRanges(languageDigitRanges) || - addPendingCharacters(); + addRanges("09"); - if (!hasEnoughCharacters) + if (res.size < min) throw new RangeError( "Failed to find enough choice options for the given criteria. " + - `${count} options are needed out of ${res.size} that are available. ` + + `${min} options are needed out of ${res.size} that are available. ` + "Reduce the number of criteria or use a different model" ); @@ -173,7 +63,3 @@ function findSingleToken(text: string, tokenizer: Tokenizer) { return undefined; } - -function isCasedCharacter(text: string): boolean { - return text.toUpperCase() !== text || text.toLowerCase() !== text; -} diff --git a/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts b/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts index 6b579eaa..650bc5b2 100644 --- a/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts +++ b/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts @@ -245,6 +245,9 @@ describe("gemma4 e2b", () => { const modelPath = await getModelFile("gemma-4-E2B-it-Q4_K_M.gguf"); const llama = await getTestLlama(); + if (llama.gpu === false) + test.skip("Logits are a bit different on different backends to cause test flakiness"); + const model = await llama.loadModel({ modelPath }); From b35f0f0695b46ced1ec8f4ac29bccf75976a04a5 Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Sat, 26 Sep 2026 11:02:43 +0200 Subject: [PATCH 05/10] test: a choice with many options --- .../utils/evaluateChoiceDecision.ts | 2 +- .../gemma4-e2b/structuredDecisions.test.ts | 156 ++++++++++++++++-- 2 files changed, 140 insertions(+), 18 deletions(-) diff --git a/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts b/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts index 4ff2edf5..e1071f7d 100644 --- a/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts +++ b/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts @@ -366,7 +366,7 @@ export async function evaluateChoiceDecision({ answer: { type: "choice", choice: maxScoreKey, - confidence: -Math.expm1((secondMaxScore ?? 0) - (maxScore ?? 0)) / totalScoreWeight, + confidence: Math.tanh(((maxScore ?? 0) - (secondMaxScore ?? 0)) / 2), probabilities } satisfies DecisionChoiceAnswer, tokenUsage: { diff --git a/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts b/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts index 650bc5b2..c65f33a6 100644 --- a/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts +++ b/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts @@ -98,8 +98,8 @@ describe("gemma4 e2b", () => { "confidence": 1, "probabilities": { "API": 1, - "codebase": 0.000028, - "database": 9.01e-7, + "codebase": 0.0000762, + "database": 0.00000208, }, "type": "choice", }, @@ -110,8 +110,8 @@ describe("gemma4 e2b", () => { "level": { "confidence": 1, "probabilities": [ - 7.93e-8, - 1.66e-7, + 0.0000264, + 0.0000152, 1, ], "score": 2, @@ -122,8 +122,8 @@ describe("gemma4 e2b", () => { "confidence": 1, "probabilities": { "engineering": 1, - "hr": 0.0000956, - "sales": 0.0000154, + "hr": 0.000242, + "sales": 0.0000568, }, "type": "choice", }, @@ -237,6 +237,128 @@ describe("gemma4 e2b", () => { expect(contextText).to.include(longText.slice(0, 64)); expect(contextText).to.not.include(longText.slice(-64)); }); + + test("many choices", {timeout: 1000 * 60 * 60 * 2}, async () => { + const modelPath = await getModelFile("gemma-4-E2B-it-Q4_K_M.gguf"); + const llama = await getTestLlama(); + + const model = await llama.loadModel({ + modelPath + }); + const context = await model.createDecisionContext({ + parallelQuestions: 1, + contextSize: 1024 + }); + + const ticket = "I can't sign in after resetting my password. My whole team is locked out."; + const res = await context.decide(ticket, { + category: { + type: "choice", + instruction: "Which category is related to this ticket?", + criteria: { + food: "Food", + travel: "Travel", + accommodation: "Accommodation", + delivery: "Delivery", + maintenance: "Maintenance", + support: "Support", + billing: "Billing", + payments: "Payments", + subscriptions: "Subscriptions", + privacy: "Privacy", + performance: "Performance", + outage: "Outage", + networking: "Networking", + hardware: "Hardware", + software: "Software", + mobile: "Mobile", + desktop: "Desktop", + website: "Website", + api: "API", + integration: "Integration", + database: "Database", + storage: "Storage", + backup: "Backup", + migration: "Migration", + installation: "Installation", + configuration: "Configuration", + permissions: "Permissions", + notifications: "Notifications", + email: "Email", + messaging: "Messaging", + communication: "Communication", + documentation: "Documentation", + training: "Training", + onboarding: "Onboarding", + cancellation: "Cancellation", + renewal: "Renewal", + pricing: "Pricing", + discount: "Discount", + promotion: "Promotion", + order: "Order", + returns: "Returns", + shipping: "Shipping", + inventory: "Inventory", + product: "Product", + availability: "Availability", + quality: "Quality", + warranty: "Warranty", + repair: "Repair", + replacement: "Replacement", + booking: "Booking", + reservation: "Reservation", + scheduling: "Scheduling", + transportation: "Transportation", + parking: "Parking", + restaurant: "Restaurant", + entertainment: "Entertainment", + events: "Events", + healthcare: "Healthcare", + insurance: "Insurance", + legal: "Legal", + finance: "Finance", + taxes: "Taxes", + employment: "Employment", + payroll: "Payroll", + humanResources: "Human Resources", + education: "Education", + childcare: "Childcare", + pets: "Pets", + utilities: "Utilities", + electricity: "Electricity", + water: "Water", + internet: "Internet", + password: "Password", // password: "Password", + phone: "Phone", + cleaning: "Cleaning", + plumbing: "Plumbing", + heating: "Heating", + cooling: "Cooling", + furniture: "Furniture", + appliances: "Appliances", + construction: "Construction", + gardening: "Gardening", + noise: "Noise", + safety: "Safety", + complaint: "Complaint", + feedback: "Feedback", + suggestion: "Suggestion", + request: "Request", + inquiry: "Inquiry", + incident: "Incident", + fraud: "Fraud", + accessibility: "Accessibility", + localization: "Localization", + // password: "Password", + compliance: "Compliance" + } + } + }); + expect(Object.keys(res.category.probabilities).length).to.be.greaterThan(10); + expect(Object.keys(res.category.probabilities).length).toMatchInlineSnapshot("94"); + expect(res.category.choice).to.be.eql("password"); + expect(res.category.confidence).to.be.greaterThanOrEqual(0.6); + }); }); }); @@ -324,7 +446,7 @@ describe("gemma4 e2b", () => { }, "relatedAnimals": { "type": "noul", - "value": 0.998, + "value": 0.999, }, } `); @@ -414,32 +536,32 @@ describe("gemma4 e2b", () => { { "animal": { "type": "noul", - "value": 0.0000739, + "value": 0.000187, }, "cookingRecipe": { "type": "noul", - "value": 0.000116, + "value": 0.000112, }, "fictionalStory": { "type": "noul", - "value": 0.00000194, + "value": 0.0000023, }, "mineralOrigin": { "type": "noul", - "value": 0.0000033, + "value": 0.0000034, }, "spaceTravel": { "type": "noul", - "value": 0.0000856, + "value": 0.0000368, }, "subject": { "choice": "materials", - "confidence": 1, + "confidence": 0.998, "probabilities": { - "brushing": 0.0000072, - "food": 0.000124, - "materials": 1, - "other": 0.0000823, + "brushing": 0.0000175, + "food": 0.000779, + "materials": 0.999, + "other": 0.00018, }, "type": "choice", }, From 5c9cc086533b50b1876d42b47d6369ebfe1254a0 Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Sat, 26 Sep 2026 11:04:47 +0200 Subject: [PATCH 06/10] docs: note probability variability across runs --- docs/guide/structured-decisions.md | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/docs/guide/structured-decisions.md b/docs/guide/structured-decisions.md index cb20bfb9..005a0dd5 100644 --- a/docs/guide/structured-decisions.md +++ b/docs/guide/structured-decisions.md @@ -192,6 +192,14 @@ In such a case you'll see that the probability of two or more items is pretty cl Try refining the criteria to make it more specific, or add an additional option in order to remove ambiguity. ::: +::: tip NOTE + +The specific values that you get for `confidence` and the rest of the probabilities could slightly vary on each evaluation (depending on your machine and setup), +but the general consensus should remain stable - the `choice` with the highest confidence stays the same, a `noul` stays as decisive as before, +but the exact confidence and probability values may fluctuate. + +::: + ### On a Chat Session {#chat-session} When using structured decisions on a [`LlamaChatSession`](../api/classes/LlamaChatSession.md) (via [`.decide()`](../api/classes/LlamaChatSession.md#decide)), From 584aae2a971556bb3b645b758b7fabd0c87134d7 Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Sat, 26 Sep 2026 11:14:43 +0200 Subject: [PATCH 07/10] chore: update module --- package-lock.json | 8 ++++---- package.json | 2 +- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/package-lock.json b/package-lock.json index 1072acb9..b75db297 100644 --- a/package-lock.json +++ b/package-lock.json @@ -23,7 +23,7 @@ "ignore": "^7.0.4", "ipull": "^3.9.5", "is-unicode-supported": "^2.1.0", - "lifecycle-utils": "^4.5.0", + "lifecycle-utils": "^4.5.2", "log-symbols": "^7.0.1", "nanoid": "^5.1.6", "node-addon-api": "^8.9.1", @@ -10802,9 +10802,9 @@ } }, "node_modules/lifecycle-utils": { - "version": "4.5.0", - "resolved": "https://registry.npmjs.org/lifecycle-utils/-/lifecycle-utils-4.5.0.tgz", - "integrity": "sha512-hFAML1VB8tMBlRjKT6Uub0k2/pEejI3fAs/ekFh+bwr4/SHUVzjlquT+b9ON/55HBreA8XBcTY451IUr5topqQ==", + "version": "4.5.2", + "resolved": "https://registry.npmjs.org/lifecycle-utils/-/lifecycle-utils-4.5.2.tgz", + "integrity": "sha512-3C1g98pNKGmtakPeIfRqvOiiBJflKAvNjo/WsDI/70vsOy/DJDoutsLM+wGuAdfcOitlbUQfl6UcfvyeIimh7g==", "license": "MIT", "funding": { "type": "github", diff --git a/package.json b/package.json index 97c4807e..122da6d0 100644 --- a/package.json +++ b/package.json @@ -201,7 +201,7 @@ "ignore": "^7.0.4", "ipull": "^3.9.5", "is-unicode-supported": "^2.1.0", - "lifecycle-utils": "^4.5.0", + "lifecycle-utils": "^4.5.2", "log-symbols": "^7.0.1", "nanoid": "^5.1.6", "node-addon-api": "^8.9.1", From 2bea8d254fdecfd7dd5d00550126cecc0bc94648 Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Sat, 26 Sep 2026 11:41:26 +0200 Subject: [PATCH 08/10] fix: bugs --- src/evaluator/LlamaChat/LlamaChat.ts | 4 ++-- .../LlamaDecisionContext/LlamaDecisionContext.ts | 6 +++--- .../utils/createQuestionInputs.ts | 2 +- .../utils/evaluateChoiceDecision.ts | 12 ++++++------ 4 files changed, 12 insertions(+), 12 deletions(-) diff --git a/src/evaluator/LlamaChat/LlamaChat.ts b/src/evaluator/LlamaChat/LlamaChat.ts index 8059f03b..7ab3a43b 100644 --- a/src/evaluator/LlamaChat/LlamaChat.ts +++ b/src/evaluator/LlamaChat/LlamaChat.ts @@ -29,7 +29,7 @@ import {defaultMaxPreloadTokens} from "../LlamaChatSession/utils/LlamaChatSessio import {LlamaLogLevel} from "../../bindings/types.js"; import {replaceRegularTextInLlamaText} from "../../chatWrappers/utils/replaceRegularTextInLlamaText.js"; import {DecisionAnswer, DecisionAnswers, DecisionQuestions} from "../LlamaDecisionContext/types.js"; -import {createQuestionInputs} from "../LlamaDecisionContext/utils/createQuestionInputs.js"; +import {createQuestionInputs, getQuestionInputMaxTokenLength} from "../LlamaDecisionContext/utils/createQuestionInputs.js"; import {createEmptyInvalidDecisionAnswer, createDecisionAnswer} from "../LlamaDecisionContext/utils/createDecisionAnswer.js"; import {trimCommonLlamaTextPrefix} from "../../utils/llamaTextUtils.js"; import {TokenMeter} from "../TokenMeter.js"; @@ -1185,7 +1185,7 @@ export class LlamaChat { return await withLock([this._chatLock, "evaluate"], signal, async (): Promise> => { const inputs = createQuestionInputs(questions, this.model.tokenizer); - const maxInputLength = Object.values(inputs).reduce((max, item) => Math.max(max, item.input.length), 0); + const maxInputLength = Object.values(inputs).reduce((max, item) => Math.max(max, getQuestionInputMaxTokenLength(item)), 0); if (maxInputLength > this.sequence.contextSize) throw new Error( "The context size is too small to fit the provided questions and/or criteria. " + diff --git a/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts b/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts index 1fc2a77f..a2234f27 100644 --- a/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts +++ b/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts @@ -7,7 +7,7 @@ import {prepareDecisionContextWindow} from "../LlamaChat/utils/prepareDecisionCo import {ChatWrapper} from "../../ChatWrapper.js"; import {resolveChatWrapper} from "../../chatWrappers/utils/resolveChatWrapper.js"; import {TokenMeter} from "../TokenMeter.js"; -import {createQuestionInputs} from "./utils/createQuestionInputs.js"; +import {createQuestionInputs, getQuestionInputMaxTokenLength} from "./utils/createQuestionInputs.js"; import {createDecisionAnswer} from "./utils/createDecisionAnswer.js"; import {evaluateChoiceDecision} from "./utils/evaluateChoiceDecision.js"; import type {DecisionAnswer, DecisionAnswers, DecisionQuestions} from "./types.js"; @@ -231,7 +231,7 @@ export class LlamaDecisionContext { } const inputs = createQuestionInputs(questions, this.model.tokenizer); - const maxInputLength = Object.values(inputs).reduce((max, item) => Math.max(max, item.input.length), 0); + const maxInputLength = Object.values(inputs).reduce((max, item) => Math.max(max, getQuestionInputMaxTokenLength(item)), 0); if (maxInputLength > this.contextSize) throw new Error( "The context size is too small to fit the provided questions and/or criteria. " + @@ -525,8 +525,8 @@ export class LlamaDecisionContext { using seqLease = await localQueue.acquire(signal); const seq = seqLease.item; - evaluationsLeft--; using drainToParentOnFinishHandle = scopeExit(() => { + evaluationsLeft--; if (evaluationsLeft > 0) return; diff --git a/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts b/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts index 1a250a42..4e449163 100644 --- a/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts +++ b/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts @@ -118,7 +118,7 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number input.push(token); pushAll(input, LlamaText([ ". ", question.criteria[criteriaKey] ?? criteriaKey, - keyIndex === choiceOptions.length - 1 + keyIndex === keys.length - 1 ? "" : "\n" ]).tokenize(tokenizer, "trimLeadingSpace")); diff --git a/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts b/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts index e1071f7d..1cbecb58 100644 --- a/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts +++ b/src/evaluator/LlamaDecisionContext/utils/evaluateChoiceDecision.ts @@ -133,8 +133,8 @@ export async function evaluateChoiceDecision({ using seqLease = await localQueue.acquire(signal); const seq = seqLease.item; - evaluationsLeft--; using drainToParentOnFinishHandle = scopeExit(() => { + evaluationsLeft--; if (evaluationsLeft > 0 || localQueue.parent == null) return; @@ -158,12 +158,12 @@ export async function evaluateChoiceDecision({ }; }); localSeqs.add(seq); + } - await using exitHandle = scopeExit(() => fixPendingSeqs(seq)); - if (needPrefixSeqs.size != 0) { - await fixPendingSeqs(seq); - signal?.throwIfAborted(); - } + await using exitHandle = scopeExit(() => fixPendingSeqs(seq)); + if (needPrefixSeqs.size != 0) { + await fixPendingSeqs(seq); + signal?.throwIfAborted(); } if (evaluateInput.length === 0) From 29d26bc48f80a5b2c98b780592357adfe8143ec7 Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Sat, 26 Sep 2026 12:59:06 +0200 Subject: [PATCH 09/10] test: fix tests --- test/modelDependent/gemma4-e2b/structuredDecisions.test.ts | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts b/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts index c65f33a6..ed2ddd9c 100644 --- a/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts +++ b/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts @@ -363,7 +363,7 @@ describe("gemma4 e2b", () => { }); describe("in a chat", () => { - test("matching", {timeout: 1000 * 60 * 60 * 2}, async () => { + test("matching", {timeout: 1000 * 60 * 60 * 2}, async (test) => { const modelPath = await getModelFile("gemma-4-E2B-it-Q4_K_M.gguf"); const llama = await getTestLlama(); @@ -622,8 +622,8 @@ describe("gemma4 e2b", () => { expect(res1.locks.confidence).to.be.greaterThan(0.8); expect(res1.locks.choice).to.equal("notDoors"); - expect(res2.locks.confidence).to.be.greaterThan(0.8); - expect(res2.locks.choice).to.equal("useful"); + expect(res2.locks.confidence).to.be.lessThan(0.8); + expect(res2.locks.choice).to.not.equal("notDoors"); }); }); }); From ba85df8cb117d03acf675cf2485eaeb07699cacc Mon Sep 17 00:00:00 2001 From: "Gilad S." Date: Sat, 26 Sep 2026 20:07:19 +0200 Subject: [PATCH 10/10] fix: bugs --- docs/guide/structured-decisions.md | 2 +- .../LlamaChat/utils/prepareDecisionContextWindow.ts | 6 ++++-- src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts | 2 +- 3 files changed, 6 insertions(+), 4 deletions(-) diff --git a/docs/guide/structured-decisions.md b/docs/guide/structured-decisions.md index 005a0dd5..b319b508 100644 --- a/docs/guide/structured-decisions.md +++ b/docs/guide/structured-decisions.md @@ -194,7 +194,7 @@ Try refining the criteria to make it more specific, or add an additional option ::: tip NOTE -The specific values that you get for `confidence` and the rest of the probabilities could slightly vary on each evaluation (depending on your machine and setup), +The specific values for `confidence` and the rest of the probabilities could slightly vary on each evaluation (depending on your machine and setup), but the general consensus should remain stable - the `choice` with the highest confidence stays the same, a `noul` stays as decisive as before, but the exact confidence and probability values may fluctuate. diff --git a/src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts b/src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts index 99b22e93..64b210da 100644 --- a/src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts +++ b/src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts @@ -102,7 +102,8 @@ export async function prepareDecisionContextWindow({ history: fullHistory, contextShiftSize: Math.max( minFreeContextTokens, - Math.min(contextShiftSize, context.contextSize - fitRegularContextWindowUnderTokenCount) + contextShiftSize, + context.contextSize - fitRegularContextWindowUnderTokenCount ), contextShiftStrategy: resolvedContextShift.strategy, contextShiftLastEvaluationMetadata: resolvedContextShift.lastEvaluationMetadata, @@ -141,7 +142,8 @@ export async function prepareDecisionContextWindow({ history: fullHistory, contextShiftSize: Math.max( minFreeContextTokens, - Math.min(contextShiftSize, context.contextSize - fitRegularContextWindowUnderTokenCount) + contextShiftSize, + context.contextSize - fitRegularContextWindowUnderTokenCount ), contextShiftStrategy: resolvedContextShift.strategy, contextShiftLastEvaluationMetadata: resolvedContextShift.lastEvaluationMetadata, diff --git a/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts b/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts index a2234f27..25e217a9 100644 --- a/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts +++ b/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts @@ -265,7 +265,7 @@ export class LlamaDecisionContext { : onOverflow === "truncateDocument" ? { lastEvaluationMetadata: null, - size: (sequence) => sequence.contextSize - 1, + size: 1, strategy({maxTokensCount, tokenizer, chatWrapper}) { const fullDocumentTokenLength = tokenizer(document, false, "trimLeadingSpace").length; const testTokenLength = chatWrapper.generateContextState({