获取文本的embedding向量。 支持多种主流模型,能自动适应不同库的调用方式。 - SentenceTransformer模型: e.g., 'all-MiniLM-L6-v2', 'Qwen/Qwen3-Embedding-0.6B' - FlagEmbedding模型: e.g., 'BAAI/bge-m3' :param text: 输入文本。 :param model_name: Hugging Face上的模型名称。 :param use_cache: 是否使用内存缓存。 :param kwargs: 传递给模型
(text, model_name="all-MiniLM-L6-v2", use_cache=True, **kwargs)
| 140 | return kwargs |
| 141 | |
| 142 | def get_embedding(text, model_name="all-MiniLM-L6-v2", use_cache=True, **kwargs): |
| 143 | """ |
| 144 | 获取文本的embedding向量。 |
| 145 | 支持多种主流模型,能自动适应不同库的调用方式。 |
| 146 | - SentenceTransformer模型: e.g., 'all-MiniLM-L6-v2', 'Qwen/Qwen3-Embedding-0.6B' |
| 147 | - FlagEmbedding模型: e.g., 'BAAI/bge-m3' |
| 148 | |
| 149 | :param text: 输入文本。 |
| 150 | :param model_name: Hugging Face上的模型名称。 |
| 151 | :param use_cache: 是否使用内存缓存。 |
| 152 | :param kwargs: 传递给模型构造函数或encode方法的额外参数。 |
| 153 | - for Qwen: `model_kwargs`, `tokenizer_kwargs`, `prompt_name="query"` |
| 154 | - for BGE-M3: `use_fp16=True`, `max_length=8192` |
| 155 | :return: 文本的embedding向量 (numpy array)。 |
| 156 | """ |
| 157 | model_config_key = json.dumps({"model_name": model_name, **kwargs}, sort_keys=True) |
| 158 | |
| 159 | if use_cache: |
| 160 | cache_key = f"{model_config_key}::{hash(text)}" |
| 161 | if cache_key in _embedding_cache: |
| 162 | return _embedding_cache[cache_key] |
| 163 | |
| 164 | # --- Model Loading --- |
| 165 | model_init_key = json.dumps({"model_name": model_name, **{k:v for k,v in kwargs.items() if k not in ['batch_size', 'max_length']}}, sort_keys=True) |
| 166 | if model_init_key not in _model_cache: |
| 167 | print(f"Loading model: {model_name}...") |
| 168 | if 'bge-m3' in model_name.lower(): |
| 169 | try: |
| 170 | from FlagEmbedding import BGEM3FlagModel |
| 171 | init_kwargs = _get_valid_kwargs(BGEM3FlagModel.__init__, kwargs) |
| 172 | print(f"-> Using BGEM3FlagModel with init kwargs: {init_kwargs}") |
| 173 | _model_cache[model_init_key] = BGEM3FlagModel(model_name, **init_kwargs) |
| 174 | except ImportError: |
| 175 | raise ImportError("Please install FlagEmbedding: 'pip install -U FlagEmbedding' to use bge-m3 model.") |
| 176 | else: # Default handler for SentenceTransformer-based models (like Qwen, all-MiniLM, etc.) |
| 177 | try: |
| 178 | from sentence_transformers import SentenceTransformer |
| 179 | init_kwargs = _get_valid_kwargs(SentenceTransformer.__init__, kwargs) |
| 180 | print(f"-> Using SentenceTransformer with init kwargs: {init_kwargs}") |
| 181 | _model_cache[model_init_key] = SentenceTransformer(model_name, **init_kwargs) |
| 182 | except ImportError: |
| 183 | raise ImportError("Please install sentence-transformers: 'pip install -U sentence-transformers' to use this model.") |
| 184 | |
| 185 | model = _model_cache[model_init_key] |
| 186 | |
| 187 | # --- Encoding --- |
| 188 | embedding = None |
| 189 | if 'bge-m3' in model_name.lower(): |
| 190 | encode_kwargs = _get_valid_kwargs(model.encode, kwargs) |
| 191 | print(f"-> Encoding with BGEM3FlagModel using kwargs: {encode_kwargs}") |
| 192 | result = model.encode([text], **encode_kwargs) |
| 193 | embedding = result['dense_vecs'][0] |
| 194 | else: # Default to SentenceTransformer-based models |
| 195 | encode_kwargs = _get_valid_kwargs(model.encode, kwargs) |
| 196 | print(f"-> Encoding with SentenceTransformer using kwargs: {encode_kwargs}") |
| 197 | embedding = model.encode([text], **encode_kwargs)[0] |
| 198 | |
| 199 | if use_cache: |
no test coverage detected