sentence-transformer==2.2.2更新后报错及模型加载问题求助
解决模型加载的RuntimeError问题
错误原因
这个报错是因为模型保存时是在带有MPS(Apple Silicon GPU)的环境中,而你的EMR环境没有MPS设备,加载时Torch尝试将模型参数映射到不存在的MPS设备,导致失败。
解决方案
直接在加载模型时指定map_location参数,强制将模型加载到当前环境可用的设备(比如CPU),修改代码如下:
import torch import io # 读取S3中的模型字节数据 model_data = s3.Bucket(bucket_name).Object(object_path).get()['Body'].read() # 使用torch.load并指定map_location为CPU,避免MPS设备依赖 sentence_transformer_model = torch.load(io.BytesIO(model_data), map_location=torch.device('cpu'))
额外建议
如果模型是通过Sentence-Transformers官方的save()方法保存的,建议使用官方的SentenceTransformer.load()方法加载,兼容性更好:
- 先将S3上的模型目录下载到EMR本地临时路径(若模型是目录形式)
- 用以下代码加载:
from sentence_transformers import SentenceTransformer local_model_dir = "/tmp/your_model_dir" # 递归下载S3模型目录到本地(可通过aws cli或s3fs实现) sentence_transformer_model = SentenceTransformer(local_model_dir)
内容的提问来源于stack exchange,提问作者Arshee Siddiqui
相关产品推荐
相关产品推荐

