| 315 | return y |
| 316 | |
| 317 | def selective_scan(self, x, delta, A, B, C, D): |
| 318 | # x : (B, L, ED) |
| 319 | # Δ : (B, L, ED) |
| 320 | # A : (ED, N) |
| 321 | # B : (B, L, N) |
| 322 | # C : (B, L, N) |
| 323 | # D : (ED) |
| 324 | |
| 325 | # y : (B, L, ED) |
| 326 | |
| 327 | deltaA = torch.exp(delta.unsqueeze(-1) * A) # (B, L, ED, N) |
| 328 | deltaB = delta.unsqueeze(-1) * B.unsqueeze(2) # (B, L, ED, N) |
| 329 | |
| 330 | BX = deltaB * (x.unsqueeze(-1)) # (B, L, ED, N) |
| 331 | |
| 332 | hs = pscan(deltaA, BX) |
| 333 | |
| 334 | y = (hs @ C.unsqueeze(-1)).squeeze( |
| 335 | 3 |
| 336 | ) # (B, L, ED, N) @ (B, L, N, 1) -> (B, L, ED, 1) |
| 337 | |
| 338 | y = y + D * x |
| 339 | |
| 340 | return y |
| 341 | |
| 342 | def selective_scan_seq(self, x, delta, A, B, C, D): |
| 343 | # x : (B, L, ED) |