(self, grad_output)
| 239 | |
| 240 | # This function has only a single output, so it gets only one gradient |
| 241 | def backward(self, grad_output): |
| 242 | |
| 243 | input, target = self.saved_variables |
| 244 | grad_input = grad_target = None |
| 245 | |
| 246 | if self.needs_input_grad[0]: |
| 247 | grad_input = grad_output * 2 * (target * self.union - self.inter) \ |
| 248 | / (self.union * self.union) |
| 249 | if self.needs_input_grad[1]: |
| 250 | grad_target = None |
| 251 | |
| 252 | return grad_input, grad_target |