MCPcopy Create free account
hub / github.com/OpenSparseLLMs/MoM / transform

Function transform

mom_naive/layers/mom_gsa.py:27–119  ·  view source on GitHub ↗

Transform input sequences into memory-organized chunks with capacity constraints. Processes input sequences by routing tokens to designated memory states according to routing_mask, sorts tokens by memory assignments, handles token truncation/padding based on memory capacity, an

(x: torch.Tensor, routing_mask: torch.Tensor, num_memories: int, selected_memories: torch.Tensor, capacity: float)

Source from the content-addressed store, hash-verified

25 from fla.models.utils import Cache
26
27def transform(x: torch.Tensor, routing_mask: torch.Tensor, num_memories: int, selected_memories: torch.Tensor, capacity: float):
28 '''
29 Transform input sequences into memory-organized chunks with capacity constraints.
30
31 Processes input sequences by routing tokens to designated memory states according to routing_mask,
32 sorts tokens by memory assignments, handles token truncation/padding based on memory capacity,
33 and returns memory-aligned tensors for parallel processing.
34
35 Key operations:
36 1. Expands input tensors when multiple memories are selected per token (top-k routing)
37 2. Sorts tokens globally by (batch_idx, memory_idx) to group memory-assigned tokens
38 3. Applies capacity-aware truncation (left-truncate oldest tokens when exceeding capacity)
39 4. Pads memory chunks to uniform length for tensorization
40
41 Args:
42 x: Input hidden states
43 Shape: (batch_size, seq_len, hidden_size)
44 routing_mask: Binary mask indicating active memory assignments
45 Shape: (batch_size, seq_len, num_memories)
46 num_memories: Total number of memories per batch
47 selected_memories: Memory indices assigned to each token. When using top-k routing,
48 this contains k memory indices per token (k >= 1)
49 Shape: (batch_size, seq_len) for k=1 or (batch_size, seq_len, topk) for k>1
50 capacity: Scaling factor for memory capacity calculation. Actual capacity per memory is
51 ceil(seq_len * capacity), maintaining proportional capacity to sequence length
52
53 Returns:
54 transformed_x: Memory-organized tensor with zero-padded capacity alignment
55 Shape: (num_memories, batch_size, capacity_len, hidden_size)
56 truncation_indices: Original indices used for gathering tokens after capacity truncation
57 Shape: (batch*num_memories, max_len)
58 sorted_indices: Sorting indices used to group tokens by memory assignments
59 Shape: (batch_size*seq_len*topk)
60 max_len: Maximum tokens per memory
61 mask: Boolean mask indicating valid (non-padded) positions in transformed_x
62 Shape: (batch*num_memories, max_len)
63 '''
64 if selected_memories.dim() == 3:
65 # (batch, seq, topk)
66 topk = selected_memories.shape[2]
67 # x (batch, seq, hidden)
68 x = x.repeat_interleave(topk, dim=1)
69 # x (batch, seq * topk, hidden)
70 # (batch, seq, topk)
71 selected_memories = selected_memories.reshape(selected_memories.shape[0], -1)
72 # (batch, seq * topk)
73
74 b, s, d = x.shape
75 x_flat = x.reshape(b * s, d) # [b*s, d]
76
77 with torch.no_grad():
78 batch_indices = torch.arange(b, device=x.device).unsqueeze(-1)
79 batch_indices = batch_indices.expand(b, s).reshape(-1)
80 # (b * s)
81 memories_flat = selected_memories.reshape(-1) # [b*s]
82
83 combined = batch_indices * (memories_flat.max() + 1) + memories_flat
84 sorted_indices = combined.argsort()

Callers 1

forwardMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected