بيان المشكلة
عنوان المشكلة: التفاضل عبر 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.pysrc/awkward/operations/ak_var.py(يصلح أيضًاak.std، المبني علىak.var)src/awkward/operations/ak_covar.py(يصلح أيضًاak.corr)src/awkward/operations/ak_moment.pysrc/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عبارة عن هوية بتكلفة صفرية في وقت التشغيل ولا يضيف أي عبء يمكن قياسه.
كيفية الاختبار
- قم بتشغيل معيد الإنتاج الخاص بالمشكلة. يجب أن يُرجع
jax.value_and_grad(lambda x: ak.mean(x))(arr)قيمة وتدرجًا بدون رفعRuntimeError. - قم بتشغيل مجموعة الاختبار الكاملة (
pytest tests/) لتأكيد أنak.mean،ak.var،ak.std،ak.covar،ak.corr،ak.moment، وak.linear_fitتتصرف بشكل مماثل على واجهة NumPy الخلفية.