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

Функция no_grad

Функция no_grad представляет собой контекстный менеджер, который отключает вычисление градиентов для всех операций внутри своего блока. Это позволяет существенно ускорить выполнение кода и уменьшить потребление памяти, так как PyTorch не сохраняет историю вычислений для градиентов. Функция не принимает параметров и используется в паре с оператором with.

Синтаксис

with torch.no_grad(): # операции без градиентов t = тензор + 1

Пример

Давайте создадим тензор с градиентом и выполним операцию внутри контекста no_grad:

import torch t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) print(t.requires_grad) with torch.no_grad(): res = t + 1 print(res.requires_grad)

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

True False

Пример

Давайте используем no_grad для вычислений без сохранения истории градиентов:

import torch t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) with torch.no_grad(): res = t * 2 res = res + 10 print(res) print(res.requires_grad)

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

tensor([12., 14., 16.]) False

Пример

Давайте сравним производительность вычислений с no_grad и без него:

import torch import time t = torch.randn(1000, 1000, requires_grad=True) start = time.time() for _ in range(100): res = t * 2 + 1 end = time.time() print(f"Without no_grad: {end - start:.4f} seconds") start = time.time() with torch.no_grad(): for _ in range(100): res = t * 2 + 1 end = time.time() print(f"With no_grad: {end - start:.4f} seconds")

Пример

Давайте используем no_grad для оценки модели на валидационных данных:

import torch model = torch.nn.Linear(10, 1) model.train() data = torch.randn(32, 10) target = torch.randn(32, 1) # обучение output = model(data) loss = torch.nn.functional.mse_loss(output, target) loss.backward() # валидация без градиентов model.eval() with torch.no_grad(): val_data = torch.randn(32, 10) val_output = model(val_data) print(val_output)

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

  • функцию enable_grad,
    которая включает вычисление градиентов в контексте
  • функцию set_grad_enabled,
    которая включает или отключает градиенты по условию
  • функцию is_grad_enabled,
    которая проверяет текущее состояние вычисления градиентов
  • функцию inference_mode,
    которая является более быстрой альтернативой для инференса
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить