(self, x)
| 286 | return output |
| 287 | |
| 288 | def ssm(self, x): |
| 289 | # x : (B, L, ED) |
| 290 | |
| 291 | # y : (B, L, ED) |
| 292 | |
| 293 | A = -torch.exp(self.A_log.float()) # (ED, N) |
| 294 | D = self.D.float() |
| 295 | # TODO remove .float() |
| 296 | |
| 297 | deltaBC = self.x_proj(x) # (B, L, dt_rank+2*N) |
| 298 | |
| 299 | delta, B, C = torch.split( |
| 300 | deltaBC, |
| 301 | [ |
| 302 | self.config.dt_rank, |
| 303 | self.config.d_state, |
| 304 | self.config.d_state, |
| 305 | ], |
| 306 | dim=-1, |
| 307 | ) # (B, L, dt_rank), (B, L, N), (B, L, N) |
| 308 | delta = F.softplus(self.dt_proj(delta)) # (B, L, ED) |
| 309 | |
| 310 | if self.config.pscan: |
| 311 | y = self.selective_scan(x, delta, A, B, C, D) |
| 312 | else: |
| 313 | y = self.selective_scan_seq(x, delta, A, B, C, D) |
| 314 | |
| 315 | return y |
| 316 | |
| 317 | def selective_scan(self, x, delta, A, B, C, D): |
| 318 | # x : (B, L, ED) |
no test coverage detected