如何在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函数等)。
实现步骤
- 预计算所有样本的token长度,避免重复计算
- 实现
Sampler子类,按总token数不超过阈值的规则生成批次索引 - 将自定义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在新版本中变化较大,且需要适配其数据格式。
实现步骤
- 定义
batch_size_fn函数,用于计算当前batch的总token数 - 将PyTorch Dataset转换为torchtext兼容的格式
- 初始化
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
相关产品推荐
相关产品推荐

