【问题标题】:How to select a batch from a preloaded dataset in Tensorflow?如何从 Tensorflow 中的预加载数据集中选择一个批次?
【发布时间】:2018-11-18 00:14:31
【问题描述】:

我有一个可以放入 GPU 内存的大型数据集,我想在训练的每个步骤中从中选择一个随机批次。数据集由两个数组组成:

data1 = np.load("data1.npy")
data2 = np.load("data2.npy")

t_data1 = tf.constant(data1)
t_data2 = tf.constant(data2)

data1data2 的形状为 (16000, 200)。批量大小为 128,所以我想从每个数组中选择 128 个具有相同索引的元素,并将其提供给优化器:

for i in range(training_steps):
    choise = np.random.choice(data1.shape[0], batch_size)

    X_batch = t_data1[choise]
    Y_batch = t_data2[choise]

    sess.run(train_step, feed_dict={X: X_batch, Y: Y_batch})

很遗憾,我收到了这个错误:

ValueError: Shape must be rank 1 but is rank 2 for 'strided_slice' (op: 'StridedSlice') with input shapes

我做错了什么?如何从已经在 GPU 上的数据生成批处理?

【问题讨论】:

    标签: python python-3.x numpy tensorflow


    【解决方案1】:

    t_data1t_data2 是 tensorflow 中的常量。在numpy中可以这样做,但是tensorflow不支持高级索引,需要使用tf.gather()

    改为:

    X_batch = tf.gather(t_data1,choise)
    Y_batch = tf.gather(t_data2,choise)
    

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2020-03-20
      • 1970-01-01
      • 1970-01-01
      • 2021-12-20
      • 1970-01-01
      • 2018-09-17
      • 2020-11-01
      相关资源
      最近更新 更多