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

Method selective_scan

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

Source from the content-addressed store, hash-verified

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)

Callers 1

ssmMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected