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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 04:59:49