DoAn / core /rag /embedding_model.py
hungnha's picture
build server
4f9286e
Raw
History Blame Contribute Delete
2.94 kB
from __future__ import annotations
import os
import logging
import time
from dataclasses import dataclass
from typing import List, Sequence
import numpy as np
from openai import OpenAI
from langchain_core.embeddings import Embeddings
logger = logging.getLogger(__name__)
@dataclass
class EmbeddingConfig:
api_base_url: str = "https://api.siliconflow.com/v1"
model: str = "Qwen/Qwen3-Embedding-8B"
dimension: int = 4096
batch_size: int = 16
_embed_config: EmbeddingConfig | None = None
def get_embedding_config() -> EmbeddingConfig:
global _embed_config
if _embed_config is None:
_embed_config = EmbeddingConfig()
return _embed_config
class QwenEmbeddings(Embeddings):
def __init__(self, config: EmbeddingConfig | None = None):
self.config = config or get_embedding_config()
api_key = os.getenv("SILICONFLOW_API_KEY", "").strip()
if not api_key:
raise ValueError("Missing SILICONFLOW_API_KEY environment variable")
self._client = OpenAI(
api_key=api_key,
base_url=self.config.api_base_url,
)
logger.info(f"Initialized QwenEmbeddings: {self.config.model}")
def embed_query(self, text: str) -> List[float]:
return self._embed_texts([text])[0]
def embed_documents(self, texts: List[str]) -> List[List[float]]:
return self._embed_texts(texts)
def _embed_texts(self, texts: Sequence[str]) -> List[List[float]]:
if not texts:
return []
all_embeddings: List[List[float]] = []
batch_size = self.config.batch_size
max_retries = 3
# Process in batches
for i in range(0, len(texts), batch_size):
batch = list(texts[i:i + batch_size])
# Retry logic for rate limits
for attempt in range(max_retries):
try:
response = self._client.embeddings.create(
model=self.config.model,
input=batch,
)
for item in response.data:
all_embeddings.append(item.embedding)
break
except Exception as e:
# Rate limit -> wait and retry
if "rate" in str(e).lower() and attempt < max_retries - 1:
wait_time = 2 ** attempt
logger.warning(f"Rate limited, waiting {wait_time}s...")
time.sleep(wait_time)
else:
raise
return all_embeddings
def embed_texts_np(self, texts: Sequence[str]) -> np.ndarray:
return np.asarray(self._embed_texts(list(texts)), dtype=np.float32)
# Backward compatibility aliases
SiliconFlowConfig = EmbeddingConfig
get_config = get_embedding_config