Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 11 additions & 2 deletions docs/guide/structured-decisions.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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)),
Expand Down
19 changes: 13 additions & 6 deletions llama/addon/AddonContext.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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);

Expand Down Expand Up @@ -315,17 +318,21 @@ class AddonContextSampleTokenWorker : public Napi::AsyncWorker {
}

void SampleToken() {
std::unique_lock<std::mutex> samplingLock(ctx->samplingMutex);
if (llama_get_logits(ctx->ctx) == nullptr) {
sampler->rebuildChainIfNeeded();

std::unique_lock<std::mutex> 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;
Expand Down
7 changes: 6 additions & 1 deletion llama/addon/AddonContext.h
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,12 @@ class AddonContext : public Napi::ObjectWrap<AddonContext> {
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;
Expand Down
9 changes: 7 additions & 2 deletions llama/addon/AddonSampler.cpp
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
#include <cmath>
#include <mutex>
#include "common/common.h"
#include "globals/addonLog.h"
#include "ggml.h"
Expand Down Expand Up @@ -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<std::mutex> & 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;
Expand Down Expand Up @@ -206,7 +208,10 @@ void AddonSampler::sample(struct llama_context* llamaContext, int32_t batchLogit
return;
}

setTokenCandidates(llamaContext, batchLogitIndex, curP);
{
std::unique_lock<std::mutex> samplingLock(samplingMutex);
setTokenCandidates(llamaContext, batchLogitIndex, curP);
}

llama_sampler_apply(grammarEvaluationState->sampler, &curP);
llama_sampler_apply(chain, &curP);
Expand Down
3 changes: 2 additions & 1 deletion llama/addon/AddonSampler.h
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
#pragma once
#include <mutex>
#include "llama.h"
#include "napi.h"
#include "RingBuffer.h"
Expand Down Expand Up @@ -64,7 +65,7 @@ class AddonSampler : public Napi::ObjectWrap<AddonSampler> {
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<std::mutex> & 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);
Expand Down
8 changes: 4 additions & 4 deletions package-lock.json

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 1 addition & 1 deletion package.json
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
73 changes: 49 additions & 24 deletions src/evaluator/LlamaChat/LlamaChat.ts
Original file line number Diff line number Diff line change
@@ -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 {
Expand Down Expand Up @@ -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";
Expand Down Expand Up @@ -473,6 +474,17 @@ export type LLamaChatLoadAndCompleteUserMessageOptions<Functions extends ChatMod
};

export type LlamaChatGenerateDecisionsOptions = {
/**
* 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 context window chat history returned from this evaluation.
*
* 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?: AbortSignal,

/**
Expand Down Expand Up @@ -1160,6 +1172,7 @@ export class LlamaChat {
options: LlamaChatGenerateDecisionsOptions = {}
): Promise<LlamaChatGenerateDecisionsResponse<Questions>> {
const {
document,
evaluationPriority = defaultEvaluationPriority,
contextShift = defaultContextShiftOptions,
functions,
Expand All @@ -1171,8 +1184,8 @@ export class LlamaChat {
} = options;

return await withLock([this._chatLock, "evaluate"], signal, async (): Promise<LlamaChatGenerateDecisionsResponse<Questions>> => {
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. " +
Expand Down Expand Up @@ -1202,7 +1215,8 @@ export class LlamaChat {
chatWrapper: this._chatWrapper,
sequence: this.sequence,
functions,
documentFunctionParams
documentFunctionParams,
injectedDocument: document
});

const answers: {[key: string]: DecisionAnswer<any>} = {} as DecisionAnswers<Questions>;
Expand All @@ -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;
}
Expand All @@ -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<LlamaContextSequence>([]);
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;
}
Expand Down
29 changes: 21 additions & 8 deletions src/evaluator/LlamaChat/utils/prepareDecisionContextWindow.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<LLamaChatContextShiftOptions>, 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,
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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]
};
}
Expand Down Expand Up @@ -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;
}
Loading
Loading