sentence-transformers CrossEncoder 自定义模型开发指南:模块链、保存加载机制与自定义模块实现
2026/9/20 1:51:49 网站建设 项目流程

sentence-transformers CrossEncoder 自定义模型开发指南:模块链、保存加载机制与自定义模块实现

【免费下载链接】sentence-transformersState-of-the-Art Embeddings, Retrieval, and Reranking项目地址: https://gitcode.com/gh_mirrors/se/sentence-transformers

本文基于 sentence-transformers 官方文档 docs/cross_encoder/usage/custom_models.rst 展开,系统讲解如何以“模块化”方式构建、保存、加载 CrossEncoder 模型:包括三种标准模块链模式(序列分类、因果语言模型 + LogitScore、特征提取 + 池化 + Dense)、modules.json等保存产物的结构解析,以及如何继承Module/InputModule基类开发自己的打分模块并实现关键词参数透传。读完后你将能够从零拼装任意 CrossEncoder 模型、理解其底层序列化机制,并具备编写可共享自定义模块的工程能力。

模块化架构:由顺序模块组成的模型

CrossEncoder 与 SentenceTransformer 一样,由一系列顺序执行的模块(module)组成。CrossEncoder 类继承自BaseModel,而BaseModel又继承自torch.nn.Sequential,因此每个 CrossEncoder 本质上就是一条“特征字典(features dict)逐模块流转”的管道。最简单的查看方式是直接打印模型:

from sentence_transformers import CrossEncoder model = CrossEncoder("Qwen/Qwen3-Reranker-0.6B") print(model) """ CrossEncoder( (0): Transformer({'transformer_task': 'text-generation', 'modality_config': {'text': {'method': 'forward', 'method_output_name': 'logits'}, 'message': {'method': 'forward', 'method_output_name': 'logits', 'format': 'flat'}}, 'module_output_name': 'causal_logits', 'architecture': 'Qwen3ForCausalLM'}) (1): LogitScore({'true_token_id': 9693, 'false_token_id': 2152, 'module_input_name': 'causal_logits'}) ) """

可以看到这个生成式重排模型只有两个模块:Transformer(因果语言模型)和LogitScore(从 logits 提取相关性分数)。

常见模块链模式

编码器式:Sequence Classification

对于传统编码器架构的 CrossEncoder(如 BERT、RoBERTa),只需单个模块即可:

  • Transformertransformer_task="sequence-classification":通过transformers.AutoModelForSequenceClassification加载模型,直接返回分类得分。
from sentence_transformers import CrossEncoder model = CrossEncoder("cross-encoder/ms-marco-MiniLM-L6-v2") print(model) """ CrossEncoder( (0): Transformer({'transformer_task': 'sequence-classification', 'modality_config': {'text': {'method': 'forward', 'method_output_name': 'logits'}}, 'module_output_name': 'scores', 'architecture': 'BertForSequenceClassification'}) ) """

在源码中,这一任务映射关系定义于 TRANSFORMER_TASK_TO_AUTO_MODEL:"feature-extraction"对应AutoModel"sequence-classification"对应AutoModelForSequenceClassification"text-generation"对应AutoModelForCausalLM,即不同的transformer_task决定了底层加载哪个 Auto 类。

因果语言模型式:Text Generation + LogitScore

对于生成式重排模型(如 Qwen、Llama),模块链通常由两部分构成:

  • Transformertransformer_task="text-generation":通过transformers.AutoModelForCausalLM加载模型,输出语言模型头(LM head)上的原始 logits;
  • LogitScore:提取最后一个 token 位置的 logits 并计算得分。若同时设置true_token_idfalse_token_id,得分为对数几率(log-odds):logit[true] - logit[false];若只设置true_token_id,得分即该 token 的原始 logit。

这一逻辑可以直接在 LogitScore 源码 中得到印证:

def forward(self, features: dict[str, torch.Tensor], **kwargs) -> dict[str, torch.Tensor]: # Left padding is enforced by Transformer, so the last position is always a real token. logits = features[self.module_input_name][:, -1] if self.false_token_id is None: scores = logits[:, self.true_token_id] else: scores = logits[:, self.true_token_id] - logits[:, self.false_token_id] features["scores"] = scores.unsqueeze(1) return features

值得注意的细节是:Transformer模块强制左填充(left padding),因此[:, -1]永远指向序列中真实的最后一个 token,这保证了“在最后一个位置读取 logits”策略的正确性。

上述模块链的自动选择在CrossEncoder._load_default_modules中完成:当模型架构名以ForCausalLM结尾时,使用transformer_task="text-generation"加载Transformer,并追加一个LogitScore模块,其true_token_id/false_token_id分别取自 tokenizer 中的"yes""no"token(源码中对这两个 token 缺失的情况会抛出明确的ValueError);否则回退到transformer_task="sequence-classification"的传统编码器方式。

特征提取 + Pooling + Dense

相比 Text Generation + LogitScore 方案,一种更省显存的替代做法是只用特征提取、不带 LM 头:

  • Transformertransformer_task="feature-extraction":仅通过transformers.AutoModel加载基座模型(无 LM 头),输出 hidden states;
  • Poolingpooling_mode="lasttoken":提取最后一个 token 的 hidden state;
  • Dense:将 hidden state 投影为单一得分。
from sentence_transformers import CrossEncoder from sentence_transformers.cross_encoder.modules import Transformer, Dense from sentence_transformers.sentence_transformer.modules import Pooling transformer = Transformer("Qwen/Qwen3.5-0.8B", transformer_task="feature-extraction") pooling = Pooling(transformer.get_embedding_dimension(), pooling_mode="lasttoken") # Initialize Dense weights to approximate LogitScore: weight = embed("1") - embed("0") true_id = transformer.tokenizer.convert_tokens_to_ids("1") false_id = transformer.tokenizer.convert_tokens_to_ids("0") embeddings = transformer.model.get_input_embeddings().weight.data init_weight = (embeddings[true_id] - embeddings[false_id]).unsqueeze(0) dense = Dense( in_features=transformer.get_embedding_dimension(), out_features=1, bias=False, activation_function=None, init_weight=init_weight, module_output_name="scores", ) model = CrossEncoder(modules=[transformer, pooling, dense])

这里的关键技巧是Dense层的权重初始化:由于大多数因果语言模型将输入嵌入与 LM 头权重绑定(weight tying),将Dense权重初始化为embed("1") - embed("0")等价于对"1""0"两个 token 计算 log-odds 的起点。这样就绕过了对整个词表做 LM 头矩阵乘的高开销计算。从 Dense 源码 可以看到,init_weight参数会直接赋给线性层的weightself.linear.weight = nn.Parameter(init_weight)),且module_output_name决定了输出写入 features 字典的键名——这里必须设为"scores"才能被模型最终识别为得分。

提示:完整对比 Text Generation + LogitScore 与 Feature Extraction + Pooling + Dense 两种方案,可参考仓库中的多模态训练示例 examples/cross_encoder/training/multimodal。

手动构建自定义模块链

你可以手动指定模块列表来构建 CrossEncoder。例如,改用"1""0"作为 true/false token 而非常规的"yes"/"no"

from sentence_transformers import CrossEncoder from sentence_transformers.cross_encoder.modules import Transformer, LogitScore transformer = Transformer("Qwen/Qwen3-Reranker-0.6B", transformer_task="text-generation") # Look up the token IDs for "1" and "0" in the tokenizer true_id = transformer.tokenizer.convert_tokens_to_ids("1") false_id = transformer.tokenizer.convert_tokens_to_ids("0") model = CrossEncoder( modules=[transformer, LogitScore(true_token_id=true_id, false_token_id=false_id)] )

所有可用于构建 CrossEncoder 的模块均可从 sentence_transformers.cross_encoder.modules 统一导入(该包重导出了DenseInputModuleModuleAsymRouterTransformerLogitScore),因此无论来自base.modules还是cross_encoder.modules,都可以通过同一条导入路径拿到。

保存 CrossEncoder 模型

调用CrossEncoder.save_pretrained时会生成三类文件:

  • modules.json:模块名称、路径与类型列表,用于重建模型;
  • config_sentence_transformers.json:模型级配置,包括模型类型、已保存的 prompts、默认 prompt 名称、激活函数等;
  • 各模块的专属文件:每个模块保存在以“模块索引_类名”命名的子目录中(如1_LogitScore),但第一个模块若设置save_in_root = True则直接保存在根目录(Transformer模块即如此,见 InputModule 的save_in_root定义)。

以保存一个 CausalLM 式 CrossEncoder 为例,生成的目录结构如下:

my-cross-encoder/ ├── 1_LogitScore │ └── config.json ├── README.md ├── chat_template.jinja ├── config.json ├── config_sentence_transformers.json ├── generation_config.json ├── model.safetensors ├── modules.json ├── sentence_bert_config.json ├── tokenizer.json └── tokenizer_config.json

其中modules.json记录每个模块的元信息:

[ { "idx": 0, "name": "0", "path": "", "type": "sentence_transformers.base.modules.transformer.Transformer" }, { "idx": 1, "name": "1", "path": "1_LogitScore", "type": "sentence_transformers.cross_encoder.modules.logit_score.LogitScore" } ]

config_sentence_transformers.json保存模型级配置:

{ "__version__": { "sentence_transformers": "5.4.0", "transformers": "5.5.0", "pytorch": "2.10.0" }, "activation_fn": "torch.nn.modules.linear.Identity", "default_prompt_name": "query", "model_type": "CrossEncoder", "prompts": { "query": "Given a web search query, retrieve relevant passages that answer the query" } }

1_LogitScore/config.json保存 LogitScore 模块自身的配置:

{ "true_token_id": 9693, "false_token_id": 2152, "module_input_name": "causal_logits" }

根目录下的sentence_bert_config.json保存的是Transformer模块的配置:

{ "transformer_task": "text-generation", "modality_config": { "text": { "method": "forward", "method_output_name": "logits" }, "message": { "method": "forward", "method_output_name": "logits", "format": "flat" } }, "module_output_name": "causal_logits" }

modality_config"message"模态下的"format"键控制聊天模板(chat template)输入的组织方式:

  • "structured":content 是类型化字典列表,如[{"type": "text", "text": "hello"}]
  • "flat":content 就是直接的文本值,如"hello"

该格式由模型聊天模板自动推断:多数视觉语言模型使用 structured 格式,纯文本因果 LM 通常使用 flat 格式。

加载 CrossEncoder 模型

加载时,系统读取modules.json确定模型由哪些模块构成:每个模块类由其type字段解析,再根据对应模块目录中的配置完成初始化。例如LogitScore就是读取1_LogitScore/config.json并把其中的值作为关键字参数传给LogitScore(...)

如果模型目录中没有modules.json(例如加载一个纯transformers格式的模型),_load_default_modules会按模型架构自动决定模块链:ForCausalLM架构得到[Transformer, LogitScore],其余架构得到[Transformer]transformer_task="sequence-classification"

与 SentenceTransformer 模型类似,CrossEncoder 模型也可以声明正确的加载所需的依赖版本,参见 声明版本要求。

多模态 CrossEncoder 模型

Transformer模块原生支持多模态输入:它会通过modality_config自动检测底层模型支持的模态(text、image、audio、video),并将输入路由到正确的处理方法。使用多模态骨干即可构建多模态 CrossEncoder:

from sentence_transformers import CrossEncoder model = CrossEncoder("Qwen/Qwen3-VL-Reranker-2B") print(model) """ CrossEncoder( (0): Transformer({'transformer_task': 'any-to-any', 'modality_config': {'text': ..., 'image': ..., 'video': ..., 'message': {..., 'format': 'structured'}}, 'module_output_name': 'causal_logits', 'processing_kwargs': {...}, 'unpad_inputs': False, 'architecture': 'Qwen3VLForConditionalGeneration'}) (1): LogitScore({'true_token_id': 9693, 'false_token_id': 2152, 'module_input_name': 'causal_logits'}) ) """ # Score text-only and image-text pairs scores = model.predict([ ("A bee on a pink flower", "This image shows a bee on a pink flower"), ("A bee on a pink flower", "bee.jpg"), # 本地图像文件路径 ])

从零(基于多模态基座模型)构建多模态 CrossEncoder 时,使用transformer_task="any-to-any"加载带语言模型头的完整因果 LM(注意:该任务要求 transformers v5+,源码中对版本不满足的情况会抛出明确提示):

from sentence_transformers import CrossEncoder from sentence_transformers.cross_encoder.modules import Transformer, LogitScore transformer = Transformer("Qwen/Qwen3.5-0.8B", transformer_task="any-to-any") score_head = LogitScore( true_token_id=transformer.tokenizer.convert_tokens_to_ids("1"), false_token_id=transformer.tokenizer.convert_tokens_to_ids("0"), ) model = CrossEncoder( modules=[transformer, score_head], prompts={ "image_to_text": "Given the image, judge whether the text matches it. Respond with 1 if they match, 0 if they don't.", "text_to_image": "Given the text, judge whether the image matches it. Respond with 1 if they match, 0 if they don't.", }, )

注意prompts参数支持按方向命名不同的指令模板(image_to_text/text_to_image),这些 prompt 会随config_sentence_transformers.json一并保存,加载时可覆盖或清空。

进阶:自定义模块

输入模块(InputModule)

管道中的第一个模块称为输入模块(input module),负责预处理输入并为后续模块生成特征。输入模块可以是任何实现了InputModule的模块(InputModule继承自Module)。它有两个必须实现的抽象方法和一个应当重写的方法:

  • forward:接收features字典,返回更新后的features字典;
  • save:将模块配置(以及可选的权重)保存到指定目录;
  • preprocess:接收输入列表和可选prompt字符串,返回传递给模块自身forwardfeatures字典,键应与forward的期望一致(如input_idsattention_maskpixel_values等)。基类提供了委托给tokenize()的默认实现,但你应该重写它。完整签名如下(见 InputModule.preprocess):
def preprocess( self, inputs: list[SingleInput | PairInput], prompt: str | None = None, **kwargs, ) -> dict[str, torch.Tensor | Any]: ...

对 CrossEncoder 而言,inputs是成对的输入。每对中的每个元素可以是:文本字符串、PIL 图像、表示音频或视频的 numpy/torch 数组、带模态键的多模态字典(如{"text": ..., "image": ...}),或聊天风格的消息列表。

可选实现:

  • modalities属性:返回该模块支持的模态列表,默认["text"]。每个条目可以是单个模态字符串(如"text")或复合模态的元组(如("image", "text")),元组应按字母序排列。BaseModel.preprocess会基于此列表在调用模块的preprocess()前校验输入模态,因此你无需自行处理不支持的模态;
  • load类方法:从保存目录或 Hub 模型加载模块;
  • max_seq_length属性:返回模块可处理的最大序列长度。

注意:tokenize()方法已被弃用,新的自定义输入模块应实现preprocess()。从源码可见,基类preprocess()会检测子类是否只重写了tokenize()并给出弃用警告后兼容回退,但新代码不应依赖这一兼容路径。

非输入模块(Subsequent Modules)

管道中其余的模块称为非输入模块。它们处理输入模块产出的 features 字典,要么转换特征,要么提取最终得分。非输入模块可以是任何实现了Module的模块,必须实现:

  • forward:接收features字典并返回更新后的字典。对 CrossEncoder,最后一个模块应写入"scores"——这正是predict读取的得分位置;
  • save:将模块配置和可选权重保存到目录。

可选实现load类方法。

示例:温度缩放打分模块

下面自定义一个打分模块,在对 logits 计算 log-odds 得分前先做温度缩放:

# temperature_logit_score.py import torch from sentence_transformers.cross_encoder.modules import Module class TemperatureLogitScore(Module): config_keys: list[str] = ["true_token_id", "false_token_id", "temperature", "module_input_name"] def __init__( self, true_token_id: int, false_token_id: int | None = None, temperature: float = 1.0, module_input_name: str = "causal_logits", **kwargs, ) -> None: super().__init__() self.true_token_id = true_token_id self.false_token_id = false_token_id self.temperature = temperature self.module_input_name = module_input_name def forward(self, features: dict[str, torch.Tensor], **kwargs) -> dict[str, torch.Tensor]: logits = features[self.module_input_name][:, -1] # Apply temperature scaling logits = logits / self.temperature if self.false_token_id is None: scores = logits[:, self.true_token_id] else: scores = logits[:, self.true_token_id] - logits[:, self.false_token_id] features["scores"] = scores.unsqueeze(1) return features def save(self, output_path: str, *args, safe_serialization: bool = True, **kwargs) -> None: self.save_config(output_path) # The default `load` method reads `config.json` and passes the values # (i.e. the `config_keys`) as kwargs to __init__. This works for us, # so no need to override it.

这个示例的三个关键点:

  • config_keys:声明应写入config.json的实例属性列表。这些值由 Module.get_config_dict 自动提取({key: getattr(self, key) for key in self.config_keys}),并在默认的load方法中作为关键字参数重建模块——所以只要属性名与__init__参数名一致,就无需重写load
  • forward:从 features 字典读取 logits 并写入"scores"键,这是CrossEncoder.predict期望的最终模块输出;
  • save:调用Module.save_config写出配置。由于该模块没有可训练权重,无需其他操作;含可训练参数的模块还应调用Module.save_torch_weights(见 save_config / save_torch_weights)。

将其装入 CrossEncoder:

from sentence_transformers import CrossEncoder from sentence_transformers.cross_encoder.modules import Transformer from temperature_logit_score import TemperatureLogitScore transformer = Transformer("Qwen/Qwen3-Reranker-0.6B", transformer_task="text-generation") score_head = TemperatureLogitScore( true_token_id=transformer.tokenizer.convert_tokens_to_ids("yes"), false_token_id=transformer.tokenizer.convert_tokens_to_ids("no"), temperature=2.0, ) model = CrossEncoder(modules=[transformer, score_head]) print(model) """ CrossEncoder( (0): Transformer({'transformer_task': 'text-generation', ...}) (1): TemperatureLogitScore({'true_token_id': 9693, 'false_token_id': 2152, 'temperature': 2.0, 'module_input_name': 'causal_logits'}) ) """ scores = model.predict([("How many people live in Berlin?", "Berlin has a population of 3,520,031.")])

此时调用save_pretrained生成的modules.json会是:

[ { "idx": 0, "name": "0", "path": "", "type": "sentence_transformers.base.modules.transformer.Transformer" }, { "idx": 1, "name": "1", "path": "1_TemperatureLogitScore", "type": "temperature_logit_score.TemperatureLogitScore" } ]

为了让temperature_logit_score.TemperatureLogitScore可以被导入,你需要把temperature_logit_score.py复制到模型保存目录;如果将模型推送到 Hub,也要把该文件一并上传到模型仓库。这样其他人就可以通过CrossEncoder("your-username/your-model-id", trust_remote_code=True)使用你的自定义模块。

注意:使用自定义模块时,无论用户从 Hub 还是本地路径加载模型,都必须显式设置trust_remote_code=True,这是防止远程代码执行攻击的安全措施。

另外,在__init__forwardsaveloadpreprocess方法中保留**kwargs参数是推荐的,以保证方法在未来库版本更新时保持兼容。

将自定义模块放入独立仓库

如果你的模型和自定义建模代码都在 Hub 上,可以考虑把自定义模块拆分到一个独立仓库:这样只需维护一份实现,就能在多个模型间复用。做法是把modules.jsontype字段改为{repository_id}--{dot_path_to_module}格式。例如temperature_logit_score.py存放在仓库my-user/my-model-implementation中,modules.json可以写成:

[ { "idx": 0, "name": "0", "path": "", "type": "sentence_transformers.base.modules.transformer.Transformer" }, { "idx": 1, "name": "1", "path": "1_TemperatureLogitScore", "type": "my-user/my-model-implementation--temperature_logit_score.TemperatureLogitScore" } ]

进阶:自定义模块中的关键词参数透传

如果希望用户在调用CrossEncoder.predict时能为你的模块传入自定义关键词参数,可以把参数名添加到modules.json中对应模块的"kwargs"字段。例如,若模块需要根据task参数改变行为:

[ { "idx": 0, "name": "0", "path": "", "type": "custom_transformer.CustomTransformer", "kwargs": ["task"] }, { "idx": 1, "name": "1", "path": "1_LogitScore", "type": "sentence_transformers.cross_encoder.modules.logit_score.LogitScore" } ]

然后在自定义模块的forward中读取task

from sentence_transformers.cross_encoder.modules import Transformer class CustomTransformer(Transformer): def forward(self, features: dict[str, torch.Tensor], task: str | None = None, **kwargs) -> dict[str, torch.Tensor]: if task == "default": # Do something ... else: # Do something else ... return features

此后用户即可在调用predict时传入task参数:

from sentence_transformers import CrossEncoder model = CrossEncoder("your-username/your-model-id", trust_remote_code=True) pairs = [("query", "document")] model.predict(pairs, task="default")

这一机制的底层实现位于 BaseModel.forward:前向传播时,框架会检查每个模块在modules.json中声明的module_kwargs(或模块自身的forward_kwargs属性),只把匹配到的键值转发给该模块,其余模块不受影响。保存时,BaseModel.save 会把self.module_kwargs写回modules.json"kwargs"字段,形成完整的存取闭环。

小结

本文以 sentence-transformers 的 CrossEncoder 自定义模型文档为主线,覆盖了模块化模型的完整生命周期:

  1. 架构:CrossEncoder 是torch.nn.Sequential的子类,模型 = 顺序模块链,features 字典在模块间流转,最终模块负责写入"scores"
  2. 模块链选型:编码器分类模型用单模块sequence-classification;生成式重排模型用text-generation+LogitScore;显存敏感场景可用feature-extraction+Pooling("lasttoken")+Dense并借助嵌入差值初始化 Dense 权重来近似 LogitScore;
  3. 序列化modules.json(模块类型与路径)、config_sentence_transformers.json(模型级配置)、各模块子目录(模块配置与权重)共同构成可复现的模型格式,加载时据此重建模块链,无modules.json时按架构自动回退到默认模块链;
  4. 自定义模块:继承InputModule/Module,实现preprocessforwardsave,用config_keys声明持久化属性;自定义模块文件需随模型分发,加载方需trust_remote_code=True
  5. 参数透传:通过modules.json"kwargs"字段 + 模块forward中的对应参数名,即可让predict的用户参数精准路由到指定模块。

相关实现可进一步深入以下文件:CrossEncoder 主实现、LogitScore、模块基类 Module、输入模块基类 InputModule、模块包入口,以及多模态对比训练示例 examples/cross_encoder/training/multimodal。

【免费下载链接】sentence-transformersState-of-the-Art Embeddings, Retrieval, and Reranking项目地址: https://gitcode.com/gh_mirrors/se/sentence-transformers

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询