Spaces:
Running
Running
File size: 10,702 Bytes
3bbbb17 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 | /**
* The bi-encoder's ONNX session, running off the main thread.
*
* `onnxruntime-node` runs inference synchronously on whichever JS thread calls
* it, so scoring a 200-passage pool used to stop the server's event loop for
* hundreds of milliseconds β no other request could be served while a single
* page-content read was ranking its passages. Here that work blocks only this
* thread. The model is loaded here for the same reason: the ~450 MB of weights
* must live in the worker's memory, not the main thread's.
*
* Node loads this file directly (type stripping), so keep it to syntax that
* erases: no enums, no parameter properties, no namespaces.
*/
import { availableParallelism } from "node:os";
import path from "node:path";
import { parentPort } from "node:worker_threads";
import type { Tokenizer } from "@huggingface/tokenizers";
import { type InferenceSession, Tensor } from "onnxruntime-node";
import {
BATCH_ROWS,
type BiEncoderRequest,
type BiEncoderResponse,
} from "./biEncoderWorkerProtocol.ts";
import { createModelLogger, loadOnnxModel } from "./utils/onnxModelLoader.ts";
const printMessage = createModelLogger(path.basename(import.meta.url));
const MODEL_HF_REPO =
"sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2";
/**
* The ONNX export. ~450 MB, multilingual (50+ languages), 384-dimensional
* embeddings. CPU cost depends on passage length and pool size; the page-content
* selector limits dense scoring to 256 passages. See docs/page-content.md.
*/
const MODEL_HF_FILE = "onnx/model.onnx";
/**
* Maximum tokens per encoding. The model was trained with a 256-token limit;
* passages longer than that are truncated from the end, which is where the
* passage content sits after the query prefix.
*/
const MAX_SEQUENCE_LENGTH = 256;
/**
* Half the reported parallelism, which is the physical core count on the
* hyperthreaded x86 hosts this runs on, and at least one thread everywhere
* else.
*
* ONNX Runtime otherwise sizes its intra-op pool from the logical core count
* and its threads spin before yielding, so it oversubscribes every core and
* leaves nothing for the thread that has to answer HTTP. Measured on a
* 200-passage pool on this host (x86_64, 32 logical cores), worst main-thread
* event-loop delay against wall time for the pass:
*
* ORT default (32): 31-38 ms, 7.1 s
* 16: 1-4 ms, 2.6 s
* 8: 1-2 ms, 4.2 s
* 4: 1-1 ms, 5.7 s
*
* Capping is not a trade here: 16 was both the quietest and the fastest, and
* the oversubscribed default was the slowest setting measured.
*/
const INTRA_OP_THREADS = Math.max(1, Math.floor(availableParallelism() / 2));
interface EncodeResult {
embeddings: Float32Array[];
/** `input_ids` dims of each forward pass, in run order. */
runDimensions: number[][];
}
/**
* Encodes a batch of texts into normalized embedding vectors with one
* `session.run()` per length bucket.
*
* Rows are sorted by token length and cut into buckets of at most
* `BATCH_ROWS`, so each batch pads to its own near-uniform width instead of
* the global max: attention is O(L^2), and padding a mixed-length chunk out
* to 256 can multiply the FLOPs several-fold and come out slower than the
* per-text path.
*
* Right-padding with the real pad id under an attention mask is score-safe for
* this fp32 export: it has no per-tensor dynamic quantization scale, so a
* padded row cannot move another row's scale, and the real tokens' hidden
* states are bit-identical to encoding the text alone. (The cross-encoder is
* dynamically quantized and cannot batch for that reason; see the note in
* `rerankerService.ts`.)
*
* Buckets run one after another with a plain `await` β deliberately NOT a
* `Promise.all` across buckets. Each await hands this thread's event loop a
* turn between buckets, which is what lets two overlapping scoring requests
* interleave instead of one waiting out the other. Do not "optimize" the
* buckets back into a `Promise.all`: onnxruntime-node runs inference
* synchronously on the calling thread regardless, so it would add no
* concurrency while taking away those between-bucket turns.
*/
async function encodeBatch(
activeSession: InferenceSession,
loadedTokenizer: Tokenizer,
activePadTokenId: number,
texts: string[],
): Promise<EncodeResult> {
const tokenized = texts
.map((text, index) => {
const { ids, attention_mask } = loadedTokenizer.encode(text);
const slicedIds = ids.slice(0, MAX_SEQUENCE_LENGTH);
const slicedMask = attention_mask.slice(0, MAX_SEQUENCE_LENGTH);
// The real-token length comes from the mask's leading 1s, not from
// trusting every returned id as real. At @huggingface/tokenizers
// 0.2.0 the mask is all 1s, so this is a no-op today; it keeps the
// row correct if a future tokenizer pads or marks truncation in the
// mask (the cached tokenizer.json already declares a padding
// strategy with pad_id 1).
const firstPad = slicedMask.indexOf(0);
const realLength = firstPad === -1 ? slicedIds.length : firstPad;
return { index, ids: slicedIds.slice(0, realLength) };
})
.sort((a, b) => a.ids.length - b.ids.length);
const embeddings: Float32Array[] = new Array(texts.length);
const runDimensions: number[][] = [];
for (let start = 0; start < tokenized.length; start += BATCH_ROWS) {
const rows = tokenized.slice(start, start + BATCH_ROWS);
// Sorted ascending, so the last row sets the bucket width.
const bucketLength = rows[rows.length - 1].ids.length;
const dimensions = [rows.length, bucketLength];
const inputIds = new BigInt64Array(rows.length * bucketLength).fill(
BigInt(activePadTokenId),
);
const attentionMask = new BigInt64Array(rows.length * bucketLength);
// The export declares `token_type_ids` and ONNX Runtime refuses to run
// with a declared input missing. Every text is one segment, so zeros.
const tokenTypeIds = new BigInt64Array(rows.length * bucketLength);
for (let row = 0; row < rows.length; row++) {
const base = row * bucketLength;
const { ids } = rows[row];
for (let t = 0; t < ids.length; t++) {
inputIds[base + t] = BigInt(ids[t]);
attentionMask[base + t] = 1n;
}
}
runDimensions.push(dimensions);
const { last_hidden_state } = await activeSession.run({
input_ids: new Tensor("int64", inputIds, dimensions),
attention_mask: new Tensor("int64", attentionMask, dimensions),
token_type_ids: new Tensor("int64", tokenTypeIds, dimensions),
});
const hidden = last_hidden_state.data as Float32Array;
const dim = last_hidden_state.dims[2];
// Take the row stride from the output tensor, not from the bucket
// length we asked for: same value today, but the pooling can never
// read the wrong row if the export's layout ever changes.
const seqStride = last_hidden_state.dims[1];
for (let row = 0; row < rows.length; row++) {
// Right-padding keeps every real token inside `ids.length`, so
// pooling over that count never touches a pad. (Checking the mask
// here instead would have to compare against 0n: BigInt64Array
// entries are never `=== 0`.)
const rowLength = rows[row].ids.length;
const pooled = new Float32Array(dim);
for (let t = 0; t < rowLength; t++) {
const offset = (row * seqStride + t) * dim;
for (let d = 0; d < dim; d++) {
pooled[d] += hidden[offset + d];
}
}
if (rowLength > 0) {
for (let d = 0; d < dim; d++) {
pooled[d] /= rowLength;
}
}
// L2 normalize.
let norm = 0;
for (let d = 0; d < dim; d++) {
norm += pooled[d] * pooled[d];
}
norm = Math.sqrt(norm);
if (norm > 0) {
for (let d = 0; d < dim; d++) {
pooled[d] /= norm;
}
}
embeddings[rows[row].index] = pooled;
}
}
return { embeddings, runDimensions };
}
/**
* Computes cosine similarity between a query embedding and passage embeddings.
* Both are assumed to be L2-normalized, so cosine similarity = dot product.
*/
function cosineSimilarities(
query: Float32Array,
passages: Float32Array[],
): number[] {
return passages.map((passage) => {
let sum = 0;
for (let d = 0; d < query.length; d++) {
sum += query[d] * passage[d];
}
return sum;
});
}
if (!parentPort) {
throw new Error(
"biEncoderWorker.ts is a worker entry point and must be started with node:worker_threads",
);
}
const port = parentPort;
printMessage("Preparing bi-encoder model...");
const loaded = await loadOnnxModel(MODEL_HF_REPO, MODEL_HF_FILE, {
sessionOptions: { intraOpNumThreads: INTRA_OP_THREADS },
});
if (loaded.padTokenId === null) {
throw new Error(
`The bi-encoder tokenizer (${MODEL_HF_REPO}) declares no pad_token; batched scoring needs the real pad id and will not assume one`,
);
}
const session = loaded.session;
const tokenizer = loaded.tokenizer;
const padTokenId = loaded.padTokenId;
// Warm up with a test encoding, so the first real request does not pay for
// the lazily allocated arenas of the first run.
await encodeBatch(session, tokenizer, padTokenId, ["test query"]);
port.on("message", async (request: BiEncoderRequest) => {
if (request.type !== "score") return;
try {
// The query rides in the same batched pass as the passages instead of
// running alone, saving one forward pass. Its position in the result is
// guaranteed by the index-based restore inside encodeBatch
// (`embeddings[rows[row].index]`), not by where the length sort happens
// to place it, so it stays correct whatever the bucketing does.
const { embeddings, runDimensions } = await encodeBatch(
session,
tokenizer,
padTokenId,
[request.query, ...request.passages],
);
const [queryEmbedding, ...passageEmbeddings] = embeddings;
const response: BiEncoderResponse = {
type: "scores",
id: request.id,
scores: cosineSimilarities(queryEmbedding, passageEmbeddings),
runDimensions,
};
port.postMessage(response);
} catch (error) {
// The caller falls back to lexical ranking on empty scores, so one bad
// request must not take the worker down with it.
const response: BiEncoderResponse = {
type: "failed",
id: request.id,
message: error instanceof Error ? error.message : String(error),
};
port.postMessage(response);
}
});
const ready: BiEncoderResponse = { type: "ready" };
port.postMessage(ready);
|