33 lines
975 B
Python
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()
|