Skip to content
This repository was archived by the owner on Jan 15, 2026. It is now read-only.
Merged
Show file tree
Hide file tree
Changes from 2 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
24 changes: 24 additions & 0 deletions models/customs/config.json
Comment thread
solarpush marked this conversation as resolved.
Outdated
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
{
"_name_or_path": "nreimers/MiniLM-L6-H384-uncased",
"architectures": [
"BertModel"
],
"attention_probs_dropout_prob": 0.1,
"gradient_checkpointing": false,
"hidden_act": "gelu",
"hidden_dropout_prob": 0.1,
"hidden_size": 384,
"initializer_range": 0.02,
"intermediate_size": 1536,
"layer_norm_eps": 1e-12,
"max_position_embeddings": 512,
"model_type": "bert",
"num_attention_heads": 12,
"num_hidden_layers": 6,
"pad_token_id": 0,
"position_embedding_type": "absolute",
"transformers_version": "4.8.2",
"type_vocab_size": 2,
"use_cache": true,
"vocab_size": 30522
}
Binary file added models/customs/mymodel.onnx
Binary file not shown.
1 change: 1 addition & 0 deletions models/customs/special_tokens_map.json
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
{"unk_token": "[UNK]", "sep_token": "[SEP]", "pad_token": "[PAD]", "cls_token": "[CLS]", "mask_token": "[MASK]"}
1 change: 1 addition & 0 deletions models/customs/tokenizer.json

Large diffs are not rendered by default.

1 change: 1 addition & 0 deletions models/customs/tokenizer_config.json
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
{"do_lower_case": true, "unk_token": "[UNK]", "sep_token": "[SEP]", "pad_token": "[PAD]", "cls_token": "[CLS]", "mask_token": "[MASK]", "tokenize_chinese_chars": true, "strip_accents": null, "name_or_path": "nreimers/MiniLM-L6-H384-uncased", "do_basic_tokenize": true, "never_split": null, "tokenizer_class": "BertTokenizer", "model_max_length": 512}
70 changes: 51 additions & 19 deletions src/fastembed.ts
Original file line number Diff line number Diff line change
@@ -1,10 +1,10 @@
import { AddedToken, Tokenizer } from "@anush008/tokenizers";
import fs, { PathLike } from "fs";
import https from "https";
import * as ort from "onnxruntime-node";
import path from "path";
import Progress from "progress";
import tar from "tar";
import { AddedToken, Tokenizer } from "@anush008/tokenizers";
import * as ort from "onnxruntime-node";

export enum ExecutionProvider {
CPU = "cpu",
Expand All @@ -22,14 +22,14 @@ export enum EmbeddingModel {
BGESmallENV15 = "fast-bge-small-en-v1.5",
BGESmallZH = "fast-bge-small-zh-v1.5",
MLE5Large = "fast-multilingual-e5-large",
CUSTOM = "custom",
}

interface InitOptions {
model: EmbeddingModel;
executionProviders: ExecutionProvider[];
maxLength: number;
cacheDir: string;
showDownloadProgress: boolean;
export interface InitOptionsBase {
executionProviders?: ExecutionProvider[];
maxLength?: number;
cacheDir?: string;
showDownloadProgress?: boolean;
}

interface ModelInfo {
Expand Down Expand Up @@ -86,7 +86,20 @@ function getEmbeddings(

// return resultArray;
// }
// Cas standard
export interface InitStandardOptions extends InitOptionsBase {
model: Exclude<EmbeddingModel, EmbeddingModel.CUSTOM>;
modelAbsoluteDirPath?: undefined;
modelName?: string;
}

// Cas custom
export interface InitCustomOptions extends InitOptionsBase {
model: EmbeddingModel.CUSTOM;
modelAbsoluteDirPath: fs.PathLike;
modelName: string;
}
export type InitOptions = InitStandardOptions | InitCustomOptions;
abstract class Embedding {
abstract listSupportedModels(): ModelInfo[];

Expand All @@ -111,28 +124,47 @@ export class FlagEmbedding extends Embedding {
) {
super();
}

static async init(options: InitStandardOptions): Promise<FlagEmbedding>;
static async init(options: InitCustomOptions): Promise<FlagEmbedding>;
static async init({
model = EmbeddingModel.BGESmallENV15,
executionProviders = [ExecutionProvider.CPU],
maxLength = 512,
cacheDir = "local_cache",
showDownloadProgress = true,
modelAbsoluteDirPath = "",
modelName = "",
}: Partial<InitOptions> = {}) {
const modelDir = await FlagEmbedding.retrieveModel(
model,
cacheDir,
showDownloadProgress
);
if (model === EmbeddingModel.CUSTOM) {
if (!modelAbsoluteDirPath) {
throw new Error(
"For custom model, modelAbsoluteDirPath is required in FlagEmbedding.init"
);
}
if (!modelName) {
throw new Error(
"For custom model, modelName is required in FlagEmbedding.init"
);
}
}
const modelDir =
model === EmbeddingModel.CUSTOM
? modelAbsoluteDirPath
: await FlagEmbedding.retrieveModel(
model,
cacheDir,
showDownloadProgress
);

const tokenizer = this.loadTokenizer(modelDir, maxLength);

const modelPath = path.join(
modelDir.toString(),
const defaultModelName =
model === EmbeddingModel.MLE5Large ||
model === EmbeddingModel.AllMiniLML6V2
model === EmbeddingModel.AllMiniLML6V2
? "model.onnx"
: "model_optimized.onnx",
: "model_optimized.onnx";
const modelPath = path.join(
modelDir.toString(),
modelName || defaultModelName
);
if (!fs.existsSync(modelPath)) {
throw new Error(`Model file not found at ${modelPath}`);
Expand Down
123 changes: 123 additions & 0 deletions tests/fastembed_custom.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,123 @@
import path from "path";
import { expect, test } from "vitest";
import { EmbeddingModel, FlagEmbedding } from "../src";

test("Init EmbeddingModel", async () => {
const pathModel = path.resolve(__dirname, "../models/customs");
console.log("pathModel", pathModel);
const model = await FlagEmbedding.init({
model: EmbeddingModel.CUSTOM,
modelAbsoluteDirPath: pathModel,
modelName: "mymodel.onnx",
});
expect(model).toBeDefined();
});

test("FlagEmbedding embed", async () => {
const pathModel = path.resolve(__dirname, "../models/customs");
const flagEmbedding = await FlagEmbedding.init({
model: EmbeddingModel.CUSTOM,
modelAbsoluteDirPath: pathModel,
modelName: "mymodel.onnx",
maxLength: 512,
});
const embeddings = (await flagEmbedding.embed(["This is a test"]).next())
.value!;
expect(embeddings).toBeDefined();
expect(embeddings.length).toBe(1);
});

test("FlagEmbedding embed batch", async () => {
const pathModel = path.resolve(__dirname, "../models/customs");
const flagEmbedding = await FlagEmbedding.init({
model: EmbeddingModel.CUSTOM,
modelAbsoluteDirPath: pathModel,
modelName: "mymodel.onnx",

maxLength: 512,
});
const embeddingsBatch = flagEmbedding.embed([
"This is a test",
"Some text",
"Some more test",
"This is a test",
"Some text",
"Some more test",
]);
for await (const embeddings of embeddingsBatch) {
expect(embeddings).toBeDefined();
expect(embeddings.length).toBe(6);
expect(embeddings[0].length).toBe(384);
}
});

test("FlagEmbedding embed small batch", async () => {
const pathModel = path.resolve(__dirname, "../models/customs");
const flagEmbedding = await FlagEmbedding.init({
model: EmbeddingModel.CUSTOM,
modelAbsoluteDirPath: pathModel,
modelName: "mymodel.onnx",
maxLength: 512,
});
const embeddingsBatch = flagEmbedding.embed(
["This is a test", "Some text"],
1
);
for await (const embeddings of embeddingsBatch) {
expect(embeddings).toBeDefined();
expect(embeddings.length).toBe(1);
expect(embeddings[0].length).toBe(384);
}
});

test("FlagEmbedding queryEmbed", async () => {
const pathModel = path.resolve(__dirname, "../models/customs");
const flagEmbedding = await FlagEmbedding.init({
model: EmbeddingModel.CUSTOM,
modelAbsoluteDirPath: pathModel,
modelName: "mymodel.onnx",
maxLength: 512,
});
const embeddings = await flagEmbedding.queryEmbed("This is a test");
expect(embeddings).toBeDefined();
expect(embeddings.length).toBe(384);
});

test("FlagEmbedding passageEmbed", async () => {
const pathModel = path.resolve(__dirname, "../models/customs");
const flagEmbedding = await FlagEmbedding.init({
model: EmbeddingModel.CUSTOM,
modelAbsoluteDirPath: pathModel,
modelName: "mymodel.onnx",

maxLength: 512,
});
const embeddings = (
await flagEmbedding.passageEmbed(["This is a test"]).next()
).value!;
expect(embeddings).toBeDefined();
expect(embeddings.length).toBe(1);
});

test("FlagEmbedding canonical values", async () => {
const pathModel = path.resolve(__dirname, "../models/customs");
const flagEmbedding = await FlagEmbedding.init({
model: EmbeddingModel.CUSTOM,
modelAbsoluteDirPath: pathModel,
modelName: "mymodel.onnx",
maxLength: 512,
});
const expected = [
0.025276897475123405, 0.013033483177423477, 0.005586996208876371,
0.04152565822005272, -0.018848471343517303, -0.05523142218589783,
0.018086062744259834, -0.000535094877704978, -0.013765564188361168,
-0.016923097893595695,
];

const embeddings = (await flagEmbedding.embed(["hello world"]).next()).value!;
console.log("embeddings", embeddings[0].slice(0, 10));
expect(embeddings).toBeDefined();
for (let i = 0; i < expected.length; i++) {
expect(embeddings[0][i]).toBeCloseTo(expected[i], 3);
}
});
3 changes: 2 additions & 1 deletion tsconfig.json
Original file line number Diff line number Diff line change
Expand Up @@ -13,5 +13,6 @@
"skipLibCheck": true,
"types": ["vitest/globals"]
},
"include": ["./src"]
"include": ["./src"],
"exclude": ["./models"]
}
Loading