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
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@

## Unreleased

- Added image inputs for vision-capable llama.cpp, Ollama, and OpenRouter models.
llama.cpp image scoring uses `labels` mode.
- Added labels-only Ollama scoring with up to 20 choices through `choosekit/ollama`.
- Added an exact `/v1/models` check to the SemIf llama.cpp benchmark, with
`--skip-model-check` for unverified runs.
Expand Down
25 changes: 24 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
# choosekit

`choosekit` scores a finite set of choices with a language model and returns a typed decision with a probability distribution. It supports local models through llama.cpp and Ollama, plus an optional OpenRouter backend.
`choosekit` scores a finite set of choices and returns a typed decision with a probability distribution. It accepts `text and images` through llama.cpp, Ollama, and OpenRouter.

![SuperGPQA direct-choice benchmark](benchmarks/supergpqa-benchmark.svg)

Expand Down Expand Up @@ -108,6 +108,29 @@ Returned probabilities are normalized across the supplied choices and are not ca

OpenRouter may route the same model through different providers. Set `provider: "provider-slug"` to use only that provider and disable fallback.

## Image inputs

llama.cpp, Ollama, and OpenRouter can score choices from images when the selected model supports vision. Pass raw base64 data with its media type:

```ts
import { readFile } from "node:fs/promises";

const decision = await choose({
context: "Inspect the attached screenshot.",
question: "Which state is the interface in?",
choices: {
ready: "The interface is ready for input.",
loading: "The interface is still loading.",
},
images: [{
mediaType: "image/png",
base64: (await readFile("screenshot.png")).toString("base64"),
}],
});
```

Supported media types are PNG, JPEG, and WebP. llama.cpp image inputs currently support `labels` mode only.

## Scoring modes

| Mode | Candidate representation | Use when |
Expand Down
2 changes: 1 addition & 1 deletion src/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@ import type { Chooser, ChooserOptions, Scorer } from "./types.js";

export type {
Choices, ChoiceKey, ChoiceRequest, Chooser, ChooserOptions, Decision,
ScoreRequest, Scorer, Scores, Usage, PromptInput,
ImageInput, ImageMediaType, ScoreRequest, Scorer, Scores, Usage, PromptInput,
} from "./types.js";
export { ScoringError } from "./validation.js";

Expand Down
24 changes: 22 additions & 2 deletions src/internal-chooser.ts
Original file line number Diff line number Diff line change
@@ -1,11 +1,28 @@
import type {
Choices, ChoiceKey, ChoiceRequest, Chooser, ChooserOptions, Decision, Scorer, Usage,
Choices, ChoiceKey, ChoiceRequest, Chooser, ChooserOptions, Decision, ImageInput, Scorer, Usage,
} from "./types.js";
import { isCount, isLogprob, isRecord, requireText, ScoringError } from "./validation.js";

export type CandidateFormat = "keys" | "labels";

const LABELS = "ABCDEFGHIJKLMNOPQRSTUVWXYZ";
const IMAGE_MEDIA_TYPES = new Set(["image/png", "image/jpeg", "image/webp"]);

function snapshotImages(value: unknown): readonly ImageInput[] | undefined {
if (value === undefined) return undefined;
if (!Array.isArray(value)) throw new TypeError("images must be an array.");
return Object.freeze(value.map((image, index) => {
if (!isRecord(image)) throw new TypeError(`images[${index}] must be an object.`);
if (!IMAGE_MEDIA_TYPES.has(image.mediaType as string)) {
throw new TypeError(`images[${index}].mediaType is not supported.`);
}
requireText(image.base64, `images[${index}].base64`);
return Object.freeze({
mediaType: image.mediaType as ImageInput["mediaType"],
base64: image.base64,
});
}));
}

function snapshot(choices: unknown): [string, string][] {
if (!isRecord(choices) || Object.getOwnPropertySymbols(choices).length !== 0) {
Expand Down Expand Up @@ -109,6 +126,7 @@ export function createFormattedChooser(score: Scorer, options: ChooserOptions,
if (typeof context !== "string") throw new TypeError("context must be a string.");
requireText(question, "question");
const entries = snapshot(choices);
const images = snapshotImages(request.images);
const keys = entries.map(([key]) => key) as ChoiceKey<C>[];
const prepared = instruction(question, entries, candidateFormat);
const prompt = formatPrompt
Expand All @@ -122,7 +140,9 @@ export function createFormattedChooser(score: Scorer, options: ChooserOptions,
}
signal?.throwIfAborted();
const scored = await score(Object.freeze({
prompt, candidates: prepared.candidates, ...(signal ? { signal } : {}),
prompt, candidates: prepared.candidates,
...(images === undefined ? {} : { images }),
...(signal ? { signal } : {}),
}));
signal?.throwIfAborted();
if (!isRecord(scored) || !Array.isArray(scored.logprobs)
Expand Down
115 changes: 104 additions & 11 deletions src/llama-cpp.ts
Original file line number Diff line number Diff line change
Expand Up @@ -12,13 +12,19 @@ export interface LlamaCppOptions extends ChooserOptions {
readonly model?: string;
readonly headers?: Readonly<Record<string, string>>;
readonly fetch?: typeof globalThis.fetch;
/** Must match the agent's tokenizer setting. Defaults to false for serialized context. */
/** Controls add_special for llama.cpp tokenization. Defaults to false; image inputs require false. */
readonly addSpecialTokens?: boolean;
}

interface Endpoints {
readonly tokenize: string;
readonly detokenize: string;
readonly completion: string;
readonly props: string;
}

interface ImageSupport {
readonly marker: string;
}

interface Probe {
Expand All @@ -36,7 +42,7 @@ interface Branch {
readonly score: number;
}

function endpoints(baseURL: string): Endpoints {
function endpoints(baseURL: string, model?: string): Endpoints {
requireText(baseURL, "baseURL");
const url = new URL(baseURL);
if ((url.protocol !== "http:" && url.protocol !== "https:") || url.username || url.password
Expand All @@ -46,8 +52,13 @@ function endpoints(baseURL: string): Endpoints {
const path = url.pathname.replace(/\/v1\/?$/, "").replace(/\/$/, "");
url.pathname = `${path}/tokenize`;
const tokenize = url.href;
url.pathname = `${path}/detokenize`;
const detokenize = url.href;
url.pathname = `${path}/completion`;
return { tokenize, completion: url.href };
const completion = url.href;
url.pathname = `${path}/props`;
if (model !== undefined) url.searchParams.set("model", model);
return { tokenize, detokenize, completion, props: url.href };
}

function requestHeaders(value: unknown): Readonly<Record<string, string>> {
Expand Down Expand Up @@ -91,12 +102,54 @@ async function post(fetchImpl: typeof globalThis.fetch, url: string,
}
}

async function get(fetchImpl: typeof globalThis.fetch, url: string,
headers: Readonly<Record<string, string>>, signal?: AbortSignal): Promise<unknown> {
signal?.throwIfAborted();
let response: Response;
try {
response = await fetchImpl(url, { method: "GET", headers, ...(signal ? { signal } : {}) });
} catch (error) {
signal?.throwIfAborted();
throw error;
}
signal?.throwIfAborted();
if (!response.ok) {
throw new ScoringError(`llama.cpp returned HTTP ${response.status} for ${new URL(url).pathname}.`);
}
try {
const value: unknown = await response.json();
signal?.throwIfAborted();
return value;
} catch {
signal?.throwIfAborted();
throw new ScoringError(`llama.cpp returned invalid JSON for ${new URL(url).pathname}.`);
}
}

function parseTokenization(value: unknown): number[] {
if (!isRecord(value)) throw new ScoringError("llama.cpp returned an invalid tokenization.");
return tokenIds(value.tokens);
}

function parseProbe(value: unknown, prompt: readonly number[], targetTokenId: number): Probe {
function parseDetokenization(value: unknown): string {
if (!isRecord(value) || typeof value.content !== "string") {
throw new ScoringError("llama.cpp returned an invalid detokenization.");
}
return value.content;
}

function parseImageSupport(value: unknown): ImageSupport {
if (!isRecord(value) || !isRecord(value.modalities) || value.modalities.vision !== true) {
throw new ScoringError("llama.cpp does not advertise vision support for this model.");
}
if (typeof value.media_marker !== "string" || value.media_marker.length === 0) {
throw new ScoringError("llama.cpp did not return a multimodal media marker.");
}
return Object.freeze({ marker: value.media_marker });
}

function parseProbe(value: unknown, expectedPromptTokens: number | null,
targetTokenId: number): Probe {
if (!isRecord(value)) throw new ScoringError("llama.cpp returned an invalid completion.");
if (value.truncated === true) {
throw new ScoringError("llama.cpp truncated the token prefix while scoring a candidate.");
Expand Down Expand Up @@ -144,7 +197,10 @@ function parseProbe(value: unknown, prompt: readonly number[], targetTokenId: nu
}

const promptTokens: unknown = value.tokens_evaluated;
if (!isCount(promptTokens) || promptTokens !== prompt.length) {
if (!isCount(promptTokens)) {
throw new ScoringError("llama.cpp returned invalid prompt-token usage.");
}
if (expectedPromptTokens !== null && promptTokens !== expectedPromptTokens) {
throw new ScoringError("llama.cpp did not evaluate the supplied numeric token prefix as sent.");
}
let completionTokens = output.length;
Expand Down Expand Up @@ -178,11 +234,20 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser {
}
const fetchImpl = options.fetch ?? globalThis.fetch;
if (typeof fetchImpl !== "function") throw new TypeError("A fetch implementation is required.");
const urls = endpoints(baseURL);
const urls = endpoints(baseURL, model);
const headers = requestHeaders(options.headers);

const score: Scorer = async ({ prompt, candidates, signal }) => {
const score: Scorer = async ({ prompt, candidates, images, signal }) => {
const hasImages = images !== undefined && images.length > 0;
if (hasImages && mode !== "labels") {
throw new TypeError("llama.cpp image inputs require labels mode.");
}
if (hasImages && addSpecialTokens) {
throw new TypeError("llama.cpp image inputs require addSpecialTokens to be false.");
}
signal?.throwIfAborted();
const imageSupport = hasImages
? parseImageSupport(await get(fetchImpl, urls.props, headers, signal))
: undefined;
const encoded: number[][] = [];
for (const content of [prompt, ...candidates.map((candidate) => prompt + candidate)]) {
const response = await post(fetchImpl, urls.tokenize, headers, {
Expand All @@ -199,7 +264,7 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser {
let promptTokens = 0;
let cachedTokens: number | null = 0;
let completionTokens = 0;
let requests = encoded.length;
let requests = encoded.length + (hasImages ? 1 : 0);

let root = treeRoot;
const rootSuffix: number[] = [];
Expand All @@ -209,6 +274,33 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser {
rootSuffix.push(first[0]);
root = first[1];
}
const materializeImagePrompt = async (numericPrefix: readonly number[]) => {
const response = await post(fetchImpl, urls.detokenize, headers, {
tokens: numericPrefix, ...(model === undefined ? {} : { model }),
}, signal);
requests++;
const detokenized = parseDetokenization(response);
if (detokenized.includes(imageSupport!.marker)) {
throw new ScoringError("The formatted prompt contains llama.cpp's multimodal media marker.");
}

const roundTripResponse = await post(fetchImpl, urls.tokenize, headers, {
content: detokenized, add_special: false,
...(model === undefined ? {} : { model }),
}, signal);
requests++;
const roundTrip = parseTokenization(roundTripResponse);
if (roundTrip.length !== numericPrefix.length
|| roundTrip.some((tokenId, index) => tokenId !== numericPrefix[index])) {
throw new ScoringError("llama.cpp could not preserve the image prompt token prefix.");
}

const value = Object.freeze({
prompt_string: `${images!.map(() => imageSupport!.marker).join("\n")}\n${detokenized}`,
multimodal_data: Object.freeze(images!.map(({ base64 }) => base64)),
});
return value;
};

const work: Branch[] = [{
node: root, indices: candidates.map((_, index) => index), suffix: rootSuffix, score: 0,
Expand All @@ -229,12 +321,13 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser {
const children = [...branch.node.children];
const siblingLogprobs = new Map<number, number>();
const numericPrefix = [...base, ...branch.suffix];
const promptValue = hasImages ? await materializeImagePrompt(numericPrefix) : numericPrefix;
for (const [targetTokenId] of children) {
if (siblingLogprobs.has(targetTokenId)) continue;
const collectSiblings = siblingLogprobs.size === 0 && children.length > 1;
signal?.throwIfAborted();
const response = await post(fetchImpl, urls.completion, headers, {
prompt: numericPrefix,
prompt: promptValue,
...(model === undefined ? {} : { model }),
n_predict: 1,
n_probs: collectSiblings ? 64 : 1,
Expand All @@ -249,7 +342,7 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser {
cache_prompt: true,
}, signal);
signal?.throwIfAborted();
const probe = parseProbe(response, numericPrefix, targetTokenId);
const probe = parseProbe(response, hasImages ? null : numericPrefix.length, targetTokenId);
promptTokens += probe.promptTokens;
cachedTokens = cachedTokens === null || probe.cachedTokens === null
? null : cachedTokens + probe.cachedTokens;
Expand Down
10 changes: 8 additions & 2 deletions src/ollama.ts
Original file line number Diff line number Diff line change
Expand Up @@ -130,13 +130,19 @@ export function fromOllama(options: OllamaOptions): Chooser {
if (typeof fetchImpl !== "function") throw new TypeError("A fetch implementation is required.");
const url = endpoint(options.baseURL);

const score: Scorer = async ({ prompt, candidates, signal }) => {
const score: Scorer = async ({ prompt, candidates, images, signal }) => {
if (candidates.length > MAX_CANDIDATES) {
throw new TypeError(`Ollama supports at most ${MAX_CANDIDATES} choices.`);
}
const response = await post(fetchImpl, url, {
model,
messages: [{ role: "user", content: prompt }],
messages: [{
role: "user",
content: prompt,
...(images === undefined || images.length === 0
? {}
: { images: images.map(({ base64 }) => base64) }),
}],
stream: false,
think: false,
logprobs: true,
Expand Down
13 changes: 11 additions & 2 deletions src/openrouter.ts
Original file line number Diff line number Diff line change
Expand Up @@ -154,13 +154,22 @@ export function fromOpenRouter(options: OpenRouterOptions): Chooser {
const fetchImpl = options.fetch ?? globalThis.fetch;
if (typeof fetchImpl !== "function") throw new TypeError("A fetch implementation is required.");

const score: Scorer = async ({ prompt, candidates, signal }) => {
const score: Scorer = async ({ prompt, candidates, images, signal }) => {
if (candidates.length > MAX_CANDIDATES) {
throw new TypeError(`OpenRouter supports at most ${MAX_CANDIDATES} choices.`);
}
const response = await post(fetchImpl, apiKey, {
model,
messages: [{ role: "user", content: prompt }],
messages: [{
role: "user",
content: images === undefined || images.length === 0 ? prompt : [
{ type: "text", text: prompt },
...images.map(({ mediaType, base64 }) => ({
type: "image_url",
image_url: { url: `data:${mediaType};base64,${base64}` },
})),
],
}],
max_tokens: 1,
stream: false,
temperature: 1,
Expand Down
10 changes: 10 additions & 0 deletions src/types.ts
Original file line number Diff line number Diff line change
@@ -1,11 +1,20 @@
export type Choices = Readonly<Record<string, string>>;
export type ChoiceKey<C extends Choices> = `${Extract<keyof C, string | number>}`;

export type ImageMediaType = "image/png" | "image/jpeg" | "image/webp";

export interface ImageInput {
readonly mediaType: ImageMediaType;
/** Raw base64 data without a data URL prefix. */
readonly base64: string;
}

export interface ChoiceRequest<C extends Choices> {
/** Existing context, copied unchanged to the beginning of the scoring prompt. */
readonly context: string;
readonly question: string;
readonly choices: C;
readonly images?: readonly ImageInput[];
readonly signal?: AbortSignal;
}

Expand Down Expand Up @@ -35,6 +44,7 @@ export interface Decision<K extends string> {
export interface ScoreRequest {
readonly prompt: string;
readonly candidates: readonly string[];
readonly images?: readonly ImageInput[];
readonly signal?: AbortSignal;
}

Expand Down
Loading
Loading