File size: 5,759 Bytes
6c3af4e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3bbbb17
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6c3af4e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3bbbb17
 
 
 
 
6c3af4e
3bbbb17
 
 
 
 
6c3af4e
 
 
 
3bbbb17
6c3af4e
 
 
6dcf170
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6c3af4e
 
6dcf170
 
6c3af4e
 
 
 
3bbbb17
 
 
 
 
6dcf170
 
 
 
 
6c3af4e
 
 
 
 
 
6dcf170
 
 
6c3af4e
 
6dcf170
6c3af4e
6dcf170
3bbbb17
6c3af4e
6dcf170
6c3af4e
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
import fs from "node:fs";
import path from "node:path";
import { fileURLToPath } from "node:url";
import { Tokenizer } from "@huggingface/tokenizers";
import debug from "debug";
import { InferenceSession } from "onnxruntime-node";
import { downloadFileFromHuggingFaceRepository } from "../downloadFileFromHuggingFaceRepository.ts";

const SERVER_DIR = path.resolve(
  path.dirname(fileURLToPath(import.meta.url)),
  "..",
);

const printMessage = createModelLogger(path.basename(import.meta.url));

/**
 * Creates a debug logger with the same enabled-always behavior used by the
 * model services.
 */
export function createModelLogger(moduleName: string) {
  const printMessage = debug(moduleName);
  printMessage.enabled = true;
  return printMessage;
}

/**
 * Session knobs a caller may override. Left unset, ONNX Runtime picks its own
 * defaults, which is what every caller but the bi-encoder worker wants.
 */
interface OnnxSessionOptions {
  intraOpNumThreads?: number;
}

export interface LoadOnnxModelOptions {
  /** Tokenizer file names, for a repo that does not use the standard ones. */
  tokenizerFile?: string;
  tokenizerConfigFile?: string;
  sessionOptions?: OnnxSessionOptions;
}

function resolveModelPath(modelRepo: string, hfRepoFile: string) {
  return path.resolve(SERVER_DIR, "models", modelRepo, hfRepoFile);
}

// Replaces a cached file when its size differs from the Hub, not just when
// missing.
async function ensureModelFileExists(modelRepo: string, hfRepoFile: string) {
  const localPath = resolveModelPath(modelRepo, hfRepoFile);
  await downloadFileFromHuggingFaceRepository(modelRepo, hfRepoFile, localPath);
  return localPath;
}

// CPU-only, errors-only logging. The reranker is a dynamically quantized
// graph, the wrong shape for WebGPU (no kernels for the integer matmuls,
// shuttles them to the CPU: 812ms against 172ms for the same work, with scores
// drifting by up to 1.15 and reordering results). `coreml` is slower than CPU
// on dynamic shapes. The bi-encoder is fp32 and was not benchmarked on
// accelerators, but inherits this choice. logSeverityLevel 3 keeps startup
// quiet: ONNX Runtime otherwise warns that it assigned shape operators to CPU,
// which is expected and not actionable. The architecture is logged because it
// selects the quantized kernel, which is the part that varies between hosts.
function createOnnxSession(
  modelRepo: string,
  modelPath: string,
  sessionOptions: OnnxSessionOptions = {},
) {
  printMessage(
    `Creating CPU session for ${modelRepo} (arch: ${process.arch}, platform: ${process.platform}${
      sessionOptions.intraOpNumThreads === undefined
        ? ""
        : `, intra-op threads: ${sessionOptions.intraOpNumThreads}`
    })...`,
  );
  return InferenceSession.create(modelPath, {
    executionProviders: ["cpu"],
    logSeverityLevel: 3,
    ...sessionOptions,
  });
}

/**
 * Resolves the tokenizer's real pad token id from its config. Returns the
 * resolved id, or `null` for anything it cannot cleanly resolve: no
 * `pad_token` key, an empty value, a shape it does not understand, or a
 * `token_to_id` miss. It never throws β€” this loader is shared, and a pad
 * id it cannot read must not take a service down at boot. The caller that
 * actually needs a pad id (the bi-encoder) decides what to do about null.
 *
 * Both serializations transformers emit are accepted: a plain string, and
 * the AddedToken dict form (e.g. `{"content": "<pad>", "lstrip": false,
 * "rstrip": false, "normalized": true}`), which the cached configs here
 * already use for `mask_token`. Padding with an assumed id is never on
 * the table: in the XLM-RoBERTa exports id 0 is `<s>`, not `<pad>`.
 */
function resolvePadTokenId(
  tokenizer: Tokenizer,
  tokenizerConfig: { pad_token?: unknown },
): number | null {
  const padToken = tokenizerConfig?.pad_token;

  let padTokenString: string | null = null;
  if (typeof padToken === "string" && padToken.length > 0) {
    padTokenString = padToken;
  } else if (
    typeof padToken === "object" &&
    padToken !== null &&
    typeof (padToken as { content?: unknown }).content === "string" &&
    (padToken as { content: string }).content.length > 0
  ) {
    padTokenString = (padToken as { content: string }).content;
  }
  if (padTokenString === null) {
    return null;
  }

  if (typeof tokenizer.token_to_id !== "function") {
    return null;
  }

  const padTokenId = tokenizer.token_to_id(padTokenString);
  return padTokenId === undefined ? null : padTokenId;
}

/**
 * Downloads model files from a Hugging Face repo and returns a ready
 * inference session, tokenizer, and the tokenizer's real pad token id. The
 * tokenizer files use the standard Hugging Face names when not overridden.
 */
export async function loadOnnxModel(
  modelRepo: string,
  modelFile: string,
  {
    tokenizerFile = "tokenizer.json",
    tokenizerConfigFile = "tokenizer_config.json",
    sessionOptions = {},
  }: LoadOnnxModelOptions = {},
): Promise<{
  session: InferenceSession;
  tokenizer: Tokenizer;
  padTokenId: number | null;
}> {
  const [modelPath, tokenizerPath, tokenizerConfigPath] = await Promise.all([
    ensureModelFileExists(modelRepo, modelFile),
    ensureModelFileExists(modelRepo, tokenizerFile),
    ensureModelFileExists(modelRepo, tokenizerConfigFile),
  ]);

  const tokenizerConfig = JSON.parse(
    fs.readFileSync(tokenizerConfigPath, "utf8"),
  ) as { pad_token?: unknown };
  const tokenizer = new Tokenizer(
    JSON.parse(fs.readFileSync(tokenizerPath, "utf8")),
    tokenizerConfig,
  );
  const padTokenId = resolvePadTokenId(tokenizer, tokenizerConfig);
  const session = await createOnnxSession(modelRepo, modelPath, sessionOptions);

  return { session, tokenizer, padTokenId };
}