如何在PyTorch中结合DistributedSampler使用类别权重?
多GPU训练时结合DistributedSampler与WeightedRandomSampler处理类别不平衡
由于PyTorch的DataLoader仅支持传入一个采样器,无法直接同时使用DistributedSampler和WeightedRandomSampler,以下提供两种实用的解决方案:
方案1:基于预先生成的加权索引实现分布式采样
先全局生成符合加权规则的样本索引,再通过DistributedSampler对这些索引进行进程划分,既保证类别加权,又满足分布式训练的样本分配逻辑。
代码示例:
import torch from torch.utils.data import Dataset, DataLoader, DistributedSampler, WeightedRandomSampler # 1. 初始化原数据集并计算样本权重 train_dataset = SampleDataset(data_root=data_root, transform=train_transforms, num_classes=num_classes) # 假设数据集有labels属性存储每个样本的类别标签 class_counts = torch.bincount(torch.tensor(train_dataset.labels)) class_weights = 1.0 / class_counts.float() sample_weights = class_weights[train_dataset.labels] # 2. 生成全局加权采样索引(可设置采样次数,此处与原数据集大小一致) weighted_sampler = WeightedRandomSampler(sample_weights, num_samples=len(train_dataset), replacement=True) weighted_indices = list(weighted_sampler) # 3. 定义索引包装数据集,用于通过索引取原数据集样本 class IndexDataset(Dataset): def __init__(self, original_dataset, indices): self.original_dataset = original_dataset self.indices = indices def __getitem__(self, idx): return self.original_dataset[self.indices[idx]] def __len__(self): return len(self.indices) index_dataset = IndexDataset(train_dataset, weighted_indices) # 4. 对索引数据集使用DistributedSampler train_sampler = DistributedSampler(index_dataset, num_replicas=world_size) train_loader = DataLoader(index_dataset, batch_size=batch_size, pin_memory=True, sampler=train_sampler, num_workers=0)
注意:若需每个epoch重新生成加权采样序列(避免重复),需在每个epoch开始时重新生成weighted_indices并更新index_dataset.indices,同时调用train_sampler.set_epoch(epoch)保证分布式采样的随机性。
方案2:自定义结合逻辑的DistributedWeightedSampler
直接继承DistributedSampler,重写__iter__方法,在生成进程专属索引时加入加权采样逻辑,无需额外包装数据集,实现更简洁。
代码示例:
from torch.utils.data import DistributedSampler import torch class DistributedWeightedSampler(DistributedSampler): def __init__(self, dataset, weights, num_replicas=None, rank=None, shuffle=True, replacement=True, num_samples=None): super().__init__(dataset, num_replicas=num_replicas, rank=rank, shuffle=shuffle) self.weights = weights self.replacement = replacement self.num_samples = num_samples if num_samples is not None else len(dataset) def __iter__(self): # 全局生成加权采样索引,用epoch作为种子保证每个epoch采样不同 g = torch.Generator() g.manual_seed(self.epoch) indices = torch.multinomial(self.weights, self.num_samples, replacement=self.replacement, generator=g).tolist() # 按照DistributedSampler规则划分当前进程的索引 indices = indices[self.rank:self.total_size:self.num_replicas] return iter(indices) # 使用自定义采样器 train_dataset = SampleDataset(data_root=data_root, transform=train_transforms, num_classes=num_classes) class_counts = torch.bincount(torch.tensor(train_dataset.labels)) class_weights = 1.0 / class_counts.float() sample_weights = class_weights[train_dataset.labels] train_sampler = DistributedWeightedSampler(train_dataset, sample_weights, num_replicas=world_size) train_loader = DataLoader(train_dataset, batch_size=batch_size, pin_memory=True, sampler=train_sampler, num_workers=0) # 训练时每个epoch必须调用set_epoch,保证采样随机性 for epoch in range(epochs): train_sampler.set_epoch(epoch) # ... 训练逻辑代码
内容的提问来源于stack exchange,提问作者let me down slowly
相关产品推荐
相关产品推荐

