मुख्य कंटेंट तक स्किप करें

समस्या विवरण

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'>
)

मीन (mean) की गणना करने वाला एक स्टैंडअलोन 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,
)

कमिट संदेश

fix: ak.mean को JAX के साथ विभेदक (differentiable) बनाएं

`ak.mean` के माध्यम से ट्रेस करते समय `jax.value_and_grad`,
`RuntimeError: Cannot differentiate through count_zero` देता है।
इसका मूल कारण यह है कि आंतरिक रूप से वज़न (weights) का योग (sum)
गणना करने के लिए उपयोग किए जाने वाले `ak.count` का कोई
JAX विभेदन नियम (differentiation rule) नहीं है।

`ak.count(x, ...)` को `ak.sum(x * 0 + 1, ...)` से बदलें, जो
समान संख्यात्मक परिणाम उत्पन्न करता है लेकिन JAX के तहत
पूरी तरह से विभेदक (differentiable) है। यही प्रतिस्थापन `ak.var`,
`ak.covar`, `ak.moment`, और `ak.linear_fit` पर भी लागू किया गया है।

Fixes #2595.

पुल रिक्वेस्ट

सारांश

उस समस्या का समाधान करता है जहाँ ak.mean को कॉल करने वाले फ़ंक्शन पर jax.value_and_grad का उपयोग करने पर त्रुटि संदेश उत्पन्न होता है। (Fixes #2595, #2591 का फॉलो-अप।)

समस्या

JAX के साथ ak.mean के माध्यम से अंतर करने पर विफलता मिलती है:

RuntimeError: Cannot differentiate through count_zero

यह त्रुटि कॉल करते समय उत्पन्न होती है:

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

जब कोई weight नहीं दिया जाता है, तो ak.mean, ak.count रिड्यूसर के माध्यम से वज़न (weights) का योग (sum) की गणना करता है। count (count_zero/count_nonzero कर्नेल के ऊपर लागू किया गया) में कोई JAX विभेदन नियम नहीं है, इसलिए ak.mean और समान पैटर्न का पालन करने वाले अन्य आँकड़ों (statistics) के संचालन के माध्यम से कोई भी ग्रेडिएंट (gradient), त्रुटि को ट्रिगर करता है।

समाधान

यह फिक्स गैर-विभेदक (non-differentiable) ak.count(x, ...) कॉल को उसके समतुल्य गणितीय ak.sum(x * 0 + 1, ...) से बदल देता है:

  • x * 0 + 1, सूची संरचना (list structure) और लापता मानों (missing values) को संरक्षित करते हुए, x के प्रत्येक तत्व (element) पर 1 का वज़न ब्रॉडकास्ट करता है।
  • ak.sum, JAX के तहत विभेदक है (चूंकि x * 0 के माध्यम से इसका ग्रेडिएंट योगदान शून्य है)।
  • संख्यात्मक परिणाम पिछले ak.count-आधारित कार्यान्वयन के समान है।

यही फिक्स उन सभी आँकड़ों के संचालन पर लागू किया गया है जो बिना वज़न वाले योग (unweighted sum-of-weights) के लिए 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

परीक्षण

  • समस्या रीप्रोड्यूसर (reproducer) पैटर्न अब काम करता है:
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 का अपेक्षित ग्रेडिएंट (gradient) और सही मान वापस कर रहा है।

  • डिफ़ॉल्ट (NumPy) बैकएंड पर सभी मौजूदा परीक्षण, ak.mean, ak.var, ak.std, ak.corr, ak.covar, ak.linear_fit, और ak.moment के लिए बिना किसी बदलाव के पास होते हैं, जिसमें axis=None, axis=-1, और अनियमित (ragged) इनपुट भी शामिल हैं।

नोट्स

  • कोई सार्वजनिक API परिवर्तन नहीं; मौजूदा कॉल सिग्नेचर सहेजे गए हैं।
  • कोई नई निर्भरता (dependencies) पेश नहीं की गई है।
  • प्रतिस्थापन x * 0 + 1 रनटाइम पर एक शून्य-लागत पहचान (zero-cost identity) है और कोई मापने योग्य ओवरहेड नहीं जोड़ता है।

कैसे परीक्षण करें

  1. समस्या से रीप्रोड्यूसर (reproducer) चलाएँ। jax.value_and_grad(lambda x: ak.mean(x))(arr) को RuntimeError के बिना मान और ग्रेडिएंट वापस करना चाहिए।
  2. पूर्ण परीक्षण सूट (pytest tests/) चलाएँ ताकि यह पुष्टि हो सके कि NumPy बैकएंड पर ak.mean, ak.var, ak.std, ak.covar, ak.corr, ak.moment, और ak.linear_fit समान रूप से व्यवहार करते हैं।