Класс WeightedRandomSampler
Класс WeightedRandomSampler предназначен для создания сэмплера,
который выбирает элементы из набора данных с вероятностью,
пропорциональной заданным весам. Первым параметром конструктор
принимает тензор или список весов weights, вторым - количество
элементов для выборки num_samples. Также можно указать
параметр replacement, который определяет, разрешена ли
выборка с повторением.
Синтаксис
torch.utils.data.WeightedRandomSampler(weights, num_samples, replacement=False)
Параметры
Конструктор класса принимает следующие параметры:
-
weights- тензор или список с весами для каждого элемента. Чем больше вес, тем выше вероятность попадания элемента в выборку. -
num_samples- количество элементов, которые необходимо выбрать. -
replacement- логический флаг, указывающий, разрешена ли выборка с повторением. По умолчаниюFalse.
Пример
Давайте создадим простой сэмплер с весами и выполним выборку:
import torch
from torch.utils.data import WeightedRandomSampler
torch.manual_seed(0)
weights = torch.tensor([0.1, 0.2, 0.7])
sampler = WeightedRandomSampler(weights, num_samples=5, replacement=True)
indices = list(sampler)
print(indices)
Результат выполнения кода:
[2, 0, 1, 0, 1]
Пример
Теперь создадим сэмплер без повторений и выполним выборку:
import torch
from torch.utils.data import WeightedRandomSampler
torch.manual_seed(0)
weights = torch.tensor([0.1, 0.8, 0.1])
sampler = WeightedRandomSampler(weights, num_samples=3, replacement=False)
indices = list(sampler)
print(indices)
Результат выполнения кода:
[1, 2, 0]
Пример
Используем сэмплер вместе с загрузчиком данных для выбора элементов из датасета:
import torch
from torch.utils.data import DataLoader, TensorDataset, WeightedRandomSampler
torch.manual_seed(0)
data = torch.tensor([10, 20, 30, 40, 50])
targets = torch.tensor([0, 1, 0, 1, 0])
dataset = TensorDataset(data, targets)
weights = torch.tensor([0.3, 0.1, 0.2, 0.1, 0.3])
sampler = WeightedRandomSampler(weights, num_samples=4, replacement=True)
loader = DataLoader(dataset, sampler=sampler, batch_size=2)
for batch_data, batch_targets in loader:
print(f"Data: {batch_data}, Targets: {batch_targets}")
Результат выполнения кода:
Data: tensor([50, 20]), Targets: tensor([0, 1])
Data: tensor([10, 30]), Targets: tensor([0, 0])
Пример
Создадим сэмплер с неравномерными весами для борьбы с дисбалансом классов:
import torch
from torch.utils.data import DataLoader, TensorDataset, WeightedRandomSampler
torch.manual_seed(42)
data = torch.randn(1000, 10)
targets = torch.cat([torch.zeros(900), torch.ones(100)]).long()
class_counts = torch.bincount(targets)
class_weights = 1.0 / class_counts.float()
sample_weights = class_weights[targets]
sampler = WeightedRandomSampler(sample_weights, num_samples=200, replacement=True)
dataset = TensorDataset(data, targets)
loader = DataLoader(dataset, sampler=sampler, batch_size=32)
selected = torch.cat([batch_targets for _, batch_targets in loader])
print(f"Selected class distribution: {torch.bincount(selected)}")
Результат выполнения кода:
Selected class distribution: tensor([100, 100])
Смотрите также
-
класс
Sampler,
который является базовым классом для всех сэмплеров -
класс
SequentialSampler,
который последовательно перебирает элементы -
класс
RandomSampler,
который выполняет случайную выборку без весов -
класс
SubsetRandomSampler,
который выполняет случайную выборку из подмножества данных