HIPS/autograd

support numpy.take

Ouverte

#743 ouverte le 24 nov. 2025

 (2 commentaires) (0 réaction) (0 personne assignée)Python (909 forks)batch import
PR welcomegood first issue

Métriques du dépôt

Stars
 (6 628 étoiles)
Métriques de merge PR
 (Métriques PR en attente)

Description

Currently tracing the gradient through take is not supported:

import numpy as np
import autograd as ag
import autograd.numpy as anp

rng = np.random.default_rng(42)
x = rng.uniform(size=(3, 4, 5))
idx = rng.integers(0, 4, size=(6,))

def foo(x, idx):
    # # works:
    # return x[:, idx, :].sum()
    # # doesn't work:
    return anp.take(x, idx, axis=1).sum()

gfoo = ag.grad(foo, argnum=0)
gfoo(x, rng.integers(0, 4, size=(6,))).shape
# NotImplementedError: VJP of take wrt argnums (0,) not defined

take is mostly just a convenient syntax/subset of getitem indexing, so I suppose there is nothing fundamentally blocking it.

Guide contributeur