| 340 | return y |
| 341 | |
| 342 | def selective_scan_seq(self, x, delta, A, B, C, D): |
| 343 | # x : (B, L, ED) |
| 344 | # Δ : (B, L, ED) |
| 345 | # A : (ED, N) |
| 346 | # B : (B, L, N) |
| 347 | # C : (B, L, N) |
| 348 | # D : (ED) |
| 349 | |
| 350 | # y : (B, L, ED) |
| 351 | |
| 352 | _, L, _ = x.shape |
| 353 | |
| 354 | deltaA = torch.exp(delta.unsqueeze(-1) * A) # (B, L, ED, N) |
| 355 | deltaB = delta.unsqueeze(-1) * B.unsqueeze(2) # (B, L, ED, N) |
| 356 | |
| 357 | BX = deltaB * (x.unsqueeze(-1)) # (B, L, ED, N) |
| 358 | |
| 359 | h = torch.zeros( |
| 360 | x.size(0), |
| 361 | self.config.d_inner, |
| 362 | self.config.d_state, |
| 363 | device=deltaA.device, |
| 364 | ) # (B, ED, N) |
| 365 | hs = [] |
| 366 | |
| 367 | for t in range(0, L): |
| 368 | h = deltaA[:, t] * h + BX[:, t] |
| 369 | hs.append(h) |
| 370 | |
| 371 | hs = torch.stack(hs, dim=1) # (B, L, ED, N) |
| 372 | |
| 373 | y = (hs @ C.unsqueeze(-1)).squeeze( |
| 374 | 3 |
| 375 | ) # (B, L, ED, N) @ (B, L, N, 1) -> (B, L, ED, 1) |
| 376 | |
| 377 | y = y + D * x |
| 378 | |
| 379 | return y |
| 380 | |
| 381 | # -------------------------- inference -------------------------- # |
| 382 | """ |