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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 19:07:17