使用tf.data.Dataset.batch遇问题:请教该方法正确使用细节
Hey there! Let's break down why your batch operation didn't work as expected and get it sorted out properly.
The Core Issue
The root problem is how you created your "dataset": when you use tf.convert_to_tensor(), you end up with a Tensor object, not a tf.data.Dataset. The batch() method you called isn't the tf.data.Dataset.batch() you're expecting—it's a tensor-specific method that behaves differently, and when you tried to iterate over it, you got the full tensor instead of a sliced batch.
Correct Approach Step-by-Step
To use tf.data.Dataset.batch() as intended, you need to first convert your tensor into a tf.data.Dataset that treats each row as an individual sample. Here's how to do it:
- Create your raw tensor (this part you had right):
import tensorflow as tf data_tensor = tf.convert_to_tensor([ [5.1, 3.3, 1.7, 0.5], [5.9, 3.0, 4.2, 1.5], [6.9, 3.1, 5.4, 2.1], [2.3, 1.3, 6.4, 9.3] ])
- Convert the tensor to a tf.data.Dataset
Usetf.data.Dataset.from_tensor_slices()—this splits the tensor along its first dimension, turning each row into a separate element in the dataset:
dataset = tf.data.Dataset.from_tensor_slices(data_tensor)
- Apply the batch operation
Now when you callbatch(2), it will group 2 consecutive samples into a single batch:
dataset = dataset.batch(2)
- Iterate to get your batch
In TensorFlow 2.x (which uses eager execution by default), you can simply grab the first batch like this:
x = next(iter(dataset)) print(x.shape) # Output: (2, 4) — exactly what you expected!
Why Your Original Code Failed
When you called batch(2) directly on the tensor, you were using tf.Tensor.batch(), which splits the tensor into chunks along a specified axis (default is axis 0). However, since you tried to iterate over the tensor as if it were a tf.data.Dataset, the iterator returned the entire tensor as a single element—hence why you got the full 4×4 array instead of a batch.
Quick Note on Dataset Creation
Avoid using tf.data.Dataset.from_tensors() here—this would wrap the entire 4×4 tensor as one element in the dataset, so batch(2) wouldn't split it into smaller samples (it would just return the single element if drop_remainder=False).
内容的提问来源于stack exchange,提问作者Garen Zhou

