tf.data.Dataset重复生成相同输出问题排查求助
tf.data.Dataset.from_generator 每次生成相同输出的问题解决
我在构建视频文件的推理数据管道时,使用tf.data.Dataset.from_generator创建批次数据集,但每次调用next(iter(inference_datasets))得到的文件集合都完全相同。
问题原因
- 生成器每次启动都从头遍历原始列表:
from_generator创建的数据集,每次调用iter()都会重新初始化生成器实例。当training=False时,生成器不会打乱视频路径列表,因此每次都会从列表第一个元素开始生成数据,导致每次获取的都是前BATCH_SIZE个文件。 - 列表引用导致的潜在bug:当前代码中
video_paths = self.video_paths直接引用了实例的原始列表,当training=True时,random.shuffle(video_paths)会直接修改原始列表,导致后续调用生成器时使用的是已打乱的列表,而非初始列表。
解决方案
方案1:保留迭代器实例,持续获取下一批数据
如果需要按顺序遍历所有视频,不要每次都创建新的迭代器,而是只初始化一次迭代器,后续每次调用next()获取下一批:
# 仅创建一次迭代器 infer_iter = iter(inference_datasets) # 每次调用获取下一批,直到遍历完所有数据 sample_inference_dataset = next(infer_iter) print(sample_inference_dataset[1])
方案2:修复生成器的列表副本问题
无论training是否为True,都应该使用视频路径列表的副本,避免修改原始列表:
class FrameGenerator: def __init__(self, video_paths, n_frames, training=False): self.video_paths = video_paths self.n_frames = n_frames self.training = training def __call__(self): # 创建列表副本,避免修改原始video_paths video_paths = list(self.video_paths) if self.training: random.shuffle(video_paths) for path in video_paths: video_frames = frames_from_video_file(path, self.n_frames) file_name = path.split('/')[-1] yield video_frames, file_name
方案3:推理时随机打乱数据(可选)
如果需要在推理时随机获取不同批次的视频,可以使用tf.data.Dataset.shuffle()方法打乱数据集:
dataset = tf.data.Dataset.from_generator( FrameGenerator(video_paths, 20, training=False), output_signature=output_signature ) # 设置buffer_size为数据集大小可完全打乱,根据内存情况调整 dataset = dataset.shuffle(buffer_size=len(video_paths)) inference_datasets = dataset.batch(BATCH_SIZE)
内容的提问来源于stack exchange,提问作者prynet
相关产品推荐
相关产品推荐

