feat(ts-sdk): add HuggingFace embedding provider (#6027)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -5,6 +5,10 @@ description: "Configure Hugging Face as an embedding provider in Mem0 for local
|
||||
|
||||
You can use embedding models from Huggingface to run Mem0 locally.
|
||||
|
||||
<Note>
|
||||
The TypeScript SDK supports Hugging Face only through a hosted [Text Embeddings Inference (TEI)](#using-text-embeddings-inference-tei) endpoint, or any OpenAI-compatible Hugging Face endpoint. The local `sentence-transformers` mode shown first is Python-only.
|
||||
</Note>
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
@@ -34,9 +38,10 @@ m.add(messages, user_id="john")
|
||||
|
||||
### Using Text Embeddings Inference (TEI)
|
||||
|
||||
You can also use Hugging Face's Text Embeddings Inference service for faster and more efficient embeddings:
|
||||
You can also use Hugging Face's Text Embeddings Inference service for faster and more efficient embeddings. This is the mode the TypeScript SDK uses.
|
||||
|
||||
```python
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
@@ -56,6 +61,24 @@ m = Memory.from_config(config)
|
||||
m.add("This text will be embedded using the TEI service.", user_id="john")
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from 'mem0ai/oss';
|
||||
|
||||
// Point at a running TEI server, or any OpenAI-compatible HF endpoint
|
||||
const config = {
|
||||
embedder: {
|
||||
provider: 'huggingface',
|
||||
config: {
|
||||
huggingfaceBaseUrl: 'http://localhost:3000/v1',
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
await memory.add("This text will be embedded using the TEI service.", { userId: "john" });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
To run the TEI service, you can use Docker:
|
||||
|
||||
```bash
|
||||
@@ -66,11 +89,22 @@ docker run -d -p 3000:80 -v huggingfacetei:/data --platform linux/amd64 \
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Huggingface embedder:
|
||||
Here are the parameters available for configuring the Hugging Face embedder:
|
||||
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `model` | The name of the model to use | `multi-qa-MiniLM-L6-cos-v1` |
|
||||
| `embedding_dims` | Dimensions of the embedding model | `selected_model_dimensions` |
|
||||
| `model_kwargs` | Additional arguments for the model | `None` |
|
||||
| `huggingface_base_url` | URL to connect to Text Embeddings Inference (TEI) API | `None` |
|
||||
| `huggingface_base_url` | URL to connect to Text Embeddings Inference (TEI) API | `None` |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `huggingfaceBaseUrl` | TEI or OpenAI-compatible endpoint URL. Required; falls back to `baseURL`, `url`, then the `HUGGINGFACE_BASE_URL` env var | `None` |
|
||||
| `model` | Model name sent to the endpoint (TEI ignores it) | `tei` |
|
||||
| `apiKey` | API key for the endpoint; falls back to the `HUGGINGFACE_API_KEY` env var | `"hf"` |
|
||||
</Tab>
|
||||
</Tabs>
|
||||
@@ -0,0 +1,78 @@
|
||||
import OpenAI from "openai";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
|
||||
/**
|
||||
* HuggingFace embedding provider (hosted inference mode).
|
||||
*
|
||||
* Mirrors the `huggingface_base_url` branch of the Python provider
|
||||
* (`mem0/embeddings/huggingface.py`): a HuggingFace Text Embeddings Inference
|
||||
* (TEI) server, or any HuggingFace OpenAI-compatible inference endpoint,
|
||||
* exposes a `/v1/embeddings` route, so this embedder reuses the existing
|
||||
* `openai` client pointed at that base URL. No new dependency is required.
|
||||
*
|
||||
* A base URL is required. The Python provider's alternative local
|
||||
* `sentence-transformers` path has no lightweight TypeScript equivalent, so
|
||||
* hosted inference is the supported TS mode.
|
||||
*/
|
||||
export class HuggingFaceEmbedder implements Embedder {
|
||||
private openai: OpenAI;
|
||||
private model: string;
|
||||
|
||||
constructor(config: EmbeddingConfig) {
|
||||
const baseURL =
|
||||
config.huggingfaceBaseUrl ||
|
||||
config.baseURL ||
|
||||
config.url ||
|
||||
process.env.HUGGINGFACE_BASE_URL;
|
||||
|
||||
if (!baseURL) {
|
||||
throw new Error(
|
||||
"HuggingFace embedder requires an inference endpoint. Set " +
|
||||
"`huggingfaceBaseUrl` (or `baseURL`) in the embedder config, or the " +
|
||||
"HUGGINGFACE_BASE_URL environment variable (e.g. a TEI server at " +
|
||||
"http://localhost:8080/v1).",
|
||||
);
|
||||
}
|
||||
|
||||
this.openai = new OpenAI({
|
||||
apiKey: config.apiKey || process.env.HUGGINGFACE_API_KEY || "hf",
|
||||
baseURL,
|
||||
});
|
||||
// TEI ignores the model field; default mirrors the Python provider.
|
||||
this.model = config.model || "tei";
|
||||
}
|
||||
|
||||
async embed(text: string): Promise<number[]> {
|
||||
const response = await this.openai.embeddings.create({
|
||||
model: this.model,
|
||||
input: text,
|
||||
});
|
||||
if (!response.data || response.data.length === 0) {
|
||||
throw new Error(
|
||||
`HuggingFace embed() returned no embeddings for model '${this.model}'`,
|
||||
);
|
||||
}
|
||||
return response.data[0].embedding;
|
||||
}
|
||||
|
||||
async embedBatch(texts: string[]): Promise<number[][]> {
|
||||
if (texts.length === 0) {
|
||||
return [];
|
||||
}
|
||||
const response = await this.openai.embeddings.create({
|
||||
model: this.model,
|
||||
input: texts,
|
||||
});
|
||||
const embeddings = response.data
|
||||
.sort((a, b) => a.index - b.index)
|
||||
.map((item) => item.embedding);
|
||||
if (embeddings.length !== texts.length) {
|
||||
throw new Error(
|
||||
`HuggingFace embedBatch() returned ${embeddings.length} embeddings ` +
|
||||
`for ${texts.length} texts using model '${this.model}'`,
|
||||
);
|
||||
}
|
||||
return embeddings;
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ export * from "./memory";
|
||||
export * from "./memory/memory.types";
|
||||
export * from "./types";
|
||||
export * from "./embeddings/base";
|
||||
export * from "./embeddings/huggingface";
|
||||
export * from "./embeddings/openai";
|
||||
export * from "./embeddings/ollama";
|
||||
export * from "./embeddings/lmstudio";
|
||||
|
||||
@@ -19,6 +19,8 @@ export interface EmbeddingConfig {
|
||||
url?: string;
|
||||
embeddingDims?: number;
|
||||
modelProperties?: Record<string, any>;
|
||||
// HuggingFace TEI / OpenAI-compatible inference endpoint base URL.
|
||||
huggingfaceBaseUrl?: string;
|
||||
}
|
||||
|
||||
export type { ValkeyConfig } from "./valkey";
|
||||
|
||||
@@ -43,6 +43,7 @@ import { AzureOpenAIEmbedder } from "../embeddings/azure";
|
||||
import { FastEmbedEmbedder } from "../embeddings/fastembed";
|
||||
import { LangchainLLM } from "../llms/langchain";
|
||||
import { LangchainEmbedder } from "../embeddings/langchain";
|
||||
import { HuggingFaceEmbedder } from "../embeddings/huggingface";
|
||||
import { LangchainVectorStore } from "../vector_stores/langchain";
|
||||
import { AzureAISearch } from "../vector_stores/azure_ai_search";
|
||||
import { PGVector } from "../vector_stores/pgvector";
|
||||
@@ -77,6 +78,8 @@ export class EmbedderFactory {
|
||||
return new FastEmbedEmbedder(config);
|
||||
case "langchain":
|
||||
return new LangchainEmbedder(config);
|
||||
case "huggingface":
|
||||
return new HuggingFaceEmbedder(config);
|
||||
default:
|
||||
throw new Error(`Unsupported embedder provider: ${provider}`);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
/// <reference types="jest" />
|
||||
/**
|
||||
* HuggingFace Embedder unit tests (mocked OpenAI client).
|
||||
* The TS provider targets a HuggingFace TEI / OpenAI-compatible inference
|
||||
* endpoint, so it reuses the `openai` client with a HuggingFace baseURL.
|
||||
* These tests verify the required base URL, request shape, and batch ordering.
|
||||
*/
|
||||
|
||||
const mockEmbeddingsCreate = jest.fn();
|
||||
const mockOpenAICtor = jest.fn();
|
||||
|
||||
jest.mock("openai", () => {
|
||||
return {
|
||||
__esModule: true,
|
||||
default: jest.fn().mockImplementation((opts: any) => {
|
||||
mockOpenAICtor(opts);
|
||||
return { embeddings: { create: mockEmbeddingsCreate } };
|
||||
}),
|
||||
};
|
||||
});
|
||||
|
||||
import { HuggingFaceEmbedder } from "../src/embeddings/huggingface";
|
||||
|
||||
const mockEmbedding = [0.1, 0.2, 0.3, 0.4, 0.5];
|
||||
|
||||
describe("HuggingFaceEmbedder (unit)", () => {
|
||||
const OLD_ENV = process.env;
|
||||
|
||||
beforeEach(() => {
|
||||
mockEmbeddingsCreate.mockReset();
|
||||
mockOpenAICtor.mockReset();
|
||||
mockEmbeddingsCreate.mockResolvedValue({
|
||||
data: [{ index: 0, embedding: mockEmbedding }],
|
||||
});
|
||||
process.env = { ...OLD_ENV };
|
||||
delete process.env.HUGGINGFACE_BASE_URL;
|
||||
});
|
||||
|
||||
afterAll(() => {
|
||||
process.env = OLD_ENV;
|
||||
});
|
||||
|
||||
describe("configuration", () => {
|
||||
it("throws when no inference endpoint is configured", () => {
|
||||
expect(() => new HuggingFaceEmbedder({ apiKey: "test-key" })).toThrow(
|
||||
/requires an inference endpoint/,
|
||||
);
|
||||
});
|
||||
|
||||
it("uses huggingfaceBaseUrl and the default model", async () => {
|
||||
const embedder = new HuggingFaceEmbedder({
|
||||
apiKey: "test-key",
|
||||
huggingfaceBaseUrl: "http://localhost:8080/v1",
|
||||
});
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(mockOpenAICtor.mock.calls[0][0]).toMatchObject({
|
||||
apiKey: "test-key",
|
||||
baseURL: "http://localhost:8080/v1",
|
||||
});
|
||||
const callArgs = mockEmbeddingsCreate.mock.calls[0][0];
|
||||
expect(callArgs).toEqual({ model: "tei", input: "hello" });
|
||||
});
|
||||
|
||||
it("falls back to baseURL and honors a custom model", async () => {
|
||||
const embedder = new HuggingFaceEmbedder({
|
||||
baseURL: "https://tei.example.com/v1",
|
||||
model: "BAAI/bge-small-en-v1.5",
|
||||
});
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(mockOpenAICtor.mock.calls[0][0]).toMatchObject({
|
||||
baseURL: "https://tei.example.com/v1",
|
||||
});
|
||||
expect(mockEmbeddingsCreate.mock.calls[0][0].model).toBe(
|
||||
"BAAI/bge-small-en-v1.5",
|
||||
);
|
||||
});
|
||||
|
||||
it("reads HUGGINGFACE_BASE_URL from the environment", async () => {
|
||||
process.env.HUGGINGFACE_BASE_URL = "http://env-host:8080/v1";
|
||||
const embedder = new HuggingFaceEmbedder({ apiKey: "test-key" });
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(mockOpenAICtor.mock.calls[0][0]).toMatchObject({
|
||||
baseURL: "http://env-host:8080/v1",
|
||||
});
|
||||
});
|
||||
|
||||
it("never forwards a dimensions parameter", async () => {
|
||||
const embedder = new HuggingFaceEmbedder({
|
||||
huggingfaceBaseUrl: "http://localhost:8080/v1",
|
||||
embeddingDims: 384,
|
||||
});
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(mockEmbeddingsCreate.mock.calls[0][0]).not.toHaveProperty(
|
||||
"dimensions",
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
describe("basic functionality", () => {
|
||||
const cfg = { huggingfaceBaseUrl: "http://localhost:8080/v1" };
|
||||
|
||||
it("embed() returns the embedding vector", async () => {
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
expect(await embedder.embed("hello")).toEqual(mockEmbedding);
|
||||
});
|
||||
|
||||
it("embedBatch() returns [] for empty input without calling the API", async () => {
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
expect(await embedder.embedBatch([])).toEqual([]);
|
||||
expect(mockEmbeddingsCreate).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("embedBatch() sorts results by index", async () => {
|
||||
mockEmbeddingsCreate.mockResolvedValue({
|
||||
data: [
|
||||
{ index: 1, embedding: [0.3, 0.4] },
|
||||
{ index: 0, embedding: [0.1, 0.2] },
|
||||
],
|
||||
});
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
expect(await embedder.embedBatch(["a", "b"])).toEqual([
|
||||
[0.1, 0.2],
|
||||
[0.3, 0.4],
|
||||
]);
|
||||
});
|
||||
|
||||
it("embedBatch() throws when the count mismatches the input", async () => {
|
||||
mockEmbeddingsCreate.mockResolvedValue({
|
||||
data: [{ index: 0, embedding: [0.1, 0.2] }],
|
||||
});
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
await expect(embedder.embedBatch(["a", "b"])).rejects.toThrow(
|
||||
/returned 1 embeddings for 2 texts/,
|
||||
);
|
||||
});
|
||||
|
||||
it("embed() throws when the endpoint returns no embeddings", async () => {
|
||||
mockEmbeddingsCreate.mockResolvedValue({ data: [] });
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
/returned no embeddings/,
|
||||
);
|
||||
});
|
||||
});
|
||||
});
|
||||
Reference in New Issue
Block a user