feat: 导入 MemRelay 初始源码
This commit is contained in:
@@ -0,0 +1,32 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user