You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.24 08:06:04