如何以16位/8位精度加载HuggingFaceCrossEncoder重排序模型?
解决HuggingFaceCrossEncoder加载16/8位精度的问题
问题描述
使用LangChain的HuggingFaceCrossEncoder加载重排序模型时,尝试通过use_fp16=True参数启用16位精度,触发ValidationError,因为该类并不直接支持use_fp16参数。
原实现代码:
from langchain.retrievers.document_compressors import CrossEncoderReranker from langchain_community.cross_encoders import HuggingFaceCrossEncoder model = HuggingFaceCrossEncoder(model_name="BAAI/bge-reranker-base") compressor = CrossEncoderReranker(model=model, top_n=4) compression_retriever = ContextualCompressionRetriever( base_compressor=compressor, base_retriever=retriever )
尝试启用16位精度的错误代码:
model = HuggingFaceCrossEncoder(model_name="BAAI/bge-reranker-base",use_fp16=True)
错误信息:
ValidationError: 1 validation error for HuggingFaceCrossEncoder
解决方案
LangChain的HuggingFaceCrossEncoder类通过model_kwargs参数传递配置项给transformers底层模型,需将精度设置放入该字典中:
1. 加载16位精度模型
from langchain.retrievers.document_compressors import CrossEncoderReranker from langchain_community.cross_encoders import HuggingFaceCrossEncoder # 通过model_kwargs传递浮点精度参数 model = HuggingFaceCrossEncoder( model_name="BAAI/bge-reranker-base", model_kwargs={"torch_dtype": "float16"} ) compressor = CrossEncoderReranker(model=model, top_n=4) compression_retriever = ContextualCompressionRetriever( base_compressor=compressor, base_retriever=retriever )
2. 加载8位精度模型(需依赖bitsandbytes库)
先安装依赖库:
pip install bitsandbytes
再配置8位加载参数:
model = HuggingFaceCrossEncoder( model_name="BAAI/bge-reranker-base", model_kwargs={"load_in_8bit": True} )
说明
model_kwargs会将内部参数传递给transformers的AutoModelForSequenceClassification.from_pretrained()方法,所有该方法支持的精度配置参数都可通过此字典传入。- 使用16位精度时,需确保硬件支持CUDA,否则可能出现性能问题或报错。
内容的提问来源于stack exchange,提问作者Debasish Pradhan
相关产品推荐
相关产品推荐

