You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.24 18:27:22