【发布时间】:2021-10-05 15:49:21
【问题描述】:
我正在尝试学习如何使用 Jax,但偶然发现了将 torch.nn.functionnal.pad 函数转换为 Jax 的问题。有一个执行填充的函数,但我想以与 PyTorch 中相同的方式在填充中使用负数(例如 F.pad(array, [-1,-1]))。
有人有想法或有同样的问题吗?
【问题讨论】:
我正在尝试学习如何使用 Jax,但偶然发现了将 torch.nn.functionnal.pad 函数转换为 Jax 的问题。有一个执行填充的函数,但我想以与 PyTorch 中相同的方式在填充中使用负数(例如 F.pad(array, [-1,-1]))。
有人有想法或有同样的问题吗?
【问题讨论】:
jax.lax.pad 函数接受负填充索引,尽管 API 与 torch.nn.functional.pad 的 API 略有不同。例如:
from jax import lax
import jax.numpy as jnp
x = jnp.ones((2, 3))
y = lax.pad(x, padding_config=[(0, 0, 0), (1, 1, 0)], padding_value=0.0)
print(y)
# [[0. 1. 1. 1. 0.]
# [0. 1. 1. 1. 0.]]
x = lax.pad(y, padding_config=[(0, 0, 0), (-1, -1, 0)], padding_value=0.0)
print(x)
# [[1. 1. 1.]
# [1. 1. 1.]]
如果您愿意,您可以使用与 torch 版本具有相似语义的函数来包装它。这是一个快速的尝试:
def jax_pad(input, pad, mode='constant', value=0):
"""JAX implementation of torch.nn.functional.pad
Warning: this has not been thoroughly tested!
"""
if mode != 'constant':
raise NotImplementedError("Only mode='constant' is implemented")
assert len(pad) % 2 == 0
assert len(pad) // 2 <= input.ndim
pad = list(zip(*[iter(pad)]*2))
pad += [(0, 0)] * (input.ndim - len(pad))
return lax.pad(
input,
padding_config=[(i, j, 0) for i, j in pad[::-1]],
padding_value=jnp.array(value, input.dtype))
x = jnp.ones((2, 3))
y = jax_pad(x, (1, 1))
print(y)
# [[0. 1. 1. 1. 0.]
# [0. 1. 1. 1. 0.]]
x = jax_pad(y, (-1, -1))
print(x)
# [[1. 1. 1.]
# [1. 1. 1.]]
【讨论】: