diff --git a/packages/mnemopi/src/core/embeddings.ts b/packages/mnemopi/src/core/embeddings.ts index 241122364..21851d2a1 100644 --- a/packages/mnemopi/src/core/embeddings.ts +++ b/packages/mnemopi/src/core/embeddings.ts @@ -11,7 +11,7 @@ import { } from "@oh-my-pi/pi-utils"; import type { EmbeddingModel } from "fastembed"; import { LRUCache } from "lru-cache/raw"; -import { ensureFastembedTokenizerSidecars } from "./fastembed-model-cache"; +import { ensureFastembedModelSidecars } from "./fastembed-model-cache"; import { loadFastembed } from "./fastembed-runtime"; import { type EmbeddingOutput, @@ -68,11 +68,13 @@ async function defaultLocalModelInitializer(options: LocalModelInitOptions): Pro } catch (error) { const message = error instanceof Error ? error.message : ""; if ( - !/(?:Tokenizer file not found at .*tokenizer|Tokens map file not found at .*special_tokens_map)/u.test(message) + !/(?:Config file not found at .*config|Tokenizer file not found at .*tokenizer|Tokens map file not found at .*special_tokens_map)/u.test( + message, + ) ) { throw error; } - if (!(await ensureFastembedTokenizerSidecars(options.model, options.cacheDir))) throw error; + if (!(await ensureFastembedModelSidecars(options.model, options.cacheDir))) throw error; return FlagEmbedding.init(options); } } diff --git a/packages/mnemopi/src/core/fastembed-model-cache.ts b/packages/mnemopi/src/core/fastembed-model-cache.ts index e88cf70a7..f31354765 100644 --- a/packages/mnemopi/src/core/fastembed-model-cache.ts +++ b/packages/mnemopi/src/core/fastembed-model-cache.ts @@ -1,6 +1,11 @@ import * as path from "node:path"; -const FASTEMBED_TOKENIZER_SIDECARS = ["tokenizer.json", "tokenizer_config.json", "special_tokens_map.json"] as const; +const FASTEMBED_MODEL_SIDECARS = [ + "config.json", + "tokenizer.json", + "tokenizer_config.json", + "special_tokens_map.json", +] as const; const FASTEMBED_HF_REPOS: Record = { "fast-all-MiniLM-L6-v2": "sentence-transformers/all-MiniLM-L6-v2", @@ -12,13 +17,13 @@ const FASTEMBED_HF_REPOS: Record = { "fast-multilingual-e5-large": "intfloat/multilingual-e5-large", }; -/** Download missing tokenizer sidecars into a fastembed model cache directory. */ -export async function ensureFastembedTokenizerSidecars(model: string, cacheDir = "local_cache"): Promise { +/** Download missing config/tokenizer sidecars into a fastembed model cache directory. */ +export async function ensureFastembedModelSidecars(model: string, cacheDir = "local_cache"): Promise { const repo = FASTEMBED_HF_REPOS[model]; if (repo === undefined) return false; const modelDir = path.join(cacheDir, model); - for (const fileName of FASTEMBED_TOKENIZER_SIDECARS) { + for (const fileName of FASTEMBED_MODEL_SIDECARS) { const target = path.join(modelDir, fileName); if (await Bun.file(target).exists()) continue; diff --git a/packages/mnemopi/test/fastembed-model-cache.test.ts b/packages/mnemopi/test/fastembed-model-cache.test.ts index 890ea21dc..3097a8608 100644 --- a/packages/mnemopi/test/fastembed-model-cache.test.ts +++ b/packages/mnemopi/test/fastembed-model-cache.test.ts @@ -2,10 +2,10 @@ import { describe, expect, it, spyOn } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; -import { ensureFastembedTokenizerSidecars } from "../src/core/fastembed-model-cache"; +import { ensureFastembedModelSidecars } from "../src/core/fastembed-model-cache"; describe("fastembed model cache repair", () => { - it("downloads missing tokenizer sidecars without overwriting cached files", async () => { + it("downloads missing config and tokenizer sidecars without overwriting cached files", async () => { const cacheDir = await fs.mkdtemp(path.join(os.tmpdir(), "mnemopi-fastembed-")); const model = "fast-bge-base-en-v1.5"; const modelDir = path.join(cacheDir, model); @@ -25,12 +25,14 @@ describe("fastembed model cache repair", () => { await fs.mkdir(modelDir, { recursive: true }); await Bun.write(path.join(modelDir, "tokenizer.json"), "cached-tokenizer"); - expect(await ensureFastembedTokenizerSidecars(model, cacheDir)).toBe(true); + expect(await ensureFastembedModelSidecars(model, cacheDir)).toBe(true); expect(requested).toEqual([ + "https://huggingface.co/BAAI/bge-base-en-v1.5/resolve/main/config.json", "https://huggingface.co/BAAI/bge-base-en-v1.5/resolve/main/tokenizer_config.json", "https://huggingface.co/BAAI/bge-base-en-v1.5/resolve/main/special_tokens_map.json", ]); + expect(await Bun.file(path.join(modelDir, "config.json")).text()).toBe("body:config.json"); expect(await Bun.file(path.join(modelDir, "tokenizer.json")).text()).toBe("cached-tokenizer"); expect(await Bun.file(path.join(modelDir, "tokenizer_config.json")).text()).toBe("body:tokenizer_config.json"); expect(await Bun.file(path.join(modelDir, "special_tokens_map.json")).text()).toBe( @@ -52,7 +54,7 @@ describe("fastembed model cache repair", () => { ), ); try { - expect(await ensureFastembedTokenizerSidecars("unknown-model", "/tmp/missing")).toBe(false); + expect(await ensureFastembedModelSidecars("unknown-model", "/tmp/missing")).toBe(false); expect(fetchSpy).not.toHaveBeenCalled(); } finally { fetchSpy.mockRestore();