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
9 changes: 6 additions & 3 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -157,9 +157,12 @@ from uniform, 0, to certain, 1):
jevos answers yes/no questions, so a choice is asked as one yes/no question per option, each listing all
the options ("... Among the candidates, is it this one? Candidate: shipping"), the way the model saw
choices in training. The option's probability is its P(yes) divided by the sum over the options.
What the option questions have in common, the state, the instructions and the list of options, is read
once; each option then costs only its last line. On the laptop where one yes/no question takes about
20 ms, this three-option choice takes about 100 ms.
All of them go to the model in one call, read as a prefix tree: what the option questions have in
common, the state, the instructions and the list of options, is read once, and each option then costs
only its last line, also when the request asks other questions beside the choice. Each option still
sees only its own question: the answers are the same as asking each option alone. On the laptop where
one yes/no question takes about 20 ms, this three-option choice takes about 100 ms, a ten-option one
about 170 ms.

On 2,676 held-out choice questions, the most probable option is the right one 78.8% of the time (74.2%
on rental questions with labels from code, 81.2% agreement with Jev on workflow questions). The
Expand Down
37 changes: 21 additions & 16 deletions src/calls.hpp
Original file line number Diff line number Diff line change
@@ -1,43 +1,48 @@
// How a request is split into model calls (plan_calls). Nothing here knows about the model.
#pragma once
#include <algorithm>
#include <cstdint>
#include <vector>

#include "tree.hpp"

// One model call of a request: the state tokens [state_from, state_to) it reads and the questions (jobs)
// it answers.
struct Call {
size_t state_from = 0, state_to = 0;
std::vector<size_t> questions;
};

// The state's tokens from `from` (the cells before come from a snapshot) up to P, then the questions
// (suffix[i] tokens each after the state), shortest first. While the rest of the state and the shortest
// question do not fit `call_budget`, a piece of at most `piece_budget` state tokens goes alone; the call
// with the state's last piece takes the questions that fit `call_budget`, the next calls the others, as
// many as fit and at most `max_count` each. A question longer than the budget has a call of its own.
inline std::vector<Call> plan_calls(size_t from, size_t P, const std::vector<size_t>& suffix, size_t piece_budget, size_t call_budget,
// The prompts all start with the same P tokens (the state, and what else they share), read from `from`
// on (the cells before come from a snapshot); after them, each call reads its questions as a prefix tree
// (tree.hpp), so a question costs only what it does not share with the others in its call. While the
// rest of the state and the cheapest question do not fit `call_budget`, a piece of at most
// `piece_budget` state tokens goes alone; the call with the state's last piece takes questions in order
// while they fit `call_budget`, the next calls the others, at most `max_count` each. A question longer
// than the budget has a call of its own.
inline std::vector<Call> plan_calls(size_t from, size_t P, const std::vector<const Tokens*>& prompts, size_t piece_budget, size_t call_budget,
size_t max_count) {
std::vector<size_t> order(suffix.size());
for (size_t i = 0; i < order.size(); ++i) order[i] = i;
std::stable_sort(order.begin(), order.end(), [&](size_t a, size_t b) { return suffix[a] < suffix[b]; });
std::vector<Call> calls;
size_t pos = from;
const size_t shortest = order.empty() ? 0 : suffix[order[0]];
while (pos < P && P - pos + shortest > call_budget) {
size_t cheapest = prompts.empty() ? 0 : SIZE_MAX;
for (auto* p : prompts) cheapest = std::min(cheapest, p->size() - P);
while (pos < P && P - pos + cheapest > call_budget) {
size_t n = std::min(piece_budget, P - pos);
calls.push_back({pos, pos + n, {}});
pos += n;
}
for (size_t k = 0; k < order.size();) {
for (size_t k = 0; k < prompts.size();) {
Call call{pos, P, {}};
size_t used = P - pos;
while (k < order.size() && (call.questions.empty() || (call.questions.size() < max_count && used + suffix[order[k]] <= call_budget))) {
used += suffix[order[k]];
call.questions.push_back(order[k++]);
while (k < prompts.size()) {
size_t cost = added_tokens(prompts, call.questions, k, P);
if (!call.questions.empty() && (call.questions.size() >= max_count || used + cost > call_budget)) break;
used += cost;
call.questions.push_back(k++);
}
calls.push_back(std::move(call));
pos = P;
}
if (order.empty() && pos < P) calls.push_back({pos, P, {}});
if (prompts.empty() && pos < P) calls.push_back({pos, P, {}});
return calls;
}
40 changes: 25 additions & 15 deletions src/model.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -30,8 +30,10 @@ struct Model::Impl {
};
using Index = SnapshotIndex<Cells>; // a snapshot's payload: the state's cells
// One block of a call: the first `past_n` cells of `past`, `shared` state tokens read now (positions
// from `shared_pos`, causal), then questions (their tokens from `skip` on) each seeing the block's
// cells, its shared tokens and itself. No questions: the logits are read at the last shared token.
// from `shared_pos`, causal), then questions (their tokens from `skip` on) as a prefix tree
// (tree.hpp): each question sees the block's cells, its shared tokens and its own tokens, the ones it
// has in common with other questions read once. No questions: the logits are read at the last shared
// token.
struct Block {
const Cells* past = nullptr;
size_t past_n = 0;
Expand Down Expand Up @@ -169,10 +171,10 @@ struct Model::Impl {
const Cells* start = sn ? &sn->payload : nullptr;
const Cells* cur = start; // holds at least the cells before the next call's state piece
Cells state;
std::vector<size_t> suffix;
for (auto& j : q.jobs) suffix.push_back(j.tokens.size() - P);
std::vector<const Tokens*> prompts;
for (auto& j : q.jobs) prompts.push_back(&j.tokens);
std::vector<double> out(q.jobs.size());
for (const Call& call : plan_calls(c, P, suffix, PIECE_TOKENS, CALL_TOKENS, q.jobs.size())) {
for (const Call& call : plan_calls(c, P, prompts, PIECE_TOKENS, CALL_TOKENS, q.jobs.size())) {
std::vector<Block> one{state_block(cur, call.state_from, prefix, call.state_from, call.state_to, q.jobs, call.questions)};
auto ps = run(one);
for (size_t s = 0; s < call.questions.size(); ++s) out[call.questions[s]] = ps[s];
Expand Down Expand Up @@ -227,11 +229,12 @@ struct Model::Impl {
// One call over `blocks`: P(yes) per question in block order (one per block without questions).
std::vector<double> run(std::vector<Block>& blocks) {
size_t C = 0, T = 0, K = 0;
std::vector<PromptTree> trees;
for (auto& b : blocks) {
C += b.past_n;
b.row = T;
T += b.shared.size();
for (auto* p : b.pieces) T += p->size() - b.skip;
trees.push_back(b.pieces.empty() ? PromptTree{} : prompt_tree(b.pieces, b.skip));
T += b.shared.size() + trees.back().tokens;
K += std::max<size_t>(1, b.pieces.size());
}
size_t W = C + T;
Expand All @@ -243,7 +246,9 @@ struct Model::Impl {
float* pb = bias.data<float>();
std::fill(pb, pb + T * W, -std::numeric_limits<float>::infinity());
size_t c0 = 0, k = 0;
for (auto& b : blocks) {
for (size_t bi = 0; bi < blocks.size(); ++bi) {
Block& b = blocks[bi];
const PromptTree& tree = trees[bi];
size_t S = b.shared.size(), at = b.row;
auto open = [&](size_t row, size_t from, size_t to) { std::fill(pb + row * W + from, pb + row * W + to, 0.0f); };
for (size_t j = 0; j < S; ++j) {
Expand All @@ -252,19 +257,24 @@ struct Model::Impl {
open(at + j, c0, c0 + b.past_n); // the block's cached cells
open(at + j, C + at, C + at + j + 1); // its shared tokens, causally
}
std::vector<size_t> node_row(tree.nodes.size());
size_t r = at + S;
for (auto* p : b.pieces) {
size_t n = p->size() - b.skip;
for (size_t j = 0; j < n; ++j) {
pi[r + j] = (*p)[b.skip + j];
pp[r + j] = static_cast<int64_t>(b.skip + j);
for (size_t n = 0; n < tree.nodes.size(); ++n) {
const TreeNode& node = tree.nodes[n];
const Tokens& p = *b.pieces[node.prompt];
node_row[n] = r;
for (size_t j = 0; j < node.to - node.from; ++j) {
pi[r + j] = p[node.from + j];
pp[r + j] = static_cast<int64_t>(node.from + j);
open(r + j, c0, c0 + b.past_n);
open(r + j, C + at, C + at + S); // the state read in this call
for (size_t a = node.parent; a != TreeNode::NO_PARENT; a = tree.nodes[a].parent)
open(r + j, C + node_row[a], C + node_row[a] + tree.nodes[a].to - tree.nodes[a].from); // what it shares
open(r + j, C + r, C + r + j + 1); // itself, causally
}
r += n;
px[k++] = static_cast<int64_t>(r - 1);
r += node.to - node.from;
}
for (size_t q = 0; q < b.pieces.size(); ++q) px[k++] = static_cast<int64_t>(node_row[tree.leaf[q]]); // a leaf is one token
if (b.pieces.empty()) px[k++] = static_cast<int64_t>(at + S - 1);
c0 += b.past_n;
}
Expand Down
2 changes: 1 addition & 1 deletion src/model.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@ class Model {
Model& operator=(const Model&) = delete;

// P(yes) per job of each request. Several requests are read in one call (the scheduler hands over
// only requests that fit CALL_TOKENS together); a request's prefix is its state (empty: none known).
// only requests that fit CALL_TOKENS together); a request's prefix is what all its prompts start with (empty: none known).
std::vector<std::vector<double>> score_batch(const std::vector<const ScoreRequest*>& reqs);
std::vector<double> score(const ScoreRequest& r) { return score_batch({&r})[0]; }

Expand Down
88 changes: 88 additions & 0 deletions src/tree.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
// The questions of one call as a prefix tree over their tokens after the part the call shares: runs of
// tokens several prompts have in common (a choice's instructions and options, a score's scale) are read
// once, and each prompt still sees only its own tokens. Nothing here knows about the model.
#pragma once
#include <algorithm>
#include <cstdint>
#include <map>
#include <vector>

#include "tokens.hpp"

// Tokens [from, to) of prompt `prompt`, at positions from..to-1, read after its parent's (NO_PARENT: right
// after the shared part). A prompt's last token is always a node of its own, the leaf its answer is read at.
struct TreeNode {
static constexpr size_t NO_PARENT = SIZE_MAX;
size_t parent;
size_t prompt;
size_t from, to;
};

struct PromptTree {
std::vector<TreeNode> nodes; // depth first, a parent before its children, children in prompt order
std::vector<size_t> leaf; // per prompt, its leaf node
size_t tokens = 0; // tokens read: the nodes' lengths
};

// The tree of `prompts` after their first `skip` tokens, which they all share and are read apart (every
// prompt is longer than `skip`).
inline PromptTree prompt_tree(const std::vector<const Tokens*>& prompts, size_t skip) {
PromptTree t;
t.leaf.resize(prompts.size());
auto add = [&](size_t parent, size_t prompt, size_t from, size_t to) {
t.nodes.push_back({parent, prompt, from, to});
t.tokens += to - from;
return t.nodes.size() - 1;
};
// members share their tokens before `d`, already in the tree up to `parent`
auto grow = [&](auto& self, const std::vector<size_t>& members, size_t d, size_t parent) -> void {
std::vector<std::vector<size_t>> children; // in order of their first member
std::vector<bool> leaf;
std::map<int32_t, size_t> by_token;
for (size_t m : members) {
const Tokens& p = *prompts[m];
if (p.size() - 1 <= d) { // everything but the last token is in the tree: the leaf
children.push_back({m});
leaf.push_back(true);
continue;
}
auto [it, fresh] = by_token.emplace(p[d], children.size());
if (fresh) {
children.emplace_back();
leaf.push_back(false);
}
children[it->second].push_back(m);
}
for (size_t c = 0; c < children.size(); ++c) {
const std::vector<size_t>& group = children[c];
const Tokens& a = *prompts[group[0]];
if (leaf[c]) {
t.leaf[group[0]] = add(parent, group[0], a.size() - 1, a.size());
continue;
}
size_t end = a.size() - 1;
for (size_t m : group) end = std::min(end, prompts[m]->size() - 1);
size_t e = d + 1; // the group shares token d
while (e < end && std::all_of(group.begin(), group.end(), [&](size_t m) { return (*prompts[m])[e] == a[e]; })) ++e;
self(self, group, e, add(parent, group[0], d, e));
}
};
std::vector<size_t> all(prompts.size());
for (size_t i = 0; i < all.size(); ++i) all[i] = i;
grow(grow, all, skip, TreeNode::NO_PARENT);
return t;
}

// The tokens prompt `p` adds to a tree that already holds `members` (all after `skip`): what it does not
// share with any of them, and always its last token.
inline size_t added_tokens(const std::vector<const Tokens*>& prompts, const std::vector<size_t>& members, size_t p, size_t skip) {
const Tokens& a = *prompts[p];
size_t body = a.size() - 1, shared = skip;
for (size_t m : members) {
const Tokens& b = *prompts[m];
size_t k = skip, n = std::min(body, b.size() - 1);
while (k < n && a[k] == b[k]) ++k;
shared = std::max(shared, k);
}
return a.size() - shared;
}
Loading
Loading