HIPS/autograd

support numpy.take

Offen

#743 geöffnet am 24.11.2025

 (2 Kommentare) (0 Reaktionen) (0 zugewiesene Personen)Python (909 Forks)batch import
PR welcomegood first issue

Repository-Metriken

Stars
 (6.628 Sterne)
PR-Merge-Metriken
 (PR-Metriken ausstehend)

Beschreibung

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.

Contributor Guide