如何从形状(890,2048,3)的NumPy数组创建tf.Dataset并验证
从NumPy数组创建TensorFlow Dataset及验证方法
一、正确创建Dataset
你的代码其实已经成功创建了TensorSliceDataset,输出的element_spec是TensorFlow数据集的正常属性——它明确告诉你数据集里每个元素的结构:每个元素是一个元组,包含两个形状为(2048, 3)的float64张量,分别对应一组源点云和目标点云。
完整创建代码如下:
import tensorflow as tf import numpy as np # 假设source和targ是你已有的(890,2048,3)格式NumPy数组 data = tf.data.Dataset.from_tensor_slices((source, targ))
二、验证Dataset是否创建成功
可以通过以下几种方式验证数据集的正确性:
1. 确认样本总数
用cardinality()方法查看数据集的样本数量,结果应与你的890组数据一致:
print("数据集样本数:", data.cardinality().numpy())
2. 对比单个样本与原NumPy数据
取出数据集的第一个样本,和原NumPy数组的对应数据做一致性校验:
# 获取第一个样本 first_source, first_targ = next(iter(data)) # 与原NumPy数组的第一个样本对比 np.testing.assert_allclose(first_source.numpy(), source[0], rtol=1e-6) np.testing.assert_allclose(first_targ.numpy(), targ[0], rtol=1e-6) print("第一个样本数据匹配成功")
3. 遍历查看元素结构
遍历前几个样本,确认每个元素的形状、数据类型符合预期:
for idx, (src, tgt) in enumerate(data.take(3)): print(f"第{idx+1}个样本:") print(f"源点云形状:{src.shape},数据类型:{src.dtype}") print(f"目标点云形状:{tgt.shape},数据类型:{tgt.dtype}")
4. 测试预处理操作(可选)
如果后续需要做打乱、分批等预处理,可以测试这些操作是否正常生效:
# 打乱数据集并按批次划分(示例batch_size=8) batch_data = data.shuffle(buffer_size=890).batch(batch_size=8) # 查看批次后的张量形状 for batch_src, batch_tgt in batch_data.take(1): print(f"批次源点云形状:{batch_src.shape}") # 应为(8,2048,3) print(f"批次目标点云形状:{batch_tgt.shape}") # 应为(8,2048,3)
补充说明
你之前看到的<TensorSliceDataset element_spec=...>不是错误,而是TensorFlow数据集的默认打印格式,它仅用于展示数据集的元素规格,说明你的数据集创建已经成功。如果需要查看具体数据,必须通过迭代器取出元素才能查看。
内容的提问来源于stack exchange,提问作者willie_1997
相关产品推荐
相关产品推荐

