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

如何从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))

关键说明

  1. 重叠程度控制:通过stride参数调整,比如stride=50表示子序列间重叠50帧,stride=100表示无重叠。
  2. 性能保障:拆分逻辑在tf.data pipeline中执行,利用TF的并行处理能力,远快于train_step中的循环拆分。
  3. 元数据绑定:每个子序列会保留原长序列的元数据,符合GPT模型训练时的输入需求。
  4. 边界处理:自动过滤长度小于子序列长度的无效样本(可根据需求改为补零处理)。

内容的提问来源于stack exchange,提问作者Dr Sokoban

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 03:14:54