| 66 | #@once_differentiable |
| 67 | @custom_bwd |
| 68 | def backward(ctx, grad): |
| 69 | |
| 70 | inputs, embeddings, offsets, dy_dx = ctx.saved_tensors |
| 71 | B, D, C, L, S, H, gridtype, interpolation = ctx.dims |
| 72 | align_corners = ctx.align_corners |
| 73 | |
| 74 | # grad: [B, L * C] --> [L, B, C] |
| 75 | grad = grad.view(B, L, C).permute(1, 0, 2).contiguous() |
| 76 | |
| 77 | grad_embeddings = torch.zeros_like(embeddings) |
| 78 | |
| 79 | if dy_dx is not None: |
| 80 | grad_inputs = torch.zeros_like(inputs, dtype=embeddings.dtype) |
| 81 | else: |
| 82 | grad_inputs = None |
| 83 | |
| 84 | _backend.grid_encode_backward(grad, inputs, embeddings, offsets, grad_embeddings, B, D, C, L, S, H, dy_dx, grad_inputs, gridtype, align_corners, interpolation) |
| 85 | |
| 86 | if dy_dx is not None: |
| 87 | grad_inputs = grad_inputs.to(inputs.dtype) |
| 88 | |
| 89 | return grad_inputs, grad_embeddings, None, None, None, None, None, None, None |
| 90 | |
| 91 | |
| 92 | |