#!/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()