MCPcopy Create free account
hub / github.com/Project-MONAI/MONAI / get_grad

Method get_grad

monai/visualize/gradient_based.py:89–98  ·  view source on GitHub ↗
(
        self, x: torch.Tensor, index: torch.Tensor | int | None, retain_graph: bool = True, **kwargs: Any
    )

Source from the content-addressed store, hash-verified

87 self._model = m # replace the ModelWithHooks
88
89 def get_grad(
90 self, x: torch.Tensor, index: torch.Tensor | int | None, retain_graph: bool = True, **kwargs: Any
91 ) -> torch.Tensor:
92 if x.shape[0] != 1:
93 raise ValueError("expect batch size of 1")
94 x.requires_grad = True
95
96 self._model(x, class_idx=index, retain_graph=retain_graph, **kwargs)
97 grad: torch.Tensor = x.grad.detach() # type: ignore
98 return grad
99
100 def __call__(self, x: torch.Tensor, index: torch.Tensor | int | None = None, **kwargs: Any) -> torch.Tensor:
101 return self.get_grad(x, index, **kwargs)

Callers 2

__call__Method · 0.95
__call__Method · 0.80

Calls

no outgoing calls

Tested by

no test coverage detected