#!/usr/bin/env python
"""Embed a user query with NVIDIA NIM embeddings."""
from __future__ import annotations
import argparse
import json
import os
import numpy as np
def load_client(model_name: str, api_key: str, truncate: str):
try:
from langchain_nvidia_ai_endpoints import NVIDIAEmbeddings
except Exception as exc: # pragma: no cover - import guard for missing deps
raise SystemExit(
"Failed to import langchain_nvidia_ai_endpoints. "
"Verify the env and run `python -c \"from langchain_nvidia_ai_endpoints import NVIDIAEmbeddings\"`.\n"
f"Import error: {exc}"
) from exc
return NVIDIAEmbeddings(model=model_name, api_key=api_key, truncate=truncate)
def normalize_embedding(embedding: np.ndarray) -> np.ndarray:
norm = np.linalg.norm(embedding)
if norm <= 1e-12:
return embedding
return embedding / norm
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Embed a query with NVIDIA NIM embeddings.")
parser.add_argument("--query", required=True, help="User query text.")
parser.add_argument(
"--output",
default=os.path.join("output", "query_embedding.json"),
help="Output path for the query embedding JSON.",
)
parser.add_argument(
"--model",
default="nvidia/llama-3.2-nv-embedqa-1b-v2",
help="NVIDIA embedding model name.",
)
parser.add_argument(
"--api-key",
default=os.getenv("NVIDIA_API_KEY"),
help="NVIDIA API key (or set NVIDIA_API_KEY).",
)
parser.add_argument(
"--truncate",
default="NONE",
help="Truncation mode for the NVIDIA endpoint (e.g., NONE, START, END).",
)
parser.add_argument(
"--normalize",
action="store_true",
help="Normalize embeddings for cosine similarity.",
)
return parser.parse_args()
def main() -> None:
args = parse_args()
if not args.api_key:
raise SystemExit(
"Missing NVIDIA API key. Set NVIDIA_API_KEY or pass --api-key."
)
os.makedirs(os.path.dirname(args.output), exist_ok=True)
client = load_client(args.model, args.api_key, args.truncate)
embedding = client.embed_query(args.query)
embedding = np.asarray(embedding, dtype=float)
if args.normalize:
embedding = normalize_embedding(embedding)
embedding = embedding.tolist()
payload = {
"query": args.query,
"embedding": [float(x) for x in embedding],
"embedding_model": args.model,
"normalize": args.normalize,
"truncate": args.truncate,
}
with open(args.output, "w", encoding="utf-8") as handle:
json.dump(payload, handle, indent=2, ensure_ascii=True)
print(f"Wrote query embedding to {args.output}")
if __name__ == "__main__":
main()