Aller au contenu principal

Énoncé du problème

Titre du problème : Différenciation à travers un ak.mean

Version d'Awkward Array

branche main

Description et code pour reproduire

Ceci est une suite du #2591 avec une configuration légèrement simplifiée. Conceptuellement, il devrait être possible de différencier en prenant une moyenne. Actuellement, cela ne fonctionne pas.

Reproducteur :

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)

Résultat :

RuntimeError: Cannot differentiate through count_zero

This error occurred while calling

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

Une version autonome de JAX effectuant une moyenne fonctionne très bien :

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)

Différences de code

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

Message de validation (Commit)

fix: rendre ak.mean différentiable avec JAX

`jax.value_and_grad` lève `RuntimeError: Cannot differentiate through
count_zero` lors du traçage à travers `ak.mean`. La cause profonde est que
`ak.count`, utilisé en interne pour calculer la somme des poids, n'a pas de
règle de différenciation dans JAX.

Remplacez `ak.count(x, ...)` par `ak.sum(x * 0 + 1, ...)`, ce qui
produit le même résultat numérique mais est entièrement différentiable sous
JAX. La même substitution s'applique à `ak.var`, `ak.covar`,
`ak.moment` et `ak.linear_fit`.

Corrige #2595.

Demande d'extraction (Pull request)

Résumé

Résout le problème où l'utilisation de jax.value_and_grad sur une fonction appelant ak.mean génère un message d'erreur. (Corrige #2595, suite de #2591.)

Problème

La différenciation à travers ak.mean avec JAX échoue avec :

RuntimeError: Cannot differentiate through count_zero

Cette erreur se déclenche lors de l'appel :

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

Lorsqu'aucun weight n'est fourni, ak.mean calcule la somme des poids via le réducteur ak.count. count (implémenté au-dessus des noyaux count_zero/count_nonzero) n'a pas de règle de différenciation JAX, donc tout gradient à travers ak.mean, et à travers les autres opérations statistiques qui suivent le même modèle, déclenche l'erreur.

Solution

Le correctif remplace l'appel non différentiable ak.count(x, ...) par son équivalent mathématique ak.sum(x * 0 + 1, ...) :

  • x * 0 + 1 diffuse un poids de 1 sur chaque élément de x, préservant la structure de la liste et les valeurs manquantes,
  • ak.sum est différentiable sous JAX (sa contribution au gradient à travers x * 0 est identiquement nulle),
  • Le résultat numérique est identique à la mise en œuvre précédente basée sur ak.count.

La même correction est appliquée à toutes les opérations statistiques qui utilisaient ak.count pour la somme des poids non pondérée :

  • src/awkward/operations/ak_mean.py
  • src/awkward/operations/ak_var.py (corrige également ak.std, qui s'appuie sur ak.var)
  • src/awkward/operations/ak_covar.py (corrige également ak.corr)
  • src/awkward/operations/ak_moment.py
  • src/awkward/operations/ak_linear_fit.py

Tests

  • Le motif reproducteur du problème fonctionne désormais :
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], ...]>)

renvoyant la valeur correcte et le gradient attendu de 1/N par élément.

  • Tous les tests existants sur le backend par défaut (NumPy) passent sans changement pour ak.mean, ak.var, ak.std, ak.corr, ak.covar, ak.linear_fit et ak.moment, y compris axis=None, axis=-1 et les entrées irrégulières (ragged).

Remarques

  • Aucune modification de l'API publique ; les signatures d'appel existantes sont préservées.
  • Aucune nouvelle dépendance n'est introduite.
  • La substitution x * 0 + 1 est une identité à coût nul lors de l'exécution et n'ajoute aucune surcharge mesurable.

Comment tester

  1. Exécutez le reproducteur du problème. jax.value_and_grad(lambda x: ak.mean(x))(arr) devrait renvoyer une valeur et un gradient sans générer de RuntimeError.
  2. Exécutez la suite de tests complète (pytest tests/) pour confirmer que ak.mean, ak.var, ak.std, ak.covar, ak.corr, ak.moment et ak.linear_fit se comportent de manière identique sur le backend NumPy.