HIPS/autograd

support numpy.take

オープン

#743 opened on 2025/11/24

 (2 件のコメント) (0 件のリアクション) (0 人の担当者)Python (909 件のフォーク)batch import
PR welcomegood first issue

Repository metrics

Stars
 (6,628 個のスター)
PR merge metrics
 (PR metrics pending)

説明

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.

コントリビューターガイド