跳到主要内容

问题陈述

Issue 标题:通过 ak.mean 求导

Awkward Array 版本

main 分支

描述和重现代码

这是 #2591 的后续,配置略微简化。从概念上讲,应该可以通过取平均值(mean)来进行求导。目前这不起作用。

重现代码:

import awkward as ak
import jax
import uproot

ak.jax.register_and_check()

ttbar_file = "https://github.com/scikit-hep/scikit-hep-testdata/"\
"raw/main/src/skhep_testdata/data/nanoAOD_2015_CMS_Open_Data_ttbar.root"

def mean_jet_pt(jets):
return ak.mean(jets.pt)

with uproot.open(ttbar_file) as f:
arr = f["Events"].arrays(["Jet_pt","Jet_eta", "Jet_phi", "Jet_mass"])
evtfilter = ak.num(arr["Jet_pt"]) >= 2
jets = ak.zip(dict(zip(["pt","eta", "phi", "mass"], ak.unzip(arr))), with_name="Momentum4D")[evtfilter]
jets = ak.to_backend(jets, "jax")

jax.value_and_grad(mean_jet_pt, argnums=0)(jets)

结果:

RuntimeError: Cannot differentiate through count_zero

This error occurred while calling

ak.mean(
<Array [[...], [...], ..., [...], [...]] type='140 * var * float32'>
)

独立运行的计算平均值的 jax 版本运行良好:

import jax.numpy as jnp

def mean(j):
return jnp.mean(j)

data = jnp.array([1, 7, 3, 5],dtype=float)

jax.value_and_grad(mean, argnums=0)(data)

代码差异

src/awkward/operations/ak_covar.py

diff --git a/src/awkward/operations/ak_covar.py b/src/awkward/operations/ak_covar.py
index a070ac68..f0decdeb 100644
--- a/src/awkward/operations/ak_covar.py
+++ b/src/awkward/operations/ak_covar.py
@@ -102,52 +102,52 @@ def _impl(x, y, weight, axis, keepdims, mask_identity, highlevel, behavior, attr
y = ctx.wrap(y_layout)
weight = ctx.wrap(weight_layout, allow_other=True)

with np.errstate(invalid="ignore", divide="ignore"):
xmean = ak.operations.ak_mean._impl(
x, weight, axis, False, mask_identity,
highlevel=True, behavior=None, attrs=None,
)
ymean = ak.operations.ak_mean._impl(
y, weight, axis, False, mask_identity,
highlevel=True, behavior=None, attrs=None,
)
if weight is None:
- sumw = ak.operations.ak_count._impl(
- x,
+ sumw = ak.operations.ak_sum._impl(
+ x * 0 + 1,
axis, keepdims, mask_identity,
highlevel=True, behavior=None, attrs=None,
)
sumwxy = ak.operations.ak_sum._impl(
(x - xmean) * (y - ymean),
axis, keepdims, mask_identity,
highlevel=True, behavior=None, attrs=None,
)

提交信息 (Commit message)

fix: 使 ak.mean 可以在 JAX 中求导

当在 `ak.mean` 中进行追踪时,`jax.value_and_grad` 会引发
`RuntimeError: Cannot differentiate through count_zero`。根本原因是
内部用于计算权重总和的 `ak.count` 没有 JAX
求导规则。

将 `ak.count(x, ...)` 替换为 `ak.sum(x * 0 + 1, ...)`,这会
产生相同的数值结果,但在 JAX 下是完全可导的。相同的替换也适用于
`ak.var`、`ak.covar`、`ak.moment` 和 `ak.linear_fit`。

修复了 #2595。

合并请求 (Pull request)

摘要

解决了在调用 ak.mean 的函数上使用 jax.value_and_grad 会引发错误消息的问题。(修复了 #2595,这是 #2591 的后续。)

问题

在 JAX 中通过 ak.mean 进行求导时失败:

RuntimeError: Cannot differentiate through count_zero

此错误在调用以下内容时触发:

ak.mean(
<Array [[...], [...], ..., [...], [...]] type='140 * var * float32'>
)

当没有给定 weight 时,ak.mean 会通过 ak.count 归约器计算权重的总和。count(实现在 count_zero/count_nonzero 核之上)没有 JAX 求导规则,因此任何通过 ak.mean 以及遵循相同模式的其他统计操作的梯度,都会触发该错误。

解决方案

此修复将不可导的 ak.count(x, ...) 调用替换为其等效的数学表达式 ak.sum(x * 0 + 1, ...)

  • x * 0 + 1 会将权重 1 广播到 x 的每个元素上,保留列表结构和缺失值,
  • ak.sum 在 JAX 下是可导的(其通过 x * 0 的梯度贡献恒为零),
  • 数值结果与以前基于 ak.count 的实现相同。

相同的修复应用于所有使用 ak.count 进行未加权权重总和计算的统计操作:

  • src/awkward/operations/ak_mean.py
  • src/awkward/operations/ak_var.py(也修复了基于 ak.var 构建的 ak.std
  • src/awkward/operations/ak_covar.py(也修复了 ak.corr
  • src/awkward/operations/ak_moment.py
  • src/awkward/operations/ak_linear_fit.py

测试

  • 该问题的重现模式现在可以正常工作:
arr = ak.Array([[1.0, 2.0, 3.0], [4.0, 5.0], [6.0]], backend="jax")
jax.value_and_grad(lambda x: ak.mean(x))(arr)
# (Array(3.5, dtype=float32), <Array [[0.1667, 0.1667, 0.1667], ...]>)

返回正确的值和每个元素预期的 1/N 梯度。

  • 默认(NumPy)后端上的所有现有测试均顺利通过(未作更改),这包括 ak.meanak.varak.stdak.corrak.covarak.linear_fitak.moment,并且涵盖了 axis=Noneaxis=-1 以及不规则(ragged)输入。

备注

  • 公共 API 没有更改;现有的调用签名被保留。
  • 没有引入新的依赖项。
  • 替换 x * 0 + 1 是运行时的零成本恒等操作,不会增加可测量的开销。

如何测试

  1. 运行此 Issue 的重现代码。jax.value_and_grad(lambda x: ak.mean(x))(arr) 应该返回一个值和梯度,而不会引发 RuntimeError
  2. 运行完整的测试套件 (pytest tests/),以确认 ak.meanak.varak.stdak.covarak.corrak.momentak.linear_fit 在 NumPy 后端上的行为相同。