【问题标题】:Inverse operation to padding in JaxJax 中填充的逆操作
【发布时间】:2021-10-05 15:49:21
【问题描述】:

我正在尝试学习如何使用 Jax,但偶然发现了将 torch.nn.functionnal.pad 函数转换为 Jax 的问题。有一个执行填充的函数,但我想以与 PyTorch 中相同的方式在填充中使用负数(例如 F.pad(array, [-1,-1]))。

有人有想法或有同样的问题吗?

【问题讨论】:

    标签: pytorch padding jax


    【解决方案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.]]
    

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2017-04-29
      • 1970-01-01
      • 2021-08-04
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2015-10-26
      相关资源
      最近更新 更多