torch.compiler.wrap_numpy
-
torch.compiler.wrap_numpy(fn)[исходный код] -
Декоратор, преобразующий функцию из
np.ndarrays вnp.ndarrays в функцию изtorch.Tensors вtorch.Tensors.Предназначен для использования с
torch.compile()иfullgraph=True. Он позволяет компилировать функцию NumPy так, как если бы она была функцией PyTorch. Это позволяет запускать код NumPy на CUDA или вычислять его градиенты.Примечание
Этот декоратор не работает без
torch.compile().Пример:
>>> # Compile a NumPy function as a Tensor -> Tensor function >>> @torch.compile(fullgraph=True) >>> @torch.compiler.wrap_numpy >>> def fn(a: np.ndarray): >>> return np.sum(a * a) >>> # Execute the NumPy function using Tensors on CUDA and compute the gradients >>> x = torch.arange(6, dtype=torch.float32, device="cuda", requires_grad=True) >>> out = fn(x) >>> out.backward() >>> print(x.grad) tensor([ 0., 2., 4., 6., 8., 10.], device='cuda:0')
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.compiler.wrap_numpy.html