diff --git a/docs/guide/structured-decisions.md b/docs/guide/structured-decisions.md index 1458eb91..b319b508 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", @@ -191,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 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)), 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); 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", diff --git a/src/evaluator/LlamaChat/LlamaChat.ts b/src/evaluator/LlamaChat/LlamaChat.ts index d792c72c..7ab3a43b 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 { @@ -29,10 +29,11 @@ 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 {createEmptyInvalidDecisionAnswer, createDecisionAnswer, decisionAnswerMinimumTopLogits} from "../LlamaDecisionContext/utils/createDecisionAnswer.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"; +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"; @@ -473,6 +474,17 @@ export type LLamaChatLoadAndCompleteUserMessageOptions> { const { + document, evaluationPriority = defaultEvaluationPriority, contextShift = defaultContextShiftOptions, functions, @@ -1171,8 +1184,8 @@ export class LlamaChat { } = options; return await withLock([this._chatLock, "evaluate"], signal, async (): Promise> => { - const inputs = createQuestionInputs(questions, this.model); - const maxInputLength = Object.values(inputs).reduce((max, item) => Math.max(max, item.input.length), 0); + const inputs = createQuestionInputs(questions, this.model.tokenizer); + 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. " + @@ -1202,7 +1215,8 @@ export class LlamaChat { chatWrapper: this._chatWrapper, sequence: this.sequence, functions, - documentFunctionParams + documentFunctionParams, + injectedDocument: document }); const answers: {[key: string]: DecisionAnswer} = {} as DecisionAnswers; @@ -1215,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; } @@ -1242,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/LlamaChat/utils/prepareDecisionContextWindow.ts b/src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts index 522b62ee..64b210da 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, @@ -100,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, @@ -139,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, @@ -174,7 +178,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 +222,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 +258,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/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 34474923..25e217a9 100644 --- a/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts +++ b/src/evaluator/LlamaDecisionContext/LlamaDecisionContext.ts @@ -7,8 +7,9 @@ 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 {createDecisionAnswer, decisionAnswerMinimumTopLogits} from "./utils/createDecisionAnswer.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"; import type {LlamaModel} from "../LlamaModel/LlamaModel.js"; import type {ChatHistoryItem, Token, Tokenizer} from "../../types.js"; @@ -229,8 +230,8 @@ export class LlamaDecisionContext { disposeAggregator.add(() => signal.removeEventListener("abort", disposeAggregator.dispose)); } - const inputs = createQuestionInputs(questions, this.model); - const maxInputLength = Object.values(inputs).reduce((max, item) => Math.max(max, item.input.length), 0); + const inputs = createQuestionInputs(questions, this.model.tokenizer); + 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. " + @@ -264,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({ @@ -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 } }; } @@ -498,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; @@ -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 1e9b69a2..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; @@ -102,12 +56,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); @@ -153,31 +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; - - const destinationLogit = textToLogit.get(text); - if (destinationLogit != null && destinationLogit > logit) - res.set(token, destinationLogit); - } - - return res; -} diff --git a/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts b/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts index 27c06c96..4e449163 100644 --- a/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts +++ b/src/evaluator/LlamaDecisionContext/utils/createQuestionInputs.ts @@ -1,21 +1,20 @@ 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) { +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, 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): QuestionInput { if ( (LlamaText.isLlamaText(question.instruction) && question.instruction.values.length === 0) || (typeof question.instruction === "string" && question.instruction.length === 0) || @@ -24,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, model); + 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() === "") @@ -56,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, model); + + 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'); @@ -76,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 === keys.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, model); + 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'); @@ -121,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" @@ -132,7 +181,7 @@ function createQuestionInput(keyName: string, question: DecisionQuestions[number type: "score" as const, tokens: scoreOptions, input - }; + } satisfies ScoreQuestionInput; } else void (question satisfies never); @@ -145,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..1cbecb58 --- /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; + + using drainToParentOnFinishHandle = scopeExit(() => { + evaluationsLeft--; + 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.tanh(((maxScore ?? 0) - (secondMaxScore ?? 0)) / 2), + 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 53a7d1d1..714c4d07 100644 --- a/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts +++ b/src/evaluator/LlamaDecisionContext/utils/generateCriteriaChoiceOptionTokens.ts @@ -1,75 +1,20 @@ -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) { - if (count <= 0) +export function generateCriteriaChoiceOptionTokens(min: number, max: number, tokenizer: Tokenizer, additionalChars: string = ""): Token[] { + if (max <= 0) return []; const res = new Set(); - const ranges = - "AZ" + - "09" + - "\u03b1\u03c1\u03c3\u03c9" + // greek symbols - "\u0531\u0556" + // hy - "\u05d0\u05d9\u05db\u05dc\u05de\u05de\u05e0\u05e2\u05e4\u05e4\u05e6\u05ea" + // he - "\u10d0\u10f0" + // 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, model); + const token = findSingleToken(char, tokenizer); if (token != null) res.add(token); - return res.size >= count; + return res.size >= max; } function pushCode(code: number) { @@ -90,68 +35,29 @@ export function generateCriteriaChoiceOptionTokens(count: number, model: LlamaMo 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; + if (additionalChars !== "") { + for (const char of additionalChars) { + if (pushCharacter(char)) + break; } - - return false; } - 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("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" ); 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; } diff --git a/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts b/test/modelDependent/gemma4-e2b/structuredDecisions.test.ts index c71e9b78..ed2ddd9c 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.0000762, + "database": 0.00000208, }, "type": "choice", }, @@ -108,10 +108,10 @@ describe("gemma4 e2b", () => { "value": 1, }, "level": { - "confidence": 0.999, + "confidence": 1, "probabilities": [ - 0.0000978, - 0.0002, + 0.0000264, + 0.0000152, 1, ], "score": 2, @@ -122,8 +122,8 @@ describe("gemma4 e2b", () => { "confidence": 1, "probabilities": { "engineering": 1, - "hr": 0.000214, - "sales": 0.000036, + "hr": 0.000242, + "sales": 0.0000568, }, "type": "choice", }, @@ -237,14 +237,139 @@ 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); + }); }); }); 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(); + if (llama.gpu === false) + test.skip("Logits are a bit different on different backends to cause test flakiness"); + const model = await llama.loadModel({ modelPath }); @@ -255,13 +380,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 +430,7 @@ describe("gemma4 e2b", () => { { "animal": { "type": "noul", - "value": 1, + "value": 0.999, }, "animalOrigins": { "type": "noul", @@ -322,7 +446,7 @@ describe("gemma4 e2b", () => { }, "relatedAnimals": { "type": "noul", - "value": 1, + "value": 0.999, }, } `); @@ -348,13 +472,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 +536,32 @@ describe("gemma4 e2b", () => { { "animal": { "type": "noul", - "value": 0.00238, + "value": 0.000187, }, "cookingRecipe": { "type": "noul", - "value": 0.0000225, + "value": 0.000112, }, "fictionalStory": { "type": "noul", - "value": 0.0000138, + "value": 0.0000023, }, "mineralOrigin": { "type": "noul", - "value": 0.0000348, + "value": 0.0000034, }, "spaceTravel": { "type": "noul", - "value": 0.0000123, + "value": 0.0000368, }, "subject": { "choice": "materials", - "confidence": 1, + "confidence": 0.998, "probabilities": { - "brushing": 0.00000313, - "food": 0.00000662, - "materials": 1, - "other": 0.0000181, + "brushing": 0.0000175, + "food": 0.000779, + "materials": 0.999, + "other": 0.00018, }, "type": "choice", }, @@ -449,6 +572,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.lessThan(0.8); + expect(res2.locks.choice).to.not.equal("notDoors"); + }); }); }); });