input: attn: (M, h), v: (N, h, C//h), index0: (M), index1: (M) output: output: [L, h, C//h]
(ctx, attn, v, index0, index1)
| 268 | class AttentionStep2_v2(Function): |
| 269 | @staticmethod |
| 270 | def forward(ctx, attn, v, index0, index1): |
| 271 | """ |
| 272 | input: attn: (M, h), v: (N, h, C//h), index0: (M), index1: (M) |
| 273 | output: output: [L, h, C//h] |
| 274 | """ |
| 275 | assert attn.is_contiguous() and v.is_contiguous() and index0.is_contiguous() and index1.is_contiguous() |
| 276 | |
| 277 | L = int(index0.max().item()) + 1 |
| 278 | |
| 279 | M, h = attn.shape |
| 280 | N, h, C_div_h = v.shape |
| 281 | C = int(C_div_h * h) |
| 282 | |
| 283 | output = torch.cuda.FloatTensor(L, h, C//h).zero_() |
| 284 | pointops_cuda.attention_step2_forward_cuda(N, M, h, C, attn, v, index0, index1, output) |
| 285 | ctx.M = M |
| 286 | |
| 287 | # print("attn[:5,:5]: ", attn[:5, :5]) |
| 288 | |
| 289 | ctx.save_for_backward(attn, v, index0, index1) |
| 290 | return output |
| 291 | |
| 292 | @staticmethod |
| 293 | def backward(ctx, grad_output): |
nothing calls this directly
no outgoing calls
no test coverage detected