Saltar al contenido principal

Planteamiento del problema

Título del issue: Diferenciación a través de un ak.mean

Versión de Awkward Array

rama main

Descripción y código para reproducir

Esta es una continuación del #2591 con una configuración un poco más simplificada. Conceptualmente debería ser posible diferenciar calculando una media. Actualmente esto no funciona.

Reproductor:

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)

Resultado:

RuntimeError: Cannot differentiate through count_zero

This error occurred while calling

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

Una versión independiente de jax calculando una media funciona 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)

Diferencia de código

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

Mensaje de commit

fix: hacer que ak.mean sea diferenciable con JAX

`jax.value_and_grad` lanza `RuntimeError: Cannot differentiate through
count_zero` cuando rastrea a través de `ak.mean`. La causa principal es que
`ak.count`, utilizado internamente para calcular la suma de pesos, no tiene una
regla de diferenciación en JAX.

Reemplazar `ak.count(x, ...)` con `ak.sum(x * 0 + 1, ...)`, lo cual
produce el mismo resultado numérico pero es completamente diferenciable bajo
JAX. La misma sustitución se aplica a `ak.var`, `ak.covar`,
`ak.moment` y `ak.linear_fit`.

Soluciona #2595.

Pull Request

Resumen

Resuelve el problema donde el uso de jax.value_and_grad en una función que llama a ak.mean genera un mensaje de error. (Soluciona #2595, continuación del #2591.)

Problema

La diferenciación a través de ak.mean con JAX falla con:

RuntimeError: Cannot differentiate through count_zero

Este error se desencadena al llamar a:

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

Cuando no se proporciona ningún weight, ak.mean calcula la suma de los pesos a través del reductor ak.count. count (implementado sobre los kernels count_zero/count_nonzero) no tiene regla de diferenciación en JAX, por lo que cualquier gradiente a través de ak.mean, y a través de las otras operaciones estadísticas que siguen el mismo patrón, desencadena el error.

Solución

La corrección reemplaza la llamada no diferenciable ak.count(x, ...) por su equivalente matemático ak.sum(x * 0 + 1, ...):

  • x * 0 + 1 transmite un peso de 1 a cada elemento de x, preservando la estructura de la lista y los valores faltantes,
  • ak.sum es diferenciable bajo JAX (su contribución al gradiente a través de x * 0 es idénticamente cero),
  • El resultado numérico es idéntico a la implementación anterior basada en ak.count.

La misma corrección se aplica a todas las operaciones estadísticas que usaban ak.count para la suma de pesos no ponderados:

  • src/awkward/operations/ak_mean.py
  • src/awkward/operations/ak_var.py (también soluciona ak.std, que se construye sobre ak.var)
  • src/awkward/operations/ak_covar.py (también soluciona ak.corr)
  • src/awkward/operations/ak_moment.py
  • src/awkward/operations/ak_linear_fit.py

Pruebas

  • El patrón del reproductor del problema ahora funciona:
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], ...]>)

devolviendo el valor correcto y el gradiente esperado de 1/N por elemento.

  • Todas las pruebas existentes en el backend predeterminado (NumPy) pasan sin cambios para ak.mean, ak.var, ak.std, ak.corr, ak.covar, ak.linear_fit y ak.moment, incluyendo axis=None, axis=-1 y entradas irregulares (ragged).

Notas

  • No hay cambios en la API pública; las firmas de las llamadas existentes se conservan.
  • No se introducen nuevas dependencias.
  • La sustitución x * 0 + 1 es una identidad de costo cero en tiempo de ejecución y no agrega una sobrecarga medible.

Cómo probar

  1. Ejecute el reproductor del problema de GitHub. jax.value_and_grad(lambda x: ak.mean(x))(arr) debería devolver un valor y un gradiente sin generar un RuntimeError.
  2. Ejecute el conjunto de pruebas completo (pytest tests/) para confirmar que ak.mean, ak.var, ak.std, ak.covar, ak.corr, ak.moment y ak.linear_fit se comportan de forma idéntica en el backend NumPy.