qwen.py 1.5 KB

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