Mac M1本地运行Databricks Dolly遇CUDA错误,求解决方案
在Apple M1 Mac上运行Databricks Dolly的CUDA错误解决方法
问题概述
在Apple M1芯片Mac上部署运行Databricks Dolly时,执行Hugging Face Transformers代码触发错误:AssertionError: Torch not compiled with CUDA enabled。原因是M1 Mac基于ARM架构,不支持CUDA,PyTorch在此环境下使用Metal后端实现GPU加速。
解决步骤
1. 安装适配Apple Silicon的PyTorch
卸载原有PyTorch,重新安装针对Apple Silicon优化的版本:
pip3 uninstall -y torch torchvision torchaudio pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu
此版本的PyTorch会自动启用Metal加速,无需CUDA依赖。
2. 修改代码中的设备配置
将代码中所有指定cuda的部分替换为mps(Metal Performance Shaders的缩写),同时确保模型加载逻辑适配Apple Silicon:
修改后的完整代码:
import numpy as np from transformers import ( AutoModelForCausalLM, AutoTokenizer, PreTrainedModel, PreTrainedTokenizer ) tokenizer = AutoTokenizer.from_pretrained("databricks/dolly-v1-6b", padding_side="left") # 自动识别设备,优先使用MPS model = AutoModelForCausalLM.from_pretrained( "databricks/dolly-v1-6b", device_map="auto", trust_remote_code=True, offload_folder='offload' ) PROMPT_FORMAT = """Below is an instruction that describes a task. Write a response that appropriately completes the request. ### Instruction: {instruction} ### Response: """ def generate_response(instruction: str, *, model: PreTrainedModel, tokenizer: PreTrainedTokenizer, do_sample: bool = True, max_new_tokens: int = 256, top_p: float = 0.92, top_k: int = 0, **kwargs) -> str: # 将输入张量转移到MPS设备 input_ids = tokenizer(PROMPT_FORMAT.format(instruction=instruction), return_tensors="pt").input_ids.to("mps") response_key_token_id = tokenizer.encode("### Response:")[0] end_key_token_id = tokenizer.encode("### End")[0] gen_tokens = model.generate(input_ids, pad_token_id=tokenizer.pad_token_id, eos_token_id=end_key_token_id, do_sample=do_sample, max_new_tokens=max_new_tokens, top_p=top_p, top_k=top_k, **kwargs)[ 0].cpu() response_positions = np.where(gen_tokens == response_key_token_id)[0] if len(response_positions) >= 0: response_pos = response_positions[0] end_pos = None end_positions = np.where(gen_tokens == end_key_token_id)[0] if len(end_positions) > 0: end_pos = end_positions[0] return tokenizer.decode(gen_tokens[response_pos + 1: end_pos]).strip() return None # 测试生成 generate_response("Write a tweet announcing Dolly, a large language model from Databricks.", model=model, tokenizer=tokenizer)
3. 验证MPS可用性
运行以下代码确认PyTorch已启用MPS支持:
import torch print(torch.backends.mps.is_available())
输出True表示配置成功。
内容的提问来源于stack exchange,提问作者Kush Verma
相关产品推荐
相关产品推荐

