【问题标题】:Google JAX 1D convolutional neural networkGoogle JAX 1D 卷积神经网络
【发布时间】:2020-06-13 08:17:55
【问题描述】:

我正在尝试使用 stax.GeneralConv() (https://jax.readthedocs.io/en/latest/_modules/jax/experimental/stax.html#GeneralConv) 在 Google Jax 中实现一维卷积神经网络。 我有一个包含 18 个条目的一维输入数组和一个包含 6 个条目的输出数组。我想实现一个内核宽度为 3 的 CNN,如下所示:

init_random_params, conv_net = stax.serial(
    GeneralConv(('NC','IO','NC'),1,(3,),padding='SAME'), # dimension_numbers = ('NC','IO','NC')
    LogSoftmax,
    Dense(6),
)

带有初始网络参数:

rng = jax.random.PRNGKey(0)
_, init_params = init_random_params(rng, (18,))

但我收到以下错误:

stax.py", line 75, in <listcomp>
    next(filter_shape_iter) for c in rhs_spec]

IndexError: tuple index out of range

stax 要求维度编号 rhs_spec 至少为 2 个字符长,但我使用一维过滤器。有人知道如何解决这个问题吗?

【问题讨论】:

    标签: python conv-neural-network jax


    【解决方案1】:

    我自己没有尝试过,但我希望一维卷积仍然需要一个方向来进行卷积,例如

    Conv2d = functools.partial(GeneralConv, ('NHWC', 'HWIO', 'NHWC'))
    Conv1d = functools.partial(GeneralConv, ('NHC', 'HIO', 'NHC'))
    

    换句话说,删除W 轴以从 2d 到 1d 卷积。

    NHC对应的输入shape是(batch_size, sequence_length, num_channels)

    请注意,即使通道数可能为 1,您仍然需要包含该轴,因为 GeneralConv 会沿着 num_channels = input_shape['NHC'.index('C')] 的行进行索引查找。

    【讨论】:

      猜你喜欢
      • 2020-11-01
      • 1970-01-01
      • 2017-01-01
      • 1970-01-01
      • 2021-12-16
      • 1970-01-01
      • 2020-10-07
      • 2020-10-19
      相关资源
      最近更新 更多