(self, x, class_idx=None, retain_graph=False, **kwargs)
| 128 | return logits[:, class_idx].squeeze() |
| 129 | |
| 130 | def __call__(self, x, class_idx=None, retain_graph=False, **kwargs): |
| 131 | train = self.model.training |
| 132 | self.model.eval() |
| 133 | logits = self.model(x, **kwargs) |
| 134 | self.class_idx = logits.max(1)[-1] if class_idx is None else class_idx |
| 135 | acti, grad = None, None |
| 136 | if self.register_forward: |
| 137 | acti = tuple(self.activations[layer] for layer in self.target_layers) |
| 138 | if self.register_backward: |
| 139 | self.score = self.class_score(logits, cast(int, self.class_idx)) |
| 140 | self.model.zero_grad() |
| 141 | self.score.sum().backward(retain_graph=retain_graph) |
| 142 | for layer in self.target_layers: |
| 143 | if layer not in self.gradients: |
| 144 | warnings.warn( |
| 145 | f"Backward hook for {layer} is not triggered; `requires_grad` of {layer} should be `True`." |
| 146 | ) |
| 147 | grad = tuple(self.gradients[layer] for layer in self.target_layers if layer in self.gradients) |
| 148 | if train: |
| 149 | self.model.train() |
| 150 | return logits, acti, grad |
| 151 | |
| 152 | def get_wrapped_net(self): |
| 153 | return self.model |
nothing calls this directly
no test coverage detected