@@ -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