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

如何在PyTorch生态中集成基于Token总数的动态Batch加载?

基于总Token数限制的动态Batch PyTorch集成方案

问题描述

正在重构代码为规范的PyTorch流水线,使用Dataset、DataLoader、collate函数和samplers。当前数据集以句子为样本,每个样本的token数可通过sample.split()获取,示例数据集如下:

from random import randint
from torch.utils.data import Dataset

class DummyDataset(Dataset):
    def __init__(self):
        data = []
        for _ in range(128):
            data.append("hello " * randint(64, 176))
        self.data = data

    def __len__(self):
        return len(self.data)

    def __getitem__(self, idx: int):
        return self.data[idx]

已实现每个batch总token数不超过250的核心逻辑(不同batch的样本数量动态变化),但需要将该逻辑集成到PyTorch生态中,使其能像常规DataLoader一样工作。同时想了解torchtext的Iterator和batch_size_fn是否能解决该问题。

核心逻辑代码:

if __name__ == '__main__':
    dataset = DummyDataset()

    def get_batch(max_tokens: int = 250):
        data_idxs = list(range(len(dataset)))

        batch = []
        total_batch_len = 0
        while data_idxs:
            sample = dataset[data_idxs[0]]
            sample_len = len(sample.split())

            if total_batch_len + sample_len <= max_tokens:
                batch.append(sample)
                total_batch_len += sample_len
                data_idxs.pop(0)
            elif batch:
                yield batch
                batch = []
                total_batch_len = 0

        yield batch

    # 验证逻辑:确保所有样本都被处理
    num_samples = 0
    num_batches = 0
    for b in get_batch():
        num_samples += len(b)
        num_batches += 1

    print(f"Created {num_batches} batches")
    assert num_samples == len(dataset)

解决方案一:自定义Sampler(PyTorch原生推荐方案)

最规范的方式是自定义Sampler子类,生成按总token数分组的索引批次,直接对接DataLoader,完全兼容PyTorch流水线的所有特性(多进程加载、shuffle、collate函数等)。

实现步骤

  1. 预计算所有样本的token长度,避免重复计算
  2. 实现Sampler子类,按总token数不超过阈值的规则生成批次索引
  3. 将自定义Sampler传入DataLoader,并设置batch_size=None(由Sampler控制批次大小)

代码示例

from torch.utils.data import Sampler, DataLoader
import random

class TokenBatchSampler(Sampler):
    def __init__(self, dataset, max_tokens=250, shuffle=True):
        self.dataset = dataset
        self.max_tokens = max_tokens
        self.shuffle = shuffle
        # 预计算每个样本的token长度
        self.sample_lengths = [len(sample.split()) for sample in dataset]

    def __iter__(self):
        # 生成样本索引列表,shuffle模式下打乱顺序
        indices = list(range(len(self.dataset)))
        if self.shuffle:
            random.shuffle(indices)

        current_batch = []
        current_total_tokens = 0

        for idx in indices:
            sample_len = self.sample_lengths[idx]
            # 当前样本加入后不超过阈值则加入batch
            if current_total_tokens + sample_len <= self.max_tokens:
                current_batch.append(idx)
                current_total_tokens += sample_len
            else:
                # 输出当前batch,重新初始化新batch
                if current_batch:
                    yield current_batch
                current_batch = [idx]
                current_total_tokens = sample_len
        # 输出最后一个非空batch
        if current_batch:
            yield current_batch

    def __len__(self):
        # 估算batch数量(实际可能有细微出入)
        total_tokens = sum(self.sample_lengths)
        return (total_tokens + self.max_tokens - 1) // self.max_tokens

# 使用示例
if __name__ == '__main__':
    dataset = DummyDataset()
    sampler = TokenBatchSampler(dataset, max_tokens=250, shuffle=True)
    dataloader = DataLoader(dataset, batch_sampler=sampler)

    # 验证逻辑
    num_samples = 0
    num_batches = 0
    for batch in dataloader:
        num_samples += len(batch)
        total_tokens = sum(len(s.split()) for s in batch)
        assert total_tokens <= 250, f"Batch exceeds token limit: {total_tokens}"
        num_batches += 1

    print(f"Created {num_batches} batches")
    assert num_samples == len(dataset)

优势

  • 完全兼容PyTorch生态,支持所有DataLoader特性
  • 逻辑清晰,易于维护和扩展(如添加padding优化、过滤超长样本等)
  • 支持shuffle,满足训练场景需求

解决方案二:使用torchtext的Iterator(可选)

torchtext的Iterator(或BucketIterator)支持通过batch_size_fn动态控制batch大小,但需注意torchtext的API在新版本中变化较大,且需要适配其数据格式。

实现步骤

  1. 定义batch_size_fn函数,用于计算当前batch的总token数
  2. 将PyTorch Dataset转换为torchtext兼容的格式
  3. 初始化Iterator并传入batch_size_fn和max_tokens参数

代码示例(基于torchtext旧版本)

from torchtext.data import Iterator, Dataset as TorchtextDataset, Field

# 定义文本处理Field
TEXT = Field(tokenize=str.split)

# 将PyTorch Dataset转换为torchtext Dataset
def convert_to_torchtext_dataset(pytorch_dataset):
    examples = [TEXT.preprocess(sample) for sample in pytorch_dataset]
    return TorchtextDataset(examples, fields=[('text', TEXT)])

# 定义batch_size_fn:计算当前batch的总token数
def batch_size_fn(new, count, sofar):
    # new为当前样本的token列表,sofar为当前batch已有的总token数
    return sofar + len(new)

# 使用示例
if __name__ == '__main__':
    pytorch_dataset = DummyDataset()
    tt_dataset = convert_to_torchtext_dataset(pytorch_dataset)
    TEXT.build_vocab(tt_dataset)

    # 初始化Iterator,通过batch_size_fn控制总token数
    iterator = Iterator(
        tt_dataset,
        batch_size=1,  # 初始batch_size设为1,由batch_size_fn控制实际大小
        batch_size_fn=batch_size_fn,
        train=True,
        shuffle=True,
        max_tokens=250
    )

    # 验证逻辑
    num_samples = 0
    num_batches = 0
    for batch in iterator:
        num_samples += batch.text.size(1)  # batch.text维度为(seq_len, batch_size)
        total_tokens = batch.text.numel()
        assert total_tokens <= 250, f"Batch exceeds token limit: {total_tokens}"
        num_batches += 1

    print(f"Created {num_batches} batches")
    assert num_samples == len(pytorch_dataset)

注意事项

  • torchtext新版本已整合到torchdata中,API存在较大变动,适配成本较高
  • 需要额外学习torchtext的数据格式和处理流程,不如原生Sampler方案灵活

总结

优先选择自定义TokenBatchSampler方案,它完全融入PyTorch生态,灵活易用且维护成本低。torchtext方案仅适合已在使用torchtext的项目,兼容性和易用性均不如原生方案。

内容的提问来源于stack exchange,提问作者Bram Vanroy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 04:45:54