基于PyTorch DataLoader实现BERT训练的动态数据采样方案问询
解决IMDB文档生成不定数量BERT样本的动态批次加载问题
针对你的场景,这里提供三种实用方案,重点落地你提到的动态采样思路:
方案1:预生成所有样本(教学场景快速实现)
如果IMDB数据集规模不大,直接预生成所有sentence pair样本,用普通Dataset+DataLoader即可,代码最简单直观:
import os import re import random import torch from torch.utils.data import Dataset, DataLoader from transformers import BertTokenizer def generate_sentence_pairs(document_sentences): """单文档生成所有sentence pair样本""" pairs = [] num_sentences = len(document_sentences) for i in range(num_sentences - 1): # 连续句子对(is_next=True) pairs.append((document_sentences[i], document_sentences[i+1], True)) # 随机非连续句子对(is_next=False) random_idx = random.randint(0, num_sentences - 1) while random_idx == i+1: random_idx = random.randint(0, num_sentences - 1) pairs.append((document_sentences[i], document_sentences[random_idx], False)) return pairs class IMDBAllPairsDataset(Dataset): def __init__(self, data_dir, split='train'): self.all_pairs = [] # 加载并预处理所有文档 for label in ['pos', 'neg']: dir_path = os.path.join(data_dir, split, label) for fname in os.listdir(dir_path): with open(os.path.join(dir_path, fname), 'r', encoding='utf-8') as f: text = f.read() sentences = re.split(r'[.!?]+', text) sentences = [s.strip() for s in sentences if s.strip()] if len(sentences) >= 2: self.all_pairs.extend(generate_sentence_pairs(sentences)) # 训练时打乱样本 if split == 'train': random.shuffle(self.all_pairs) def __len__(self): return len(self.all_pairs) def __getitem__(self, idx): return self.all_pairs[idx] # 定义collate_fn处理批次tokenization def collate_fn(batch): sents_a = [pair[0] for pair in batch] sents_b = [pair[1] for pair in batch] is_next = [pair[2] for pair in batch] tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') tokenized = tokenizer(sents_a, sents_b, padding=True, truncation=True, return_tensors='pt') tokenized['is_next'] = torch.tensor(is_next, dtype=torch.long) return tokenized # 使用示例 dataset = IMDBAllPairsDataset(data_dir='path/to/imdb', split='train') dataloader = DataLoader(dataset, batch_size=32, collate_fn=collate_fn, num_workers=4, shuffle=True)
方案2:动态采样(实现你的思路,适合大数据集)
如果数据集太大无法预存所有样本,用IterableDataset实现动态缓存凑批次,完美匹配你提出的"每个文档返回不定数量样本,动态补充凑批次"需求:
1. 定义IterableDataset
class IMDBIterableDataset(torch.utils.data.IterableDataset): def __init__(self, data_dir, split='train', batch_size=32): self.data_dir = data_dir self.split = split self.batch_size = batch_size # 预加载所有文档的句子列表(也可在迭代时动态加载,进一步节省内存) self.documents = [] for label in ['pos', 'neg']: dir_path = os.path.join(data_dir, split, label) for fname in os.listdir(dir_path): with open(os.path.join(dir_path, fname), 'r', encoding='utf-8') as f: text = f.read() sentences = re.split(r'[.!?]+', text) sentences = [s.strip() for s in sentences if s.strip()] if len(sentences) >= 2: self.documents.append(sentences) # 训练时打乱文档顺序 if split == 'train': random.shuffle(self.documents) def __iter__(self): cache = [] # 多进程下划分文档给不同worker,避免重复处理同一文档 worker_info = torch.utils.data.get_worker_info() if worker_info is not None: assigned_docs = self.documents[worker_info.id::worker_info.num_workers] else: assigned_docs = self.documents # 遍历文档生成样本,缓存凑批次 for doc in assigned_docs: pairs = generate_sentence_pairs(doc) cache.extend(pairs) # 缓存足够时返回批次 while len(cache) >= self.batch_size: yield cache[:self.batch_size] cache = cache[self.batch_size:] # 处理剩余样本(可选,若不需要小批次可注释) if cache: yield cache
2. 搭配DataLoader使用
# 复用之前的collate_fn tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') dataset = IMDBIterableDataset(data_dir='path/to/imdb', split='train', batch_size=32) # 注意batch_size设为None,因为IterableDataset直接返回批次 dataloader = DataLoader(dataset, batch_size=None, collate_fn=collate_fn, num_workers=4) # 测试加载 for batch in dataloader: print(f"Input IDs shape: {batch['input_ids'].shape}") print(f"Is_next labels shape: {batch['is_next'].shape}") break
关键细节说明
- 多进程兼容:通过
worker_info给每个进程分配不同的文档子集,避免重复处理 - 缓存机制:每个文档生成的样本先存入缓存,凑够批次大小再返回,完全符合你提出的动态补充逻辑
- 内存友好:无需预存所有样本,适合大规模数据集
方案3:自定义BatchSampler(基于普通Dataset)
如果你想基于现有普通Dataset扩展动态采样,可自定义BatchSampler管理样本流:
class IMDBDocumentDataset(Dataset): def __init__(self, data_dir, split='train'): self.data = [] # 加载所有文档的句子列表 for label in ['pos', 'neg']: dir_path = os.path.join(data_dir, split, label) for fname in os.listdir(dir_path): with open(os.path.join(dir_path, fname), 'r', encoding='utf-8') as f: text = f.read() sentences = re.split(r'[.!?]+', text) sentences = [s.strip() for s in sentences if s.strip()] if len(sentences) >= 2: self.data.append(sentences) def __len__(self): return len(self.data) def __getitem__(self, idx): # 返回该文档生成的所有sentence pair样本 return generate_sentence_pairs(self.data[idx]) class DynamicBatchSampler(torch.utils.data.Sampler): def __init__(self, dataset, batch_size=32, drop_last=False): self.dataset = dataset self.batch_size = batch_size self.drop_last = drop_last # 预计算每个文档的样本数,方便采样规划 self.doc_sample_counts = [len(sample) for sample in dataset] def __iter__(self): cache = [] doc_indices = iter(range(len(self.dataset))) while True: try: idx = next(doc_indices) cache.extend(self.dataset[idx]) # 凑批次返回 while len(cache) >= self.batch_size: yield cache[:self.batch_size] cache = cache[self.batch_size:] except StopIteration: break # 处理剩余样本 if cache and not self.drop_last: yield cache def __len__(self): total_samples = sum(self.doc_sample_counts) return total_samples // self.batch_size if self.drop_last else (total_samples + self.batch_size -1) // self.batch_size # 使用示例 dataset = IMDBDocumentDataset(data_dir='path/to/imdb', split='train') sampler = DynamicBatchSampler(dataset, batch_size=32) dataloader = DataLoader(dataset, batch_sampler=sampler, collate_fn=collate_fn, num_workers=4)
选择建议
- 教学演示优先选方案1,代码简洁易懂,快速验证模型逻辑
- 追求真实工业场景还原选方案2,内存效率高,多进程支持完善
- 若需基于现有普通Dataset扩展,选方案3
内容的提问来源于stack exchange,提问作者Ali Haider Ahmad
相关产品推荐
相关产品推荐

