diff --git a/conversion/base.py b/conversion/base.py index 5561481e77..221aa8093b 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -2326,6 +2326,12 @@ class TextModel(ModelBase): raise NotImplementedError("Only MEAN, CLS, and LAST pooling types supported") self.gguf_writer.add_pooling_type(pooling_type) + # pooling before a classification head (e.g. ModernBertForSequenceClassification) + if (classifier_pooling := self.hparams.get("classifier_pooling")) is not None: + if classifier_pooling not in ("cls", "mean"): + raise NotImplementedError(f"Unsupported classifier_pooling: {classifier_pooling}") + self.gguf_writer.add_classifier_pooling_type(mode_mapping[classifier_pooling]) + def _set_vocab_glmedge(self): from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained(self.dir_model) diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index 9a7a5e5bfa..e193e99957 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -313,6 +313,7 @@ class Keys: class Classifier: OUTPUT_LABELS = "{arch}.classifier.output_labels" + POOLING_TYPE = "{arch}.classifier.pooling_type" class ShortConv: L_CACHE = "{arch}.shortconv.l_cache" diff --git a/gguf-py/gguf/gguf_writer.py b/gguf-py/gguf/gguf_writer.py index cf7b367e52..543d9b744c 100644 --- a/gguf-py/gguf/gguf_writer.py +++ b/gguf-py/gguf/gguf_writer.py @@ -1331,6 +1331,9 @@ class GGUFWriter: def add_classifier_output_labels(self, labels: Sequence[str]) -> None: self.add_array(Keys.Classifier.OUTPUT_LABELS.format(arch=self.arch), labels) + def add_classifier_pooling_type(self, value: PoolingType) -> None: + self.add_uint32(Keys.Classifier.POOLING_TYPE.format(arch=self.arch), value.value) + # for vision models def add_clip_has_vision_encoder(self, value: bool) -> None: diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index 8f1e239dae..2c5a228f35 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -361,6 +361,7 @@ static const std::map LLM_KV_NAMES = { { LLM_KV_CONVNEXT_BLOCK_COUNT, "%s.convnext.block_count" }, { LLM_KV_CLASSIFIER_OUTPUT_LABELS, "%s.classifier.output_labels" }, + { LLM_KV_CLASSIFIER_POOLING_TYPE, "%s.classifier.pooling_type" }, { LLM_KV_TARGET_LAYERS, "%s.target_layers" }, { LLM_KV_TARGET_HIDDEN_SIZE, "%s.target_hidden_size" }, diff --git a/src/llama-arch.h b/src/llama-arch.h index 23b6b38100..3b8bd64287 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -407,6 +407,7 @@ enum llm_kv { LLM_KV_CONVNEXT_BLOCK_COUNT, LLM_KV_CLASSIFIER_OUTPUT_LABELS, + LLM_KV_CLASSIFIER_POOLING_TYPE, LLM_KV_TARGET_LAYERS, LLM_KV_TARGET_HIDDEN_SIZE, diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index a806126efa..abf3069bbc 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -3727,8 +3727,8 @@ void llm_graph_context::build_pooling( } break; case LLAMA_POOLING_TYPE_RANK: { - if (arch == LLM_ARCH_MODERN_BERT) { - // modern bert gte reranker builds mean first then applies prediction head and classifier + if (hparams.pooling_type_cls == LLAMA_POOLING_TYPE_MEAN) { + // modern bert with classifier_pooling = "mean" builds mean first then applies prediction head and classifier // https://github.com/huggingface/transformers/blob/main/src/transformers/models/modernbert/modular_modernbert.py#L1404-1411 ggml_tensor * inp_mean = build_inp_mean(); cur = ggml_mul_mat(ctx0, ggml_cont(ctx0, ggml_transpose(ctx0, inp)), inp_mean); diff --git a/src/llama-hparams.h b/src/llama-hparams.h index 73dffcc9f7..d8fbfbdd8a 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -350,6 +350,7 @@ struct llama_hparams { uint32_t dec_n_layer = 0; enum llama_pooling_type pooling_type = LLAMA_POOLING_TYPE_NONE; + enum llama_pooling_type pooling_type_cls = LLAMA_POOLING_TYPE_UNSPECIFIED; // pooling before the classifier head (RANK) enum llama_rope_type rope_type = LLAMA_ROPE_TYPE_NONE; enum llama_rope_scaling_type rope_scaling_type_train = LLAMA_ROPE_SCALING_TYPE_NONE; diff --git a/src/llama-model.cpp b/src/llama-model.cpp index ab5e744b5d..151a3a2a89 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -1320,6 +1320,7 @@ void llama_model_base::load_hparams(llama_model_loader & ml) { ml.get_key(LLM_KV_EMBEDDING_LENGTH_OUT, hparams.n_embd_out_impl, false); ml.get_key(LLM_KV_ATTENTION_CAUSAL, hparams.causal_attn, false); ml.get_key(LLM_KV_POOLING_TYPE, hparams.pooling_type, false); + ml.get_key(LLM_KV_CLASSIFIER_POOLING_TYPE, hparams.pooling_type_cls, false); ml.get_key(LLM_KV_BLOCK_COUNT, hparams.n_layer_all); GGML_ASSERT(hparams.n_layer_all > 0 && hparams.n_layer_all <= LLAMA_MAX_LAYERS); ml.get_key(LLM_KV_NEXTN_PREDICT_LAYERS, hparams.n_layer_nextn, false); diff --git a/src/models/modern-bert.cpp b/src/models/modern-bert.cpp index b7542d59bd..158e3160c5 100644 --- a/src/models/modern-bert.cpp +++ b/src/models/modern-bert.cpp @@ -20,6 +20,11 @@ void llama_model_modern_bert::load_arch_hparams(llama_model_loader & ml) { hparams.llm_ffn_op = llm_ffn_op_type_from_string(hidden_act, LLM_FFN_GEGLU); } + // GGUFs without a classifier pooling type use mean (gte-reranker-modernbert-base) + if (hparams.pooling_type_cls == LLAMA_POOLING_TYPE_UNSPECIFIED) { + hparams.pooling_type_cls = LLAMA_POOLING_TYPE_MEAN; + } + switch (hparams.n_layer()) { case 12: type = LLM_TYPE_47M; break; // granite-embedding-small