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

Атрибут training

Атрибут training класса Module определяет, в каком режиме находится модель: обучения или оценки. Значение True соответствует режиму обучения, False - режиму оценки. Этот атрибут влияет на поведение некоторых слоев, например, Dropout и BatchNorm, которые по-разному работают во время обучения и инференса.

Синтаксис

# Получение текущего режима is_training = model.training # Установка режима model.training = True # режим обучения model.training = False # режим оценки

Пример

Давайте создадим простую модель с слоем Dropout и посмотрим на её поведение в разных режимах:

import torch import torch.nn as nn # Создание модели с Dropout class SimpleModel(nn.Module): def __init__(self): super().__init__() self.linear = nn.Linear(10, 10) self.dropout = nn.Dropout(0.5) def forward(self, x): x = self.linear(x) x = self.dropout(x) return x model = SimpleModel() x = torch.ones(1, 10) # Режим обучения model.training = True res1 = model(x) print("Режим обучения:", res1) # Режим оценки model.training = False res2 = model(x) print("Режим оценки:", res2)

Результат выполнения кода:

Режим обучения: tensor([[...]], grad_fn=<MulBackward0>) Режим оценки: tensor([[...]], grad_fn=<MulBackward0>)

В режиме обучения Dropout случайным образом обнуляет часть нейронов, а в режиме оценки масштабирует выходные значения без обнуления.

Пример

Рассмотрим влияние атрибута training на слой BatchNorm, который накапливает статистику по батчам во время обучения:

import torch import torch.nn as nn class BatchNormModel(nn.Module): def __init__(self): super().__init__() self.bn = nn.BatchNorm1d(5) def forward(self, x): return self.bn(x) model = BatchNormModel() x = torch.randn(2, 5) # Режим обучения - обновление статистики model.training = True res1 = model(x) print("Режим обучения (статистика обновляется):") print(res1) # Режим оценки - использование накопленной статистики model.training = False res2 = model(x) print("Режим оценки (используется накопленная статистика):") print(res2)

Результат выполнения кода:

Режим обучения (статистика обновляется): tensor([[...]], grad_fn=<NativeBatchNormBackward0>) Режим оценки (используется накопленная статистика): tensor([[...]], grad_fn=<NativeBatchNormBackward0>)

Пример

Практический пример использования атрибута training для создания пользовательского слоя с разным поведением:

import torch import torch.nn as nn class CustomLayer(nn.Module): def __init__(self, size): super().__init__() self.size = size self.scale = nn.Parameter(torch.ones(size)) def forward(self, x): if self.training: # В режиме обучения добавляем шум noise = torch.randn_like(x) * 0.1 return x * self.scale + noise else: # В режиме оценки просто масштабируем return x * self.scale model = CustomLayer(3) x = torch.ones(1, 3) # Режим обучения model.training = True res1 = model(x) print("Режим обучения:", res1) # Режим оценки model.training = False res2 = model(x) print("Режим оценки:", res2)

Результат выполнения кода:

Режим обучения: tensor([[1.0875, 0.9563, 1.0421]], grad_fn=<AddBackward0>) Режим оценки: tensor([[1., 1., 1.]], grad_fn=<MulBackward0>)

Смотрите также

  • метод train,
    который устанавливает режим обучения для модели и всех её подмодулей
  • метод eval,
    который устанавливает режим оценки для модели и всех её подмодулей
  • класс Module,
    базовый класс для всех нейронных сетей в PyTorch
  • метод forward,
    который определяет проход данных через модель
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить