Files
MemRelay/scripts/preload_basic_memory_model.py
T

33 lines
975 B
Python

"""Download and validate the local embedding model used by Basic Memory."""
from __future__ import annotations
import argparse
from fastembed import TextEmbedding
DEFAULT_MODEL = "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--model", default=DEFAULT_MODEL)
parser.add_argument("--cache-dir", default="/data/basic-memory/cache")
parser.add_argument("--threads", type=int, default=2)
args = parser.parse_args()
embedding = TextEmbedding(
model_name=args.model,
cache_dir=args.cache_dir,
threads=args.threads,
)
vectors = list(embedding.embed(["记忆服务", "memory service"]))
if len(vectors) != 2 or len(vectors[0]) != 384:
raise RuntimeError("Unexpected embedding model output")
print(f"model={args.model} vectors={len(vectors)} dimensions={len(vectors[0])}")
if __name__ == "__main__":
main()