| 436 | return output, cache |
| 437 | |
| 438 | def ssm_step(self, x, h): |
| 439 | # x : (B, ED) |
| 440 | # h : (B, ED, N) |
| 441 | |
| 442 | # y : (B, ED) |
| 443 | # h : (B, ED, N) |
| 444 | |
| 445 | A = -torch.exp( |
| 446 | self.A_log.float() |
| 447 | ) # (ED, N) # todo : ne pas le faire tout le temps, puisque c'est indépendant de la timestep |
| 448 | D = self.D.float() |
| 449 | # TODO remove .float() |
| 450 | |
| 451 | deltaBC = self.x_proj(x) # (B, dt_rank+2*N) |
| 452 | |
| 453 | delta, B, C = torch.split( |
| 454 | deltaBC, |
| 455 | [ |
| 456 | self.config.dt_rank, |
| 457 | self.config.d_state, |
| 458 | self.config.d_state, |
| 459 | ], |
| 460 | dim=-1, |
| 461 | ) # (B, dt_rank), (B, N), (B, N) |
| 462 | delta = F.softplus(self.dt_proj(delta)) # (B, ED) |
| 463 | |
| 464 | deltaA = torch.exp(delta.unsqueeze(-1) * A) # (B, ED, N) |
| 465 | deltaB = delta.unsqueeze(-1) * B.unsqueeze(1) # (B, ED, N) |
| 466 | |
| 467 | BX = deltaB * (x.unsqueeze(-1)) # (B, ED, N) |
| 468 | |
| 469 | if h is None: |
| 470 | h = torch.zeros( |
| 471 | x.size(0), |
| 472 | self.config.d_inner, |
| 473 | self.config.d_state, |
| 474 | device=deltaA.device, |
| 475 | ) # (B, ED, N) |
| 476 | |
| 477 | h = deltaA * h + BX # (B, ED, N) |
| 478 | |
| 479 | y = (h @ C.unsqueeze(-1)).squeeze(2) # (B, ED, N) @ (B, N, 1) -> (B, ED, 1) |
| 480 | |
| 481 | y = y + D * x |
| 482 | |
| 483 | # todo : pq h.squeeze(1) ?? |
| 484 | return y, h.squeeze(1) |
| 485 | |
| 486 | |
| 487 | class Mamba(nn.Module): |