ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

MLX 数组框架入门:3 步在 Apple 芯片上跑通并训练模型

MLX 数组框架入门:3 步在 Apple 芯片上跑通并训练模型 MLX 数组框架入门3 步在 Apple 芯片上跑通并训练模型【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx在 Mac 上跑机器学习模型最常见的卡点有两类框架不支持 Apple 芯片或者要在 CPU 和 GPU 之间手动搬运数据。MLX 是 Apple 出品、面向 Apple 芯片机器学习的数组框架NumPy 风格的 API、统一内存模型切换计算设备不需要移动数据。下面用 4 段不超过 15 行的代码讲清安装自查、延迟计算与设备切换并训练一个线性回归模型。安装 MLX 前先核对三件事从 PyPI 安装 MLX 有三个硬性条件按顺序核对芯片与系统Apple 硅芯片M 系列 macOS 14.0 及以上。Python 架构必须是原生nativePython 3.10 及以上。运行python -c import platform; print(platform.processor())输出应为arm如果是i386说明你用的是 Rosetta 下的 x86 Pythonpip 会找不到匹配的 mlx 包。安装目标macOS 直接pip install mlxLinux 用户按场景选pip install mlx[cuda12]CUDA 后端或pip install mlx[cpu]纯 CPU。需要编译 C API 或做定制构建时要求 C20 编译器、CMake 3.25、macOS 上 Xcode 15完整构建参数见 docs/src/install.rst。跑通第一个 MLX 数组安装完成后导入mlx.core即可开始import mlx.core as mx a mx.array([1, 2, 3, 4]) # 推断为 int32 b mx.array([1.0, 2.0, 3.0, 4.0]) # 推断为 float32 c a b # 此时尚未计算 print(c) # 打印触发求值 # array([2, 4, 6, 8], dtypefloat32)三处值得注意a推断为int32、b为float32相加后结果自动提升到float32c a b这一步并不真正计算打印时结果才被算出来。基础用法更多细节见 docs/src/usage/quick_start.rst。讲透延迟计算与 eval 的时机延迟计算lazy evaluation是 MLX 的核心设计每个操作只是把子图记入计算图真正执行发生在结果被需要时。它带来两个直接好处只算你用到的部分函数里生成了但从未被使用的数组不会真的被计算。函数变换有对象可操作mx.grad自动微分、mx.vmap自动向量化都作用在记录好的计算图上并且可以任意嵌套组合例如mx.grad(mx.vmap(f))。什么算被需要四类情况会隐式触发求值打印数组、调用.item()取标量、转换成 NumPy 数组、用mx.save存盘。想显式触发则用mx.eval(c)。eval 放哪每次求值都有固定开销计算图过大也有成本。通行做法是在外层循环如训练迭代的边界求值一次单次求值覆盖几十到几千个操作都可以。完整讨论见 docs/src/usage/lazy_evaluation.rst。不搬数据切换计算设备统一内存模型指数组不属于任何设备而存在 CPU 和 GPU 共享的内存池中。你在执行操作时指定由哪个设备干活即可a mx.random.normal((100,)) b mx.random.normal((100,)) mx.add(a, b, streammx.cpu) # CPU 执行 mx.add(a, b, streammx.gpu) # GPU 执行数据原地不动这两个加法互不依赖可以并行执行。若存在依赖第二个操作使用第一个的结果调度器会自动插入依赖关系第二个操作等第一个完成后才开始无需你手动同步。对比其他框架先显式把数组搬到目标设备的写法这里只需在操作参数里写stream。原理与示例见 docs/src/usage/unified_memory.rst。20 行训练一个线性回归模型完整示例在 examples/python/linear_regression.py生成带噪声数据、定义损失、梯度下降迭代核心如下import mlx.core as mx X mx.random.normal((1000, 100)) y X mx.random.normal((100,)) w 1e-2 * mx.random.normal((100,)) def loss_fn(w): return 0.5 * mx.mean(mx.square(X w - y)) for _ in range(10_000): grad mx.grad(loss_fn)(w) # 对整个计算图自动微分 w w - 0.01 * grad mx.eval(w) # 每轮迭代求值一次运行后会打印 loss、学到的参数与真实参数的 L2 距离以及训练吞吐it/s。注意反向传播没有手写mx.grad对整张计算图自动完成。想使用mlx.nn的层和mlx.optimizers的优化器可参考 examples/python/logistic_regression.py。热点函数想再提速用mx.compile包一层def fun(x, y): return mx.exp(-x) y compiled_fun mx.compile(fun) # 首次调用时编译之后走缓存 print(compiled_fun(x, y))首次调用会构建计算图、优化并生成编译代码比较慢相同输入签名下再次调用直接命中缓存。注意输入形状、类型或参数个数变化都会触发重新编译所以不要给频繁创建销毁的函数做编译。细节见 docs/src/usage/compile.rst。避开这三个高频坑pip 找不到 mlx九成是 Python 不是原生 arm。用前面提到的命令确认输出为arm必要时换一套原生 Python 再装。mx.eval 调用过频每做一个小操作就求值一次会把固定开销叠加得训练明显变慢攒到迭代边界统一求值。mx.compile 滥用改形状、改类型都会触发重编译变化后的首次调用会明显卡顿只编译你会反复调用的函数。延伸 GPU 调试与多设备推理需要分析 GPU 性能时可用CMAKE_ARGS-DMLX_METAL_DEBUGON编译运行时设置MTL_CAPTURE_ENABLED1再用mx.metal.start_capture(mlx_trace.gputrace)和stop_capture()捕获 GPU 工作负载最后在 Xcode 里打开.gputrace文件查看每个操作的依赖关系方法见 docs/src/dev/metal_debugger.rst。若想在多设备间做推理与训练MLX 支持张量并行与数据并行可以把线性层的权重按列或行切分到多个 GPU 上分布式用法见 docs/src/usage/distributed.rst模型保存与加载mx.savez、mx.load及 safetensors、GGUF 格式见 docs/src/usage/saving_and_loading.rst。下一步运行python examples/python/linear_regression.py把num_iters从 10000 改为 1000、把lr改为 0.05对比两次输出的 loss 与 L2 距离观察学习率对收敛的影响。【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表