Skip to content

Commit b9913a2

Browse files
committed
feat(example): auto-detect embedding model mode
1 parent bd0837e commit b9913a2

3 files changed

Lines changed: 72 additions & 8 deletions

File tree

‎examples/server/README.md‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -213,8 +213,8 @@ Most model runtime fields map to `llama_model_params` or `llama_context_params`
213213
| `threads` | Decode thread count. |
214214
| `threads_batch` | Prefill and batch thread count. |
215215
| `kv_unified` | Selects unified or per-sequence memory layout. |
216-
| `embedding` | Enables llama.cpp embedding extraction for `/v1/embeddings`. |
217-
| `pooling_type` | Selects pooled embedding behavior for embedding models, such as `1` for mean pooling. |
216+
| `embedding` | Overrides embedding mode; omit to auto-detect pooled embedding GGUFs from model metadata. |
217+
| `pooling_type` | Overrides pooled embedding behavior for embedding models, such as `1` for mean pooling. |
218218
| `store_logits` | Keeps logits after decode when needed by sampling or diagnostics. |
219219
| `use_mmap` | Memory maps model weights. |
220220
| `use_mlock` | Attempts to lock model pages into RAM. |

‎examples/server/configs/bge-small-en-v1.5.json‎

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,6 @@
99
"repo_id": "CompendiumLabs/bge-small-en-v1.5-gguf",
1010
"filename": "bge-small-en-v1.5-q4_k_m.gguf"
1111
},
12-
"embedding": true,
1312
"n_ctx": 512,
1413
"n_seq_max": 16,
1514
"n_batch": 512,

‎examples/server/server.py‎

Lines changed: 70 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -3335,7 +3335,7 @@ class ModelOptions(BaseModel):
33353335
rope_scaling_type: Optional[int] = None
33363336
pooling_type: Optional[int] = None
33373337
attention_type: Optional[int] = None
3338-
embedding: bool = False
3338+
embedding: Optional[bool] = None
33393339
rope_freq_base: Optional[float] = None
33403340
rope_freq_scale: Optional[float] = None
33413341
yarn_ext_factor: Optional[float] = None
@@ -10530,7 +10530,7 @@ def __init__(
1053010530
rope_scaling_type: Optional[int] = None,
1053110531
pooling_type: Optional[int] = None,
1053210532
attention_type: Optional[int] = None,
10533-
embedding: bool = False,
10533+
embedding: Optional[bool] = None,
1053410534
rope_freq_base: Optional[float] = None,
1053510535
rope_freq_scale: Optional[float] = None,
1053610536
yarn_ext_factor: Optional[float] = None,
@@ -10566,7 +10566,6 @@ def __init__(
1056610566
self.chat_template_override = chat_template
1056710567
self.response_schema = response_schema
1056810568
self.store_logits = store_logits
10569-
self.embedding = embedding
1057010569
self.max_output_tokens = max_output_tokens
1057110570
self.draft_model_max_batch_size = draft_model_max_batch_size
1057210571
self.draft_provider: Optional[DraftProvider] = None
@@ -10597,6 +10596,8 @@ def __init__(
1059710596
if vocab is None:
1059810597
raise RuntimeError("failed to access model vocabulary")
1059910598
self.vocab = vocab
10599+
embedding = self.resolve_embedding_mode(llama_model, embedding)
10600+
self.embedding = embedding
1060010601
self.has_encoder = bool(llama_cpp.llama_model_has_encoder(llama_model))
1060110602
self.has_decoder = bool(llama_cpp.llama_model_has_decoder(llama_model))
1060210603
if self.has_encoder and not embedding:
@@ -11074,14 +11075,34 @@ def close(self) -> None:
1107411075
llama_cpp.llama_backend_free()
1107511076
self.backend_initialized = False
1107611077

11077-
def _meta_value(self, key: str) -> Optional[str]:
11078+
@staticmethod
11079+
def _model_meta_key_by_index(llama_model: Any, index: int) -> Optional[str]:
11080+
capacity = 256
11081+
while True:
11082+
buffer = ctypes.create_string_buffer(capacity)
11083+
count = int(
11084+
llama_cpp.llama_model_meta_key_by_index(
11085+
llama_model,
11086+
index,
11087+
cast(Any, buffer),
11088+
capacity,
11089+
)
11090+
)
11091+
if count < 0:
11092+
return None
11093+
if count < capacity:
11094+
return buffer.value.decode("utf-8", errors="ignore")
11095+
capacity = count + 1
11096+
11097+
@staticmethod
11098+
def _model_meta_value(llama_model: Any, key: str) -> Optional[str]:
1107811099
encoded = key.encode("utf-8")
1107911100
capacity = 256
1108011101
while True:
1108111102
buffer = ctypes.create_string_buffer(capacity)
1108211103
count = int(
1108311104
llama_cpp.llama_model_meta_val_str(
11084-
self.llama_model,
11105+
llama_model,
1108511106
encoded,
1108611107
cast(Any, buffer),
1108711108
capacity,
@@ -11093,6 +11114,50 @@ def _meta_value(self, key: str) -> Optional[str]:
1109311114
return buffer.value.decode("utf-8", errors="ignore")
1109411115
capacity = count + 1
1109511116

11117+
@staticmethod
11118+
def _parse_pooling_type(value: str) -> Optional[int]:
11119+
normalized = value.strip().lower()
11120+
try:
11121+
return int(normalized)
11122+
except ValueError:
11123+
return {
11124+
"none": llama_cpp.LLAMA_POOLING_TYPE_NONE,
11125+
"mean": llama_cpp.LLAMA_POOLING_TYPE_MEAN,
11126+
"cls": llama_cpp.LLAMA_POOLING_TYPE_CLS,
11127+
"last": llama_cpp.LLAMA_POOLING_TYPE_LAST,
11128+
"rank": llama_cpp.LLAMA_POOLING_TYPE_RANK,
11129+
}.get(normalized)
11130+
11131+
@classmethod
11132+
def detect_embedding_model(cls, llama_model: Any) -> bool:
11133+
for index in range(int(llama_cpp.llama_model_meta_count(llama_model))):
11134+
key = cls._model_meta_key_by_index(llama_model, index)
11135+
if key is None or not key.endswith(".pooling_type"):
11136+
continue
11137+
value = cls._model_meta_value(llama_model, key)
11138+
if value is None:
11139+
continue
11140+
pooling_type = cls._parse_pooling_type(value)
11141+
return pooling_type in {
11142+
llama_cpp.LLAMA_POOLING_TYPE_MEAN,
11143+
llama_cpp.LLAMA_POOLING_TYPE_CLS,
11144+
llama_cpp.LLAMA_POOLING_TYPE_LAST,
11145+
}
11146+
return False
11147+
11148+
@classmethod
11149+
def resolve_embedding_mode(
11150+
cls,
11151+
llama_model: Any,
11152+
embedding: Optional[bool],
11153+
) -> bool:
11154+
if embedding is not None:
11155+
return embedding
11156+
return cls.detect_embedding_model(llama_model)
11157+
11158+
def _meta_value(self, key: str) -> Optional[str]:
11159+
return self._model_meta_value(self.llama_model, key)
11160+
1109611161
def _build_chat_formatter(self) -> Optional[Jinja2ChatFormatter]:
1109711162
template_text = self.chat_template_override
1109811163
if template_text is None:

0 commit comments

Comments
 (0)