Callback function to be executed before the `loss.backward()` call.
(
self,
params: Union[Dict[str, torch.nn.Parameter], torch.nn.ParameterDict],
optimizers: Dict[str, torch.optim.Optimizer],
state: Dict[str, Any],
step: int,
info: Dict[str, Any],
)
| 75 | assert key in params, f"{key} is required in params but missing." |
| 76 | |
| 77 | def step_pre_backward( |
| 78 | self, |
| 79 | params: Union[Dict[str, torch.nn.Parameter], torch.nn.ParameterDict], |
| 80 | optimizers: Dict[str, torch.optim.Optimizer], |
| 81 | state: Dict[str, Any], |
| 82 | step: int, |
| 83 | info: Dict[str, Any], |
| 84 | ): |
| 85 | """Callback function to be executed before the `loss.backward()` call.""" |
| 86 | assert ( |
| 87 | "means2d" in info |
| 88 | ), "The 2D means of the Gaussians is required but missing." |
| 89 | info["means2d"].retain_grad() |
| 90 | |
| 91 | def step_post_backward( |
| 92 | self, |
nothing calls this directly
no outgoing calls
no test coverage detected