clu*_*ess 9 pytorch oversampling pytorch-dataloader
我正在 PyTorch 中训练一个用于二元分类的深度学习模型,并且我有一个包含不平衡类比例的数据集。10%我的少数派课程由给定的观察结果组成。为了避免模型学习只预测多数类,我想WeightedRandomSampler在torch.utils.data我的DataLoader.
假设我有1000观察结果(900在类中0,100在类中1),并且我的数据加载器的批量大小100为。
如果没有加权随机抽样,我预计每个训练周期将包含 10 个批次。
这取决于您想要什么,请查看torch.utils.data.WeightedRandomSampler文档以了解详细信息。
有一个参数num_samples允许您指定Dataset与 结合时实际创建的样本数量torch.utils.data.DataLoader(假设您对它们进行了正确的加权):
len(dataset)你将得到第一个案例1800(根据您的情况),您将得到第二种情况使用此采样器时,每个时期仅对 10 个批次进行采样 - 因此,模型会在每个时期期间“错过”大多数类别的很大一部分 [...]
是的,但是新的样本将在这个纪元过去后返回
使用采样器是否会导致每个 epoch 采样超过 10 个批次(这意味着相同的少数类观察结果可能会出现多次,并且训练速度也会减慢)?
训练不会减慢,每个时期将花费更长的时间,但收敛应该大致相同(因为每个时期的数据更多,因此需要更少的时期)。
| 归档时间: |
|
| 查看次数: |
12404 次 |
| 最近记录: |