segment_anything: 将注意力与 mx.fast.scaled_dot_product_attention 相融合
在构建基于 torch-mlx 的 SAM 端口(一个与之无关的项目,在此不提议)时,我阅读了此仓库自己的原生 `segment_anything` 实现,以供参考,并注意到其两个注意力块都是手动编写的 matmul+softmax+matmul,而不是 `mx.fast.scaled_dot_product_attention`: - `segment_anything/image_encoder.py` 的 `Attention.__call__`(约在第 224 行):计算 `attn = (q * scale) @ k.transpose(...)`,可选地添加 `add_decomposed_rel_pos` 的相对位置偏置,然后 `mx.softmax(attn, axis=-1)`,然后 `attn @ v`。 - `segment_anything/Transformer.py` 的 `Attention.__call__`(约在第 213 行,由 mask 解码器中的 `TwoWayTransformer` 使用):相同的手动模式,没有相对位置偏置。 `mx.fast.scaled_dot_product_attention` 的 `mask` 参数接受与 `[B, N, T_q, T_kv]` 相兼容的任意加法数组(通过其文档进行了验证),而不仅仅是一个布尔/因果性掩码 - 因此,图像编码器的 `add_decomposed_rel_pos` 输出可以直接作为 `mask=` 传递,而不需要将其添加到手动生成的注意力矩阵中。 Transformer 注意力没有这种偏置,并且会更直接地融合(`mask=None`)。 这种改进将避免生成完整的 `H*W x H*W` 注意力矩阵(在 ViT 后端固定的 64x64 切片网格中,每个头为 4096x4096),并运行单独的 `mx.softmax` 步骤,而优先使用单个融合内核调用 - 与我最近处理的几个 MLX 端口中实现的真实、可测量的优化(depth-anything-mlx:从 fp16 结合此融合获得的 1.35-1.5 倍;类似建议适用于 Blaizzy/mlx-vlm 的 Video-Depth-Anything 端口,#2180)的同一类融合(`mx.fast.scaled_dot_product_attention` / `mx.fast.layer_norm`)。 我尚未将此特定更改与您仓库的原始实现进行对比(我目前没有将其本地转换/端到端运行),因此我无法像在实际提出差异之前那样给出可测量的数值 - 将其标记为值得尝试的内容,而不是经过验证的修复。 如果这对您有用,我很乐意提交实际的 PR,或者很乐意关闭此问题,如果它已经被考虑并因为某些原因被排除(例如,在 fp32 和 fp16 之间的添加相对位置掩码的数值稳定性问题)。
内容来源: ml-explore/mlx-examples