TF中Dataset使用window()后无法获取文件名及调用numpy报错如何解决?
问题1:调用.numpy()报错AttributeError: '_VariantDataset' object has no attribute 'numpy'
原因是tf.data.Dataset.window()返回的WindowDataset的每个元素是嵌套的数据集对象,并非直接存储张量,需要先通过打平操作将窗口转为连续的张量批次:
window_size = 3 # 对窗口数据集做打平处理,将每个窗口的3个元素打包为一个张量块 win_train_dataset = win_train_dataset.flat_map( lambda img_ds, label_ds: tf.data.Dataset.zip( (img_ds.batch(window_size), label_ds.batch(window_size)) ) ) # 调整后即可正常遍历读取张量 for imgs, labels in win_train_dataset: print(imgs.numpy().shape, labels.numpy().shape)
问题2:无法获取窗口内元素的文件名
tf.keras.utils.image_dataset_from_directory默认生成的数据集仅包含(图片张量、标签)二元组,没有存储文件路径信息,你需要自定义数据集构造逻辑,提前将路径信息嵌入数据集元素中:
import os import tensorflow as tf # 提前定义你的参数 class_names = ["class1", "class2"] # 替换为你的分类名 num_classes = len(class_names) window_size = 3 # 1. 读取所有图片路径并排序,保证同视频的帧连续排列 img_paths = sorted(tf.io.gfile.glob(os.path.join(training_data_dir, "*/*.jpg"))) # 替换为你的图片后缀 # 2. 自定义解析函数,返回值包含图片、标签、文件路径三个字段 def parse_img(path): # 从路径中提取分类标签,和image_dataset_from_directory逻辑对齐 label_str = tf.strings.split(path, os.path.sep)[-2] label = tf.one_hot(tf.argmax(label_str == class_names), depth=num_classes) # 解码并处理图片 img = tf.io.read_file(path) img = tf.image.decode_jpeg(img, channels=3) img = tf.image.resize(img, IMG_SIZE) return img, label, path # 生成带路径的无批次数据集 unb_training_dataset = tf.data.Dataset.from_tensor_slices(img_paths)\ .map(parse_img, num_parallel_calls=tf.data.AUTOTUNE) # 3. 生成滑动窗口并打平,保留路径字段 win_train_dataset = unb_training_dataset.window( window_size, shift=1, stride=1, drop_remainder=True ) win_train_dataset = win_train_dataset.flat_map( lambda img_ds, label_ds, path_ds: tf.data.Dataset.zip( (img_ds.batch(window_size), label_ds.batch(window_size), path_ds.batch(window_size)) ) ) # 4. 过滤跨视频的窗口 def filter_same_video(imgs, labels, paths): # 从路径中提取视频标识,按需调整提取规则:示例中文件名格式为video1_frame001.jpg,提取下划线前的video1作为视频ID video_ids = tf.strings.split(paths, "_")[:, 0] # 窗口内所有帧属于同一视频才保留 return tf.reduce_all(video_ids == video_ids[0]) win_train_dataset = win_train_dataset.filter(filter_same_video) # 最终训练时可丢弃路径字段,只保留模型需要的图片和标签 win_train_dataset = win_train_dataset.map( lambda imgs, labels, paths: (imgs, labels) ).batch(hyperparameters["BATCH_SIZE"])
内容的提问来源于stack exchange,提问作者ghylander
相关产品推荐
相关产品推荐

