M1 Mac下tf.sort对长度超16的tf.float32张量排序末尾值变为-0
问题原因
这是M1芯片早期适配版本TensorFlow(v2.5.0)的Metal加速内核存在的已知bug:
- TensorFlow的排序实现对长度<=16的短数组有单独的CPU优化分支,逻辑没有问题;当数组长度>16时会自动调用M1的GPU(Metal)加速排序内核,该内核的float32版本存在边界访问错误,会误写最后一位元素的符号位,导致正常数值变成-0。
- float64的排序内核实现逻辑不同,没有触发这个边界错误,所以运行完全正常。
- numpy排序默认走CPU实现,没有调用Metal加速内核,因此也不会出现该问题。
是否需要提交bug报告
不需要单独提交:这个问题是早期M1适配版本的通用问题,在后续的TensorFlow版本(v2.10及以上的官方苹果适配版tensorflow-macos)中已经被修复,你可以先升级TensorFlow版本验证,升级后即可解决该问题。
临时解决方案
如果你暂时无法升级TensorFlow版本,可以用以下两种方法规避:
- 排序前将数组转为float64,排序完成后再转回float32:
a_sorted = tf.cast(tf.sort(tf.cast(a, tf.float64)), tf.float32)
如果需要完全消除-0的影响,可以额外加一步符号修正:
a_sorted = tf.where(a_sorted == 0, 0.0, a_sorted)
- 用
tf.math.top_k间接实现排序,该接口的Metal内核没有该bug:
# 升序排序实现 a_sorted = tf.reverse(tf.math.top_k(a, k=tf.shape(a)[0]).values, axis=[-1])
内容的提问来源于stack exchange,提问作者Marshall Mueller
相关产品推荐
相关产品推荐

