如何在TensorFlow中加载PluggableDevice以使用Mac M1 GPU?
在Rust中调用TensorFlow C库时加载MPS设备插件的方法
问题背景
我已按指定教程在Mac M1上成功构建了TensorFlow v2.9 C库,但在Rust代码中检测设备时,仅能识别CPU设备,bundle.session.device_list()仅返回CPU设备:
let bundle = SavedModelBundle::load( &SessionOptions::new(), &["serve"], &mut graph, export_dir ).expect("Unable to load model from disk"); println!("{:?}", bundle.session.device_list() )
输出结果:
Device { name: "/job:localhost/replica:0/task:0/device:CPU:0", device_type: "CPU", memory_bytes: 268435456, incarnation: 10072007419359857694 }]
该Rust代码使用TensorFlow C API的绑定(例如device_list对应TF_DeviceList)。Apple M1的GPU由MPS插件支持,经测试在Python中可正常工作。MPS是作为TensorFlow的PluggableDevice实现的,无需修改TensorFlow源码即可加载。
解决方案
1. 确认MPS插件路径
MPS插件的动态库文件(libmetal_plugin.dylib)通常位于Python的site-packages目录下,路径类似~/Library/Python/3.9/lib/python/site-packages/tensorflow_plugins/libmetal_plugin.dylib,请根据你的Python版本和实际安装路径调整。
2. 设置环境变量加载插件
在启动Rust程序前,通过环境变量指定插件路径,让TensorFlow C库加载MPS插件:
TF_PLUGGABLE_DEVICE_LIBRARY_PATH=/path/to/libmetal_plugin.dylib cargo run
也可以在Rust代码中直接设置环境变量:
use std::env; fn main() { env::set_var("TF_PLUGGABLE_DEVICE_LIBRARY_PATH", "/path/to/libmetal_plugin.dylib"); // 后续加载模型的代码 }
3. 验证设备加载结果
修改代码后重新运行,检查输出的设备列表,应该会包含MPS对应的GPU设备(名称格式类似/job:localhost/replica:0/task:0/device:GPU:0)。
关键注意事项
- 确保构建TensorFlow v2.9 C库时开启了PluggableDevice支持,该特性在v2.9版本中默认启用。
- 如果仍无法识别设备,检查插件路径是否正确,动态库文件是否存在且兼容你的M1架构。
内容的提问来源于stack exchange,提问作者Mr.Wang from Next Door
相关产品推荐
相关产品推荐

