文档导航

Laya MLX 性能研究

研究日期:2026-09-19。目标:Apple M3 Max,40 个 GPU 核心,128 GiB 统一内存,MLX/MLX Metal 0.32.2。这是对原生运行时、已安装的 MLX 实现、官方文档和现有基准 JSON 的静态审阅。本次研究没有运行任何 GPU 基准或模型推理。本报告中下面提出的任何优化都没有实测加速。

最初的实验应该是整模型编译和有代表性的批调度,然后是选择性的量化矩阵乘法。它们处理的是占主导的重复工作。针对局部注意力的专用 kernel 对长输入来说是一个可信的长期项目。精确裁剪最后决策头层是可行的,但它对整模型的算术节省只有几个百分点。在不改变 checkpoint 的前提下取得大幅改进,需要改进稠密主干、消除真正冗余的请求,或找出一个经实测的实现瓶颈;仅仅替换一个激活函数或再打开一个 attention 标志不太可能够用。

现有测量确立了哪些结论

下面这些是现有的端到端中位延迟,包含提示准备和结果格式化,带同步的 GPU 完成、五次预热和 50 次测量迭代。模型加载和下载不计入。基准允许 64 个问题的批;公开运行时的默认值是 16,所以 50 问题的结果不是默认 API 配置。

Checkpoint / 精度 短 1 问题 短 10 问题 短 50 问题 长 1 问题 长 10 问题
Laya MLX FP16 13.421 ms 71.068 ms 336.030 ms 44.927 ms 420.987 ms
Laya MLX FP32 15.954 ms 98.820 ms 450.712 ms 61.331 ms 534.242 ms
Laya stock Torch MPS FP32 24.918 ms 95.265 ms 497.856 ms 65.581 ms 586.594 ms
Multilingual MLX FP16 7.390 ms 27.386 ms 127.565 ms 37.635 ms 389.487 ms
Multilingual MLX FP32 7.988 ms 32.337 ms 151.387 ms 47.208 ms 451.331 ms
Multilingual stock Torch MPS FP32 19.349 ms 43.158 ms 194.171 ms 52.939 ms 534.492 ms

来源:Laya FP16、Laya FP32、Laya MPS、multilingual FP16、multilingual FP32,以及 multilingual MPS。长输入对 Laya 含 512 个 token,对 multilingual 含 1024 个;因此比较它们的长输入延迟并不是在相同序列长度下比较。短输入填充长度分别是 93 和 91。Torch 的比较必须保留 FP32 标签:与 MLX FP16 相比时,它们把后端变化和精度变化合在了一起。

运行之间存在明显的波动。例如,multilingual FP16 长 10 问题运行的 p50 是 389.487 ms、p95 是 462.319 ms、最大值是 619.663 ms。它短单问题的前向中位数是 8.023 ms,而独立测得的端到端中位数是 7.390 ms。把这些中位数相减会得到一个荒谬的负预处理时间。当前文件没有隔离 tokenizer、Python 调度、单个 GPU kernel 或同步的开销。它们确立了有用的基线,而不是 kernel 级的瓶颈诊断。

验证报告 记录了三个 checkpoint 在 FP32 和 FP16 下各 63/63 的 argmax 一致,以及每个变体 100 次有限、确定性的重复调用:共 378/378 个答案一致和 600 次重复调用。这是在一个小型固定样本语料上的回归检查,其中包含重复的问题。它们不是「未来的量化或架构改动会保持通用任务准确率」的证据。

为什么稠密矩阵乘法值得优先

模型实现 对每一个填充过的 token 施加 QKV 和输出投影、一个带门控的编码器 MLP,以及两个常规的 Transformer 决策头层。设 D 为隐藏维度,I 为编码器中间维度,N 为编码器层数,H 为决策头层数。这些块中每个 token 用到的矩阵权重数量是:

A = N * (4 * D^2 + 3 * D * I) + H * 12 * D^2
dense FLOPs per batch ~= 2 * B * L * A
dense attention FLOPs ~= 4 * B * (N + H) * L^2 * D

这些估计把乘法和加法分开计数,并排除归一化、激活、嵌入、评分、掩码、内存搬运和 kernel 开销。它们是一个算术模型,不是运行时剖面。

Checkpoint 系列 D / I / 编码器层数 Token 嵌入权重 主要的逐 token 矩阵权重 A 编码器全局 / 局部层
Laya / typed decisions 1024 / 2624 / 28 51,576,832 368,312,320 10 / 18
Multilingual 768 / 1152 / 22 196,608,000 124,452,864 8 / 14

multilingual checkpoint 约 3.22 亿的总参数里有 1.966 亿是嵌入参数。推理时只收集选中的嵌入行;这不是一次完整的词表投影。它主要的逐 token 矩阵工作量大约是英文模型的三分之一,尽管两者的总参数量看起来接近得多。这与实测的短批延迟差距一致,尽管它并不能证明某个特定的硬件瓶颈。只量化嵌入主要会减小常驻权重的体积,对 multilingual 尤其如此;它不必改善推理延迟。

在当前稠密注意力路径下,attention 乘积在 Laya 的 L=93 时约占建模 FLOPs 的 1.5%,Laya 的 L=512 时占 7.9%,multilingual 的 L=1024 时占 23.3%。它们占真实时间的比例可能大不相同。在投入自定义 Metal 工作之前,剖析器应当区分矩阵 kernel、注意力 kernel、逐元素 kernel、CPU 图构造和空闲间隙。

按优先级排列的实验

优先级 实验 最佳目标 主要取舍 / 验收条件
P0 对一短一长两个形状做剖析;编译模型或编码器块 单请求延迟和 Python 调度 在现有精度容差内保持输出一致;单独测量首次使用编译
P0 按实际 token 预算和长度分布调优批处理 许多不同问题和混合长度流量 在 p95 延迟和内存预算约束下优化吞吐;把排队计入
P1 量化选定的主干线性层,先 8-bit 后 4-bit 权重流量,或许还有稠密推理 质量与校准门槛;实际的 M3 Max 速度可能倒退
P1 缓存共享的 state 分词和稳定的问句模板 许多问题共享一个 state 或重复的评分细则 精确的 token 同一性;有界缓存;把缓存和未缓存的结果分开
P1 只计算最后一个头层所需的输出 每一个工作负载,尤其是更长的序列 精确的依赖裁剪;整模型算术节省不大
P2 带 tile 边界的真正局部窗口注意力 512/1024 token 工作负载 新 kernel 的复杂度;保留双向窗口和填充语义
P2 只在剖析支持的地方融合 residual/norm 或 GELU/gate 小 kernel 开销或激活流量 现有的快速 kernel 已经覆盖了其中大部分;保持精确的 GELU 语义
单独的产品特性 对相同的 forward 输入去重 确实重复问题的负载 报告唯一推理数和缓存命中;不要把它当作通用的 kernel 加速来呈现

编译完整的、被评估的推理路径

Agent.forward() 目前构造数组、调用 self.model 并求值结果。在模型或块层级没有外层的 mx.compile。固定形状编译可以减少 Python 图构造并融合受支持的操作。MLX 记录了形状特化和显式状态捕获;它也警告说,无形状编译无法安全地保留任意依赖形状的 Python 操作。见官方编译指南。

从加载、转换并求值权重之后创建的已编译可调用对象开始,使用常规的形状特化。保留现有的未编译可调用对象,用于一致性比较和 CPU 兼容。对于一个冻结的推理模型,权重可以保持为该模型实例所捕获;如果权重或模块结构改变,重建该可调用对象或显式捕获相关状态。不要在替换 checkpoint 时复用一个已编译闭包。

测试一个整模型包装器;如果追踪限制或编译成本让它不划算,就分别编译编码器块和头。当前代码把 x.shape 读进 Python 整数、用显式的批/长度值做 reshape、创建 arange(length) 掩码,并用由形状推导出的行范围索引标记。在不重新设计的情况下对这个完整图应用 shapeless=True 是不安全的。把动态掩码构造移到已编译块之外,并使用与形状无关的展平/还原操作,可能使之后的无形状变体成为可能;要在变化的 B、L 和标记数量上验证它。

形状桶可以限制重追踪,但填充有计算成本。把 L=93 填充到 96 会增加约 3.2% 的逐 token 工作;填充到 128 会增加约 37.6%。比较精确形状编译与小的长度倍数,以及一组有界、来自工作负载的桶。把 (batch size, padded length, marker slots, dtype, device/model instance) 纳入缓存策略决策,并测量形状频繁变动下的冷编译延迟和保留内存。

worker.py 里独立的 forward 基准直接调用 agent.model。如果只在 Agent.forward 里加入编译,当前的 forward 基准会绕过它,而端到端基准会使用它。为了让比较有意义,两条路径都必须显式选择同一个候选实现。在计时流程中保留 mx.eval 和 GPU 同步:只测量图构造并不会测量推理。

按有用的 token 批处理,然后检查矩阵调度

collate_items 把每个块右侧填充到它最长的序列;运行时按插入顺序对问题分组。对异构流量,按准备后的长度排序或分桶,在问题数量上限之外再加一个 token 预算,并恢复原始的问题 ID 和输出顺序。只在能代表该服务的地方比较 1、2、4、8、16、32 和 64 的批。对在线请求,要把等待一个批的时间算进去;仅看离线的问题/秒会掩盖不可接受的延迟。

当前的短 50 问题工作负载对 Laya 浪费约 8.9% 的填充 token,对 multilingual 浪费 5.8%。长基准行没有填充浪费。因此,仅靠去填充或排序在这些固定样本上的算术收益有限。要揭示生产收益,需要一种混合长度的分布,比如许多短问题中夹一个长问题。移除填充必须保留逐样本的 RoPE 位置、标记位置和注意力边界;在没有隔离掩码的情况下把样本拼进一个序列会改变模型。

对稠密 kernel,检查实际的形状和步长。QKV 已经是一次投影,编码器 MLP 的两个输入分支已经共享一次投影。不加区分地拆分会增加 launch。只有在剖析器或 MLX 调度追踪显示有不理想的批 GEMM 时,才比较把连续的 [B,L,D] 输入显式展平成 [B*L,D];框架可能已经高效地展平了。仅凭源码审查不能成为「漏掉了一个 GEMM 优化」的理由。

不要在生产实现里每一层之后插入同步。当前的运行时每个块只求值一次。额外的等待可能消除 CPU/GPU 的重叠并掩盖调度改进;层级的剖析应当是一次单独的诊断运行。

量化:瞄准主干并实现其存储契约

已安装的 MLX 0.32.2 量化层实现 提供 nn.quantize(..., class_predicate=...) 和 QuantizedLinear,经由 mx.quantized_matmul 做仅权重的矩阵乘法。分组仿射量化支持 8-bit 和 4-bit 实验。从组大小 64 的编码器线性层开始,把激活、归一化、类型嵌入、评分器和动作头保留为 FP16。然后独立地加入决策头线性层,并可选地加入嵌入量化。测量每个变体;在高 token 数量下,仅权重 kernel 可能不如 FP16 GEMM。

当前的加载器和模型里有具体的集成隐患:

  1. Agent.__init__ 把每一个存储的权重转换为浮点 dtype,并在严格加载之前只实例化稠密模块。量化 checkpoint 需要描述所选模块、组大小、位宽和模式的元数据;在加载前实例化匹配的量化模块并保留打包的整数权重。把打包权重转换成浮点并不是合法的加载。
  2. 动作头的第一个线性层输入宽度是 D+4,即 1028 或 772,不能被 32、64 或 128 的仿射组大小整除。因此一刀切的量化调用不合适。转换前检查每一个所选层的输入宽度。
  3. DecisionModel.__call__ 从 self.act_head.layers[0].weight.dtype 选择动作头的输入 dtype。对量化层来说,那个权重会是打包的整数存储,而不是所期望的激活 dtype。最初排除动作头可以避开这条路径;之后要支持它需要一个显式的激活 dtype 契约。
  4. 量化会改变 logits 和校准概率。现有小型 FP16 固定样本的一致性不足以证明 4-bit 质量。使用留出的带标签 choice、score 和 noul 任务、多语言输入、接近的决策、不同选项数量和升级例子。跟踪 argmax 一致、任务准确率、评分误差、概率漂移、校准和动作概率。饱和的动作输出可能掩盖巨大的动作 logit 变化。

对组大小 64、FP16 仿射 scale 和 offset 的情况,近似的矩阵存储是每个参数 bits/8 + 4/64 字节:8-bit 时 1.0625 字节,4-bit 时 0.5625 字节,对比 FP16 的 2 字节。这些是量化矩阵的存储估计,不含其他张量和打包开销;它们不是加速估计。官方 quantize API 描述了组整除性和格式。

更新的低位格式应当针对实际的 M3 Max 后端来评估,而不是假定它会使用来自更晚 Apple 芯片的硬件。MLX 的 NAX 可用性检查要求比所记录的 applegpu_g15s 设备更新的架构代。见 MLX 0.32.2 设备检查。

在输入确实相同时复用 CPU 准备

build_sequence 为每一个问题分别把同一个 state 序列化、清洗和分词。它还分别对每条指令和每个选项分词。Rust tokenizer 已经被直接使用;替换 Transformers 分词不是一个尚未做的优化。

每次 prepare 调用把 state 只序列化并清洗一次、编码一次,再把它的 token ID 切片到每个问题可用的空间。当同一个评分细则跨 state 使用时,缓存不可变的已准备问题前缀,键要包含 tokenizer 身份/revision、问题类型、有序 criteria、指令序列化、特殊 token 清洗和 token 预算。有界缓存在 tokenizer 或配置变化后不得复用结果。批量的 tokenizer 编码是另一个实验,前提是它的输出与当前那一串独立编码完全一致。

不要把对新拼接的提示分词当作「拼接独立编码片段」的替代:子词边界可能改变。逐字节校验输入 ID、注意力掩码、标记位置、qtypes、截断行为、结构化 criteria、mask 字面量、空输入和输出映射。

对长基准来说,state 在截断前把一句话重复了 200 次。避免 N 次重复的 state 编码可能有助于 CPU 准备,但现有的端到端/前向计时差并没有测量这项节省。分别对 prepare、collate/数组构造、forward 和后处理做基准,然后用一个不重复的留出语料确认端到端结果。

解码器 KV 缓存不适用于这个编码器。 它的第一层是全局双向注意力;state token 表示取决于问题、选项及其位置。在不同的问题之间复用 state 隐藏状态或 K/V 会改变结果。分词和完全相同的整输入结果可以缓存;任意的上下文编码器状态不能。

精确裁剪最终决策头的输出

在最后一个 HeadLayer 之后,只有 [CLS] token 和选项标记 token 会被用到。更早的头层仍必须产生所有 token,因为最后一层读取它们的 K/V。只在最后一层中:

  1. 归一化所有输入 token 并计算所有 K/V。
  2. 在 [CLS] 和合法选项位置收集 Q,并把这些 query 跑在整个带掩码的 K/V 序列上。
  3. 只对这些选定位置施加输出投影、残差、第二次归一化和前馈网络。
  4. 用选中的 [CLS] 输出给动作头,用选中的标记输出做评分;保留标记填充和原始顺序。

第一版实现可以保留融合的完整 QKV 投影,之后再收集 Q。更激进的变体把它的权重拆成一个全长 KV 投影和一个选定 token 的 Q 投影。那会节省更多算术,但可能让 GEMM 调度效率更低。重复的填充标记索引只有在被掩码的结果保持不可见时才无害。1 选项、多选项和可变标记的情形需要显式的保真检查。由于更少的 query 可能选中不同的 SDPA kernel,数学等价并不意味着浮点结果逐位相同。

设 R = 1 + number of option slots,保留完整 QKV 会从最后一个头层移除约 18 * B * (L-R) * D^2 的稠密 FLOPs 和 4 * B * L * (L-R) * D 的注意力 FLOPs。拆分 Q/KV 会把稠密系数从 18 变成 20。相对于上面的整模型算术估计,取 R=5 得到:

Checkpoint / 长度 保留融合的完整 QKV 也只计算选中的 Q
Laya, L=93 2.44% 2.70%
Laya, L=512 2.60% 2.86%
Typed decisions, L=1024 2.66% 2.90%
Multilingual, L=93 4.03% 4.47%
Multilingual, L=1024 4.22% 4.58%

这些是静态的 FLOP 削减,不是预测的延迟削减。该技术从一个头层移除了大部分工作,而不是从模型移除大部分工作。它有用,是因为它保留了依赖关系并能不改训练地实现,而不是因为它承诺整模型速度的某个倍数。

只有在测量其贡献之后,才构建真正的局部注意力

所有编码器注意力调用已经使用 mx.fast.scaled_dot_product_attention;RoPE 已经是 mx.fast.rope;nn.LayerNorm 调用快速归一化原语。MLX 的注意力 API 接受布尔掩码并在 FP32 中做 softmax。当前的 head 维度是 64。MLX 0.32.2 Metal 调度 支持这种形状配数组掩码,并且在推理时不会为它选择未融合的兜底实现。没有证据表明 Laya 的布尔掩码会禁用融合注意力。 已安装版本中可用的 force_fused=True 作为诊断断言很有用,但不该在这里被宣传成一条新的快速路径。

剩下的限制是结构化稀疏。该实现构造一个形状为 [B,1,L,L] 的稠密局部布尔掩码。在常规的 Metal 注意力 kernel 中,非因果循环遍历完整的 KV tile 范围;数组掩码在 QK 乘法之后作用于分数。它保留了局部注意力语义,却没有利用局部 tile 范围。

一个精确的专用 kernel 可以把每个 query tile 限制在重叠的 K/V 窗口内、保持 FP32 softmax 累加,并避免一个稠密的 L×L 掩码。正确的窗口是双向且包含端点的:abs(query_position - key_position) <= 64。内部 query 能看到 129 个位置,尽管配置名是 local_attention=128。全注意力层和两个决策头层必须保持全局。填充的 key 必须继续被排除,未使用的填充 query 需要有定义的有限行为。

一个省力的原型可以把 query 块与重叠的 K/V 切片分组,并用一个更小的精确掩码调用现有的 SDPA。使用已经定位好的 RoPE Q/K,或显式保留绝对偏移。优先使用成批的块而不是大量 Python 调用,并考虑重复的 K/V 物化。这个原型在短长度下可能不如当前 kernel;在维护自定义 Metal 之前,它是一个正确性和盈亏平衡实验。

移除所有被禁止的局部注意力对所带来的建模整模型 FLOPs 最大削减,对短输入很小,对长输入更有希望:

Checkpoint / 长度 精确局部稀疏带来的理想总 FLOP 削减
Laya, L=93 0.086%
Laya, L=512 3.61%
Typed decisions, L=1024 7.69%
Multilingual, L=93 0.147%
Multilingual, L=1024 11.92%

这些估计对 L>r 使用 local_pairs = L*(2*r+1) - r*(r+1),取 r=64。它们包含所有稠密投影和两个全注意力决策头层。它们排除掩码生成和内存流量。运行时收益可能高于或低于 FLOP 比例,因为注意力和 GEMM 的效率不同;只有剖析才能确定。在 8192 token 时权衡会不同,但随附的 agent 把输入限制在 512 或 1024,所以一个 8192 token 的说法需要一个单独支持的工作负载。

编译之外的融合

已安装的精确 nn.gelu 已经用无形状编译装饰,nn.Linear 已经在合适时使用带偏置感知的 addmm。整块编译仍可能把 GELU 与它的门控乘法、残差加法、类型转换、掩码和小型评分特征操作融合。在实现等价的专用 kernel 之前,检查已编译的 kernel 图。

如果激活流量仍然显著,原型化精确的 GELU-and-gate 或 residual-and-LayerNorm 融合。保留当前精确的基于 erf 的 GELU;tanh 或 sigmoid 近似会改变模型,需要单独的质量测量。在添加布局转换之前检查实际的 Q/K/V 步长和拷贝:MLX 的全注意力实现接受一个连续 head 维度配其他步长,并写出便于合并 head 的输出布局。无条件的连续拷贝可能增加工作。

标记 softmax、top-two 排序、熵、小型动作头和 NumPy 结果格式化只有在实测过之后才是合理的后续目标。与数百次全宽 Transformer 操作相比,选项位置很少,所以先优化它们不太可能触及占主导的路径。

在追求激进收益的同时保护基准的含义

当前的工作负载生成器 循环使用三个问题定义来构造 5、10 或 50 个问题。那些批里至多只有三个唯一的模型输入。逐调用的精确去重可以在真实应用里避免冗余推理,但它会不成比例地改善这些固定样本。把这项特性与 kernel 优化分开,并报告 questions、unique_forward_inputs、缓存命中和实际求值的 token。用各自的标签、有序 criteria 和校准元数据重建每一个原始答案。保留一个含 50 个真正不同问题的套件和一个刻意重复输入的套件。

不要在主推理基准里使用跨调用结果缓存:它会反复调用完全相同的请求。任何缓存实验都要如实标注。tokenizer 缓存基准应当同时包含重复评分细则的场景和全新输入的场景。

对每个候选,使用这样的实验设计:

  1. 保持目标不变。 记录源 revision/哈希、模型 revision、dtype、编译标志、所选量化模块、token ID 或其哈希、形状、标记数量、批策略、预热、同步和设备。现有的 input_sha256 哈希的是 state/questions,而不是实际的 token 张量;它不能确立跨 tokenizer 的张量同一性。从同一个最终源 revision 重跑基线,因为历史 JSON 的源哈希不同。
  2. 分离工作负载类别。 受控的 kernel 比较用固定形状;调度和编译用真实的可变长度;吞吐用不同的问题;合法的 CPU 缓存用重复的评分细则;去重用刻意重复的问题。包含长长度尾部区间和不同的选项数量。在每个后端比较内部保持源文本和截断一致。
  3. 测量冷热行为。 分别记录模型加载和首次编译。测量准备、图构造、同步求值的前向、输出转换和端到端延迟,不要把独立的中位数当作可加的。在单独一次运行里剖析 kernel,因为追踪会扰动延迟。
  4. 一次只改一个优化。 用现有的迭代次数筛选,然后在交替的基线/候选块里重复入围者,并收集足够样本以获得可信的 p95,例如每个工作负载至少 200 个计时请求。只跑一个活动的 GPU 基准,保持一致的功耗/散热条件,并保留原始样本。要求改进大于实测的运行波动。
  5. 检查正确性和稳定性。 在相同 dtype 下与 MLX 以及现有 FP32 参考比较;对调度和 CPU 准备的改动强制执行 token 同一性。测试 64、128 附近以及桶边界处的长度;块上限附近的批;单选项和多选项;多语言输入;填充;以及编译后的形状变化。重复变化形状的请求,在预热后同时观察活动内存和缓存内存,并验证每种配置内的结果有限且确定。
  6. 对近似改动施加更强的门槛。 量化、激活近似、token 裁剪、提前退出和蒸馏需要小型回归固定样本之外的留出任务与校准结果。改变训练行为时要保留单独的模型身份和基准标签。更快的 multilingual checkpoint 或蒸馏模型是另一种模型,而不是同一个 Laya checkpoint 的加速。

对任何占据端到端时间比例 f 并以因子 s 加速的实测热点,用 Amdahl 上界 1 / (1 - f + f/s) 评估预期的总影响。f 要用实测的时间比例;上面的算术比例不能替代它。下一个具体的工程决策应当跟随编译/批处理消融和短对长的 kernel 剖析,而不是一个未经验证的倍数。

文档与可复现性说明

文档查询使用了规定的 Context7 工作流:一次 library MLX 解析,然后分别对编译和快速注意力做官方文档查询,使用 /websites/ml-explore_github_io_mlx_build_html(共三条命令)。检查了已安装的 .pyi 和 Python 源码以核实 MLX 0.32.2 的行为,包括 force_fused、量化 API、已编译的 GELU 和快速 LayerNorm。阅读了锁定版本的上下游 C++/Metal 源码以了解调度和 tile 循环细节。本次研究没有升级任何库。

用于详细观察的两个 FP16 基线文件在检查时的 SHA-256 摘要如下:

laya-mlx-float16.json
63146dd664d039dde1a728b17aad896e491bd01ea36ea0786953691180e55b09

laya-multilingual-mlx-float16.json
9af74bd5a11e4edc15e6a8c9dc929a7b9fd2d19cb06f348cd0e04a076f912473

在本研究进行期间,主基准进程仍在产生额外的工件。表格刻意引用审阅时已经可用的完整基线文件,并且不对未测量的候选实现作任何声称。