CoolFace
Modelpublic

KaLM-Embedding/KaLM-Reranker-V1-Nano-R2

sourceHugging Faceapache-2.0updated 3d agoView on Hugging Face
2likes92downloads
kalm_cross_encoder.py299 linesDownload Raw Back to root
1from __future__ import annotations2 3from typing import Any, ClassVar4 5try:6    from typing import Self7except ImportError:8    from typing_extensions import Self9 10import torch11from sentence_transformers.base.modules import InputModule12from transformers import AutoConfig, AutoModelForSeq2SeqLM, AutoTokenizer13 14from .kalm_reranker_utils import (15    DEFAULT_INSTRUCTION,16    DEFAULT_SYSTEM_INSTRUCTION,17    answer_token_id,18    build_decoder_text,19    cast_floating_parameters,20    extract_yes_no_logits,21    forward_reranker_model,22    normalize_requested_dtype,23    validate_text_pairs,24)25 26 27class KaLMCrossEncoderModule(InputModule):28    """Sentence Transformers input module for KaLM encoder-decoder rerankers."""29 30    config_file_name = "kalm_cross_encoder_config.json"31    config_keys: ClassVar[list[str]] = [32        "query_max_length",33        "document_max_length",34        "encoder_chunk_size",35        "system_instruction",36    ]37    save_in_root = True38 39    def __init__(40        self,41        model_name_or_path: str,42        *,43        query_max_length: int = 512,44        document_max_length: int = 1024,45        encoder_chunk_size: int | None = 4,46        system_instruction: str = DEFAULT_SYSTEM_INSTRUCTION,47        model_kwargs: dict[str, Any] | None = None,48        processor_kwargs: dict[str, Any] | None = None,49        config_kwargs: dict[str, Any] | None = None,50        backend: str = "torch",51    ) -> None:52        super().__init__()53        if backend != "torch":54            raise ValueError(55                "KaLMCrossEncoderModule only supports backend='torch'; "56                f"received {backend!r}."57            )58        if not isinstance(model_name_or_path, str) or not model_name_or_path:59            raise ValueError("model_name_or_path must be a non-empty string.")60        if not isinstance(query_max_length, int) or query_max_length <= 0:61            raise ValueError("query_max_length must be a positive integer.")62        if not isinstance(document_max_length, int) or document_max_length <= 0:63            raise ValueError("document_max_length must be a positive integer.")64        if encoder_chunk_size is not None and (65            not isinstance(encoder_chunk_size, int) or encoder_chunk_size <= 066        ):67            raise ValueError("encoder_chunk_size must be a positive integer or None.")68        if not isinstance(system_instruction, str):69            raise TypeError("system_instruction must be a string.")70 71        self.query_max_length = query_max_length72        self.max_seq_length = document_max_length73        self.encoder_chunk_size = encoder_chunk_size74        self.system_instruction = system_instruction75        self.backend = backend76 77        model_kwargs = dict(model_kwargs or {})78        processor_kwargs = dict(processor_kwargs or {})79        config_kwargs = dict(config_kwargs or {})80 81        num_labels = config_kwargs.pop("num_labels", 1)82        if num_labels != 1:83            raise ValueError(84                "KaLM reranking produces one relevance score; num_labels must be 1."85            )86 87        config = AutoConfig.from_pretrained(model_name_or_path, **config_kwargs)88        self.tokenizer = AutoTokenizer.from_pretrained(89            model_name_or_path, **processor_kwargs90        )91        if self.tokenizer.pad_token_id is None:92            if self.tokenizer.eos_token_id is None:93                raise ValueError(94                    "The tokenizer must define a pad token or an EOS token."95                )96            self.tokenizer.pad_token = self.tokenizer.eos_token97        self.tokenizer.padding_side = "right"98        self.processor = self.tokenizer99 100        requested_dtype = normalize_requested_dtype(101            model_kwargs.get("dtype", model_kwargs.get("torch_dtype"))102        )103        self.model = AutoModelForSeq2SeqLM.from_pretrained(104            model_name_or_path,105            config=config,106            **model_kwargs,107        )108        cast_floating_parameters(self.model, requested_dtype)109 110        self.yes_token_id = answer_token_id(self.tokenizer, "yes")111        self.no_token_id = answer_token_id(self.tokenizer, "no")112 113    @property114    def document_max_length(self) -> int:115        return self.max_seq_length116 117    @document_max_length.setter118    def document_max_length(self, value: int) -> None:119        if not isinstance(value, int) or value <= 0:120            raise ValueError("document_max_length must be a positive integer.")121        self.max_seq_length = value122 123    @property124    def encoder_chunk_size(self) -> int | None:125        return self._encoder_chunk_size126 127    @encoder_chunk_size.setter128    def encoder_chunk_size(self, value: int | None) -> None:129        if value is not None and (not isinstance(value, int) or value <= 0):130            raise ValueError("encoder_chunk_size must be a positive integer or None.")131        self._encoder_chunk_size = value132 133    @property134    def chunk_size(self) -> int | None:135        """Alias for the encoder token mean-pooling compression rate."""136        return self.encoder_chunk_size137 138    @chunk_size.setter139    def chunk_size(self, value: int | None) -> None:140        self.encoder_chunk_size = value141 142    def preprocess(143        self,144        inputs: list[Any],145        prompt: str | None = None,146        **kwargs: Any,147    ) -> dict[str, torch.Tensor]:148        pairs = validate_text_pairs(inputs)149        if not pairs:150            return {}151 152        instruction = DEFAULT_INSTRUCTION if prompt is None else prompt153        if not isinstance(instruction, str):154            raise TypeError("prompt must be a string or None.")155 156        encoder_texts = [f"<Document>: {document}" for _, document in pairs]157        decoder_texts = [158            build_decoder_text(159                self.tokenizer,160                query,161                instruction,162                self.system_instruction,163                self.query_max_length,164            )165            for query, _ in pairs166        ]167 168        encoder_batch = self.tokenizer(169            encoder_texts,170            padding=True,171            truncation=True,172            max_length=self.document_max_length,173            add_special_tokens=False,174            return_tensors="pt",175        )176        decoder_batch = self.tokenizer(177            decoder_texts,178            padding=True,179            pad_to_multiple_of=8,180            add_special_tokens=False,181            return_tensors="pt",182        )183        return {184            "input_ids": encoder_batch["input_ids"],185            "attention_mask": encoder_batch["attention_mask"],186            "decoder_input_ids": decoder_batch["input_ids"],187            "decoder_attention_mask": decoder_batch["attention_mask"],188        }189 190    def forward(191        self,192        features: dict[str, torch.Tensor | Any],193        **kwargs: Any,194    ) -> dict[str, torch.Tensor | Any]:195        outputs = forward_reranker_model(196            self.model,197            input_ids=features["input_ids"],198            attention_mask=features["attention_mask"],199            decoder_input_ids=features["decoder_input_ids"],200            decoder_attention_mask=features["decoder_attention_mask"],201            encoder_chunk_size=self.chunk_size,202        )203        yes_no_logits = extract_yes_no_logits(204            outputs.logits,205            features["decoder_attention_mask"],206            self.yes_token_id,207            self.no_token_id,208        )209        features["scores"] = (yes_no_logits[:, 0] - yes_no_logits[:, 1]).unsqueeze(1)210        return features211 212    def save(213        self,214        output_path: str,215        *args: Any,216        safe_serialization: bool = True,217        **kwargs: Any,218    ) -> None:219        self.model.save_pretrained(output_path, safe_serialization=safe_serialization)220        self.tokenizer.save_pretrained(output_path)221        self.save_config(output_path)222 223    @classmethod224    def load(225        cls,226        model_name_or_path: str,227        subfolder: str = "",228        token: bool | str | None = None,229        cache_folder: str | None = None,230        revision: str | None = None,231        local_files_only: bool = False,232        trust_remote_code: bool = False,233        model_kwargs: dict[str, Any] | None = None,234        processor_kwargs: dict[str, Any] | None = None,235        config_kwargs: dict[str, Any] | None = None,236        backend: str = "torch",237        **kwargs: Any,238    ) -> Self:239        module_config = cls.load_config(240            model_name_or_path,241            subfolder=subfolder,242            token=token,243            cache_folder=cache_folder,244            revision=revision,245            local_files_only=local_files_only,246        )247 248        supplied_model_kwargs = dict(model_kwargs or {})249        supplied_config_kwargs = dict(config_kwargs or {})250        supplied_module_kwargs = dict(kwargs)251        chunk_size_values: list[tuple[str, int | None]] = []252        for source_name, source in (253            ("model_kwargs", supplied_model_kwargs),254            ("config_kwargs", supplied_config_kwargs),255            ("module kwargs", supplied_module_kwargs),256        ):257            for key in ("chunk_size", "encoder_chunk_size"):258                if key in source:259                    chunk_size_values.append((f"{source_name}.{key}", source.pop(key)))260        if chunk_size_values:261            first_name, first_value = chunk_size_values[0]262            for current_name, current_value in chunk_size_values[1:]:263                if current_value != first_value:264                    raise ValueError(265                        "Conflicting encoder chunk sizes: "266                        f"{first_name}={first_value!r}, "267                        f"{current_name}={current_value!r}."268                    )269            module_config["encoder_chunk_size"] = first_value270 271        hub_kwargs = {272            "subfolder": subfolder,273            "token": token,274            "cache_dir": cache_folder,275            "revision": revision,276            "local_files_only": local_files_only,277            "trust_remote_code": trust_remote_code,278        }279        effective_model_kwargs = {**hub_kwargs, **supplied_model_kwargs}280        effective_processor_kwargs = {**hub_kwargs, **(processor_kwargs or {})}281        effective_config_kwargs = {**hub_kwargs, **supplied_config_kwargs}282 283        if "model_max_length" in effective_processor_kwargs:284            module_config["document_max_length"] = effective_processor_kwargs[285                "model_max_length"286            ]287 288        return cls(289            model_name_or_path,290            model_kwargs=effective_model_kwargs,291            processor_kwargs=effective_processor_kwargs,292            config_kwargs=effective_config_kwargs,293            backend=backend,294            **module_config,295        )296 297 298__all__ = ["KaLMCrossEncoderModule"]299