MCPcopy Create free account
hub / github.com/pytorch/pytorch / backward

Method backward

torch/utils/checkpoint.py:265–325  ·  view source on GitHub ↗
(ctx, *args)

Source from the content-addressed store, hash-verified

263
264 @staticmethod
265 def backward(ctx, *args):
266 if not torch.autograd._is_checkpoint_valid():
267 raise RuntimeError(
268 "Checkpointing is not compatible with .grad() or when an `inputs` parameter"
269 " is passed to .backward(). Please use .backward() and do not pass its `inputs`"
270 " argument."
271 )
272 # Copy the list to avoid modifying original list.
273 inputs = list(ctx.inputs)
274 tensor_indices = ctx.tensor_indices
275 tensors = ctx.saved_tensors
276 device_module = _get_device_module(ctx.device)
277
278 # Fill in inputs with appropriate saved tensors.
279 for i, idx in enumerate(tensor_indices):
280 inputs[idx] = tensors[i]
281
282 # Stash the surrounding rng state, and mimic the state that was
283 # present at this time during forward. Restore the surrounding state
284 # when we're done.
285 rng_devices = []
286 if ctx.preserve_rng_state and ctx.had_device_in_fwd:
287 rng_devices = ctx.fwd_devices
288 with torch.random.fork_rng(
289 devices=rng_devices, enabled=ctx.preserve_rng_state, device_type=ctx.device
290 ):
291 if ctx.preserve_rng_state:
292 torch.set_rng_state(ctx.fwd_cpu_state)
293 if ctx.had_device_in_fwd:
294 set_device_states(ctx.fwd_devices, ctx.fwd_device_states)
295 detached_inputs = detach_variable(tuple(inputs))
296
297 device_autocast_ctx = device_module.amp.autocast(
298 **ctx.device_autocast_kwargs
299 ) if _supports_autocast(ctx.device) else contextlib.nullcontext()
300 with torch.enable_grad(), device_autocast_ctx, \
301 torch.cpu.amp.autocast(**ctx.cpu_autocast_kwargs):
302 outputs = ctx.run_function(*detached_inputs)
303
304 if isinstance(outputs, torch.Tensor):
305 outputs = (outputs,)
306
307 # run backward() with only tensor that requires grad
308 outputs_with_grad = []
309 args_with_grad = []
310 for i in range(len(outputs)):
311 if torch.is_tensor(outputs[i]) and outputs[i].requires_grad:
312 outputs_with_grad.append(outputs[i])
313 args_with_grad.append(args[i])
314 if len(outputs_with_grad) == 0:
315 raise RuntimeError(
316 "none of output has requires_grad=True,"
317 " this checkpoint() is not necessary"
318 )
319 torch.autograd.backward(outputs_with_grad, args_with_grad)
320 grads = tuple(
321 inp.grad if isinstance(inp, torch.Tensor) else None
322 for inp in detached_inputs

Callers

nothing calls this directly

Calls 12

listFunction · 0.85
set_device_statesFunction · 0.85
detach_variableFunction · 0.85
_supports_autocastFunction · 0.85
isinstanceFunction · 0.85
set_rng_stateMethod · 0.80
enable_gradMethod · 0.80
is_tensorMethod · 0.80
_get_device_moduleFunction · 0.70
rangeFunction · 0.50
appendMethod · 0.45
backwardMethod · 0.45

Tested by

no test coverage detected