У torch.compile другая семантика autograd
Создано: 26 июня 2025 г. | Последнее обновление: 2 июня 2026 г.
Когда вы применяете torch.compile к функции в прямом проходе вашей модели, для скомпилированной функции автоматически будет сгенерирован обратный проход. Во время компиляции будет построен граф обратного прохода, который используется при каждом вызове autograd. Мы называем компонент внутри torch.compile, отвечающий за это, AOTDispatcher (иногда его называют AOTAutograd).
Таким образом, torch.compile во время компиляции функции в прямом проходе встраивает детали вычислений в построенный граф обратного прохода. Однако в PyTorch в eager-режиме вычисления обратного прохода выполняются динамически: вне прямого прохода вы можете обернуть вызов tensor.backward() или torch.autograd.grad(...) в менеджер контекста, который может изменить его поведение.
На этой странице описано, чем семантика autograd в torch.compile отличается от семантики в PyTorch в eager-режиме и как это обойти.
Поведение Autocast
torch.compile встраивает предположение о том, будет ли обратный проход выполняться в активном контексте autocast. Используйте torch._functorch.config.backward_pass_autocast, чтобы управлять этим предположением; неверное предположение может привести к незаметным ошибкам.
Предупреждение
AMP рекомендует, чтобы torch.autocast оборачивал только прямой проход и вычисление функции потерь. Не рекомендуется выполнять обратные проходы в контексте autocast; см. рекомендации torch.autocast в разделе Autocasting. Настройка компилятора по умолчанию, "same_as_forward", намеренно сохраняет существующее поведение torch.compile, предполагая, что скомпилированный обратный проход выполняется в том же контексте autocast, что и скомпилированный прямой проход. Если ваш код следует рекомендациям AMP и выполняет обратный проход вне autocast, задайте для torch._functorch.config.backward_pass_autocast значение "off" в скомпилированной области.
Возможны следующие варианты:
-
"same_as_forward"(по умолчанию). Предполагается, что обратный проход области, скомпилированной с помощьюtorch.compile, будет выполняться в том же менеджере контекста autocast, в котором выполнялась эта область (если он был). Используйте этот вариант, если ваш код выглядит следующим образом:with torch.amp.autocast(...): y = torch.compile(region)(x) ... # backward pass run under the same autocast context as the compiled region z.backward() -
"off". Предполагается, что обратный проход области, скомпилированной с помощью torch.compile, не будет выполняться ни в одном менеджере контекста autocast. Используйте этот вариант, если ваш код выглядит следующим образом:with torch.amp.autocast(...): y = torch.compile(region)(x) ... # Backward pass runs under no autocast. z.backward() -
Есть и третий вариант. Если задать для
torch._functorch.config.backward_pass_autocastсписок аргументов kwargs, предполагается, что обратный проход выполняется в контексте autocast, созданном с помощью этих kwargs.Например, если ваш код выглядит следующим образом:
y = torch.compile(region)(x) ... # Backward pass runs under special context manager with torch.amp.autocast(**kwargs): z.backward()то задайте
torch._functorch.config.backward_pass_autocast = kwargs.
Используйте patch, чтобы применить этот вариант к конкретному вызову torch.compile:
with torch.amp.autocast(...):
with torch._functorch.config.patch(backward_pass_autocast="same_as_forward")
y = torch.compile(region)(x)
...
# backward pass run under the same autocast context as the compiled region
z.backward()
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/user_guide/torch_compiler/torch.compiler_backward.html