M1 Mac用MPS后端运行T5模型遇int64错误,如何转为float?
在M1 Mac上解决T5模型MPS后端的int64类型错误
问题背景
我尝试在M1 Mac上使用MPS后端运行T5 transformer模型,代码如下:
import torch import json from transformers import T5Tokenizer, T5ForConditionalGeneration, T5Config #Make sure sentencepiece is installed device = torch.device('mps') model = T5ForConditionalGeneration.from_pretrained('t5-3b').to("mps") tokenizer = T5Tokenizer.from_pretrained('t5-3b')#, device = device) preprocess_text = full_text.strip().replace("\n",".") t5_prepared_Text = "summarize: "+preprocess_text print ("original text preprocessed: \n", preprocess_text) tokenized_text = tokenizer.encode(t5_prepared_Text, return_tensors="pt").to(device) # summmarize summary_ids = model.generate(tokenized_text, num_beams=6, no_repeat_ngram_size=3, min_length=30, max_length=9000, early_stopping=True) output = tokenizer.decode(summary_ids[0], skip_special_tokens=True) print ("\n\nSummarized text: \n",output)
其中full_text是提前定义的字符串。代码在CPU上运行正常,但切换到MPS加速时触发以下错误:
TypeError: Operation 'abs_out_mps()' does not support input type 'int64' in MPS backend.
需要找到让模型自动转换为支持的float类型、避免崩溃的方法。
可行解决方案
1. 针对特定张量手动转换类型
MPS后端不支持int64类型的abs操作,你可以定位触发错误的张量,在计算前显式转换为float32。注意T5的嵌入层需要整数类型的token ID输入,所以不能直接转换整个tokenized_text,而是在模型内部涉及abs计算的模块前转换:
# 示例:重写模型forward方法处理特定张量 class MPSFixedT5(T5ForConditionalGeneration): def forward(self, *args, **kwargs): # 先调用原forward获取输出 outputs = super().forward(*args, **kwargs) # 检查并转换触发错误的中间张量(根据报错栈调整) if hasattr(outputs, 'logits') and outputs.logits.dtype == torch.int64: outputs.logits = outputs.logits.to(torch.float32) return outputs model = MPSFixedT5.from_pretrained('t5-3b').to("mps")
2. 全局自动类型转换钩子
注册一个张量转换钩子,自动将MPS设备上的int64张量转为float32。这种方式无需修改模型结构,但要注意跳过嵌入层依赖的token ID张量:
def mps_auto_cast(tensor): if tensor.device.type == 'mps' and tensor.dtype == torch.int64: # 跳过输入ID张量,避免嵌入层报错 if not hasattr(tensor, 'names') or 'input_ids' not in tensor.names: return tensor.to(torch.float32) return tensor # 替换Tensor的to方法,添加自动转换逻辑 original_to = torch.Tensor.to def wrapped_to(self, *args, **kwargs): result = original_to(self, *args, **kwargs) return mps_auto_cast(result) torch.Tensor.to = wrapped_to
3. 强制模型使用float32精度加载
加载模型时指定torch_dtype=torch.float32,确保模型参数和大部分计算使用float32,减少整数类型的出现:
model = T5ForConditionalGeneration.from_pretrained( 't5-3b', torch_dtype=torch.float32 ).to("mps")
同时将辅助输入张量(如attention mask)转为float32:
# 生成完整tokenized输入(含attention mask) tokenized_input = tokenizer( t5_prepared_Text, return_tensors="pt", padding=True, truncation=True ).to(device) # 转换attention mask类型 tokenized_input['attention_mask'] = tokenized_input['attention_mask'].to(torch.float32) # 传入处理后的输入生成摘要 summary_ids = model.generate( **tokenized_input, num_beams=6, no_repeat_ngram_size=3, min_length=30, max_length=9000, early_stopping=True )
注意事项
- 优先测试手动转换特定张量的方式,避免全局转换带来的未知问题。
- 若遇到其他MPS不支持的操作,可临时将对应模块回退到CPU运行:
# 示例:将某个层强制放在CPU model.some_layer.to('cpu')
内容的提问来源于stack exchange,提问作者DrRaspberry
相关产品推荐
相关产品推荐

