如何从TFRecord的视频Latent序列生成重叠训练子序列?
解决方案:在TF Dataset Pipeline中拆分长序列为重叠子序列
无需重写TFRecord文件,直接在数据读取的Pipeline中完成长序列到重叠子序列的拆分,利用tf.data的并行处理能力保证效率。
核心修改思路
针对每个TFRecord条目(包含830×200的Latent序列和4元素元数据),通过flat_map结合自定义拆分逻辑,生成多个长度为100的重叠子序列,同时保持元数据与每个子序列绑定。
修改后的TFRecordReader代码
在原TFRecordReader类中添加拆分函数,并修改get_dataset方法:
class TFRecordReader(object): """ Class handling reading the TFRecord into a dataset to use for training Please check TFRecordWriter to see how it was saved to disk. The TFRecord to be read must have been created with TFRecordWriter """ def __init__(self, tfrecordfile, batch_size = 64, subseq_len=100, stride=50): self.BATCH_SIZE = batch_size self.tfrecordfile = tfrecordfile self.AUTOTUNE = tf.data.AUTOTUNE self.subseq_len = subseq_len # 子序列长度 self.stride = stride # 滑动步长,控制重叠程度 self.dataset = self.get_dataset() # this will set self.dataset self.dataset_iter = iter(self.dataset) def decode_frames(self, frames): parsed_data = tf.io.parse_tensor(frames, tf.float32) parsed_data = tf.reshape(parsed_data, [832, 200]) # explicit size needed for TPU return parsed_data def read_tfrecord(self, example): TFREC_FORMAT = { "frames": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring "recipe": tf.io.FixedLenFeature([4], tf.float32) } example = tf.io.parse_single_example(example, TFREC_FORMAT) video_latent_vectors = self.decode_frames(example['frames']) return video_latent_vectors, example['recipe'] def load_dataset(self): """ Loads a TFRecord and uses map to parse it, and stores it into self.dataset Check https://keras.io/examples/keras_recipes/tfrecord/ "define load methods" because this is basically a copy paste of that code with small modifications Args: properties (list, optional): Check parse_fn above Returns: dataset: Loadad TFRecord """ ignore_order = tf.data.Options() ignore_order.experimental_deterministic = False # disable order, increase speed dataset = tf.data.TFRecordDataset( self.tfrecordfile ) # automatically interleaves reads from multiple files dataset = dataset.with_options( ignore_order ) # uses data as soon as it streams in, rather than in its original order dataset = dataset.map( self.read_tfrecord, num_parallel_calls=self.AUTOTUNE ) # returns the dataset as loaded return dataset def split_into_overlapping_subsequences(self, latent_seq, recipe): """将单条长序列拆分为重叠子序列,同时保留元数据""" seq_len = tf.shape(latent_seq)[0] # 计算有效子序列数量(确保最后一个子序列不越界) max_start = seq_len - self.subseq_len if max_start <= 0: # 如果序列长度小于子序列长度,直接返回空(或根据需求pad) return tf.data.Dataset.from_tensor_slices([]) num_subseqs = tf.cast(tf.math.ceil((max_start) / self.stride) + 1, tf.int32) starts = tf.range(0, num_subseqs * self.stride, self.stride) starts = tf.clip_by_value(starts, 0, max_start) # 生成每个子序列和对应的元数据 def get_single_subseq(start): end = start + self.subseq_len return latent_seq[start:end], recipe return tf.data.Dataset.from_tensor_slices(starts).map(get_single_subseq) def get_dataset(self): """Loads the TFRecord from the paths (filenames), and then shuffles the data and divides it into batches. """ dataset = self.load_dataset() # 拆分每个长序列为重叠子序列 dataset = dataset.flat_map(lambda seq, recipe: self.split_into_overlapping_subsequences(seq, recipe)) dataset = dataset.shuffle(2048) dataset = dataset.prefetch(buffer_size=self.AUTOTUNE) dataset = dataset.batch(self.BATCH_SIZE, drop_remainder=True) return dataset # .repeat() def visualise_latent_reconstructions_and_recipes(self, vaepath): batch_size = 32 data = next(self.dataset_iter)[0] # returns 576,200 (or whatever latent size) ds = tf.data.Dataset.from_tensor_slices(data.numpy()[0]) ds = ds.batch(batch_size) # returns for example 9,32,200 vae, _ = load_vae_model(vaepath) for entry in ds.take(1): generated_images = vae.decoder(entry) for i in range(batch_size): img = utils.array_to_img(generated_images[i]) img.save("reader_img_%03d.png" % (i))
关键说明
- 重叠程度控制:通过
stride参数调整,比如stride=50表示子序列间重叠50帧,stride=100表示无重叠。 - 性能保障:拆分逻辑在
tf.datapipeline中执行,利用TF的并行处理能力,远快于train_step中的循环拆分。 - 元数据绑定:每个子序列会保留原长序列的元数据,符合GPT模型训练时的输入需求。
- 边界处理:自动过滤长度小于子序列长度的无效样本(可根据需求改为补零处理)。
内容的提问来源于stack exchange,提问作者Dr Sokoban
相关产品推荐
相关产品推荐

