MCPcopy Create free account
hub / github.com/kyegomez/BitNet / selective_scan_seq

Method selective_scan_seq

bitnet/bit_mamba.py:342–379  ·  view source on GitHub ↗
(self, x, delta, A, B, C, D)

Source from the content-addressed store, hash-verified

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 """

Callers 1

ssmMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected