(ctx, *args)
| 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 |
nothing calls this directly
no test coverage detected