РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
597 of 769 menu

Класс 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,
    который выполняет случайную выборку из подмножества данных
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить