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);