from sentence_transformers.base.modules import InputModule from transformers import AutoModel, AutoTokenizer class SpladeSTModule(InputModule): save_in_root = True def __init__(self, model_name_or_path: str, **kwargs): super().__init__() self.model = AutoModel.from_pretrained(model_name_or_path, trust_remote_code=True) self.tokenizer = AutoTokenizer.from_pretrained(model_name_or_path) # for SparseEncoder.decode def preprocess(self, inputs, prompt=None, **kwargs): prefix = prompt or "" cfg = self.model.config max_length = cfg.query_max_length if prefix == cfg.query_prefix else cfg.doc_max_length ids, attn, pool, _ = self.model._tokenize(list(inputs), prefix, max_length) return {"input_ids": ids, "attention_mask": attn, "pooling_mask": pool} def forward(self, features, **kwargs): features["sentence_embedding"] = self.model( features["input_ids"], features["attention_mask"], features["pooling_mask"] ) return features def get_embedding_dimension(self): return self.model.config.vocab_size @classmethod def load(cls, model_name_or_path, **kwargs): return cls(model_name_or_path) def save(self, output_path, **kwargs): pass # repo is assembled by export.py, never by ST save