Spaces:
Running
Running
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 };
}
|