qwen.py 1.4 KB

123456789101112131415161718192021222324252627282930
  1. from __future__ import annotations
  2. import hashlib
  3. from typing import Any
  4. import requests
  5. class QwenEmbeddingError(RuntimeError):
  6. pass
  7. class QwenEmbeddingClient:
  8. def __init__(self, *, api_key: str, base_url: str, model: str, dimension: int = 1024, timeout: int = 30):
  9. if not api_key or not base_url or not model:
  10. raise QwenEmbeddingError("Qwen embedding configuration is incomplete")
  11. self.api_key, self.base_url, self.model = api_key, base_url.rstrip("/"), model
  12. self.dimension, self.timeout = dimension, timeout
  13. def cache_key(self, text: str) -> str:
  14. return hashlib.sha256(f"{self.model}:{self.dimension}:{text}".encode()).hexdigest()
  15. def embed(self, texts: list[str]) -> list[list[float]]:
  16. try:
  17. response = requests.post(f"{self.base_url}/embeddings", headers={"Authorization": "Bearer " + self.api_key}, json={"model": self.model, "input": texts, "dimensions": self.dimension}, timeout=self.timeout)
  18. response.raise_for_status()
  19. vectors = [item["embedding"] for item in sorted(response.json()["data"], key=lambda item: item["index"])]
  20. except requests.RequestException as exc:
  21. raise QwenEmbeddingError("Qwen embedding request failed") from exc
  22. if len(vectors) != len(texts) or any(len(vector) != self.dimension for vector in vectors):
  23. raise QwenEmbeddingError("Qwen embedding dimension mismatch")
  24. return vectors