إنتقل إلى المحتوى الرئيسي

بيان المشكلة

عنوان المشكلة: التفاضل عبر 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)

تغييرات الكود (Code Diff)

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

يطرح `jax.value_and_grad` خطأ `RuntimeError: Cannot differentiate through
count_zero` عند التتبع عبر `ak.mean`. السبب الجذري هو أن
`ak.count`، المستخدم داخليًا لحساب مجموع الأوزان، لا يمتلك قاعدة
تفاضل في JAX.

استبدل `ak.count(x, ...)` بـ `ak.sum(x * 0 + 1, ...)`، والذي
ينتج نفس النتيجة الرقمية ولكنه قابل للتفاضل بالكامل تحت
JAX. يتم تطبيق نفس الاستبدال على `ak.var`، `ak.covar`،
`ak.moment`، و `ak.linear_fit`.

يصلح #2595.

طلب السحب (Pull Request)

ملخص

يحل المشكلة حيث يثير استخدام jax.value_and_grad في دالة تستدعي ak.mean رسالة خطأ. (يصلح #2595، متابعة لـ #2591.)

المشكلة

يفشل التفاضل عبر ak.mean باستخدام JAX مع:

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.std، المبني على ak.var)
  • 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.mean، ak.var، ak.std، ak.corr، ak.covar، ak.linear_fit، و ak.moment، بما في ذلك axis=None، axis=-1، والمدخلات غير المنتظمة (ragged).

ملاحظات

  • لا توجد تغييرات في واجهة برمجة التطبيقات (API) العامة؛ يتم الاحتفاظ بتوقيعات الاستدعاء الحالية.
  • لم يتم إدخال تبعيات جديدة.
  • الاستبدال x * 0 + 1 عبارة عن هوية بتكلفة صفرية في وقت التشغيل ولا يضيف أي عبء يمكن قياسه.

كيفية الاختبار

  1. قم بتشغيل معيد الإنتاج الخاص بالمشكلة. يجب أن يُرجع jax.value_and_grad(lambda x: ak.mean(x))(arr) قيمة وتدرجًا بدون رفع RuntimeError.
  2. قم بتشغيل مجموعة الاختبار الكاملة (pytest tests/) لتأكيد أن ak.mean، ak.var، ak.std، ak.covar، ak.corr، ak.moment، و ak.linear_fit تتصرف بشكل مماثل على واجهة NumPy الخلفية.