Skip to content

Token Merging

Overview

Token Merging (ToMe) reduces computation in transformer inference by merging similar tokens before attention, reducing sequence length. Achieves 2-4x speedup with minimal quality loss by removing redundant information.

  • Paper: "Token Merging for Fast Stable Diffusion" (Bolya et al., 2023)
  • Key Idea: Merge similar tokens before expensive attention computation
  • Speedup: 2-4x with <2% quality loss
  • Advantage: Works during inference without retraining
  • Trade-off: Slight information loss from merging

The Problem: Sequence Length Drives Attention Cost

Attention Complexity

Attention computation cost:

Q @ K^T has complexity O(N²) where N = sequence length

Concrete example:
Model: LLaMA 7B
Sequence: 4096 tokens
Attention at each layer:
 - Q shape: (32 heads, 4096, 128)
 - K^T shape: (32 heads, 128, 4096)
 - Output: (32 heads, 4096, 4096)
 - Operations: 4096 × 4096 × 128 = 2.1B operations per layer
 - 32 layers × 2.1B = 67B operations (just attention!)
 - Takes ~50% of inference time!

Key insight:
Problem is O(N²), not the actual computations

If we reduce N from 4096 to 2048:
 - New cost: 2048 × 2048 × 128 = 0.5B per layer
 - Total: 8 × 0.5B = 16B (4x reduction!)
 - Speedup: ~2-4x possible!

Question: Can we reduce sequence length without losing information?
Answer: Yes! Through token merging!

Why Merging is Feasible

Observation: Many tokens are redundant

Example: "The quick brown fox jumped over the lazy dog"

Token importance:
 - "The" → common word, low information
 - "quick" → descriptive, important
 - "brown" → descriptive, important 
 - "fox" → main subject, very important
 - "jumped" → action verb, very important
 - "lazy" → descriptive, important

Redundancy:
- Multiple "the" are similar
- Adjacent descriptors can overlap
- Some tokens add little new information

Merging strategy:
 - Identify similar tokens
 - Merge their representations
 - Reduce sequence length
 - Attention now operates on merged sequence!

Result:
Original: [The, quick, brown, fox, jumped, over, the, lazy, dog]
Merged: [The+quick+brown, fox, jumped, over, lazy+dog]
 - Reduced from 9 tokens to 6 (33% reduction)

-

How Token Merging Works

Core Algorithm

Input sequence: [t₁, t₂, t₃, t₄, t₅, t₆, t₇, t₈]

Step 1: Compute token similarity
 - For each pair of adjacent tokens: compute cosine similarity
 - Similarity matrix:
 - t₁-t₂: 0.95 (very similar)
 - t₂-t₃: 0.70 (moderately similar)
 - t₃-t₄: 0.30 (dissimilar)
 - t₄-t₅: 0.85 (similar)
 -...

Step 2: Decide which to merge
 - Use bipartite matching or greedy algorithm
 - Merge similar pairs

Step 3: Merge representations
 - t₁-t₂ → (t₁ + t₂) / 2 (average representation)
 - Keep t₃ (dissimilar to others)
 - t₄-t₅ → (t₄ + t₅) / 2
 -...

Output: [merged(t₁,t₂), t₃, merged(t₄,t₅), t₆, t₇, t₈]
Reduced from 8 to 6 tokens!

Bipartite Matching Algorithm

def token_merge(hidden_states, reduction_ratio=0.5):
 """
 Merge tokens using bipartite matching

 Args:
 hidden_states: (batch, seq_len, hidden_dim)
 reduction_ratio: Target reduction (0.5 = keep 50% of tokens)

 Returns:
 merged_states: (batch, merged_len, hidden_dim)
 indices: which tokens were kept/merged
 """
 batch, seq_len, hidden_dim = hidden_states.shape

 # Compute token similarity using cosine similarity
 normalized = F.normalize(hidden_states, dim=-1) # (batch, seq_len, hidden_dim)
 similarity = torch.matmul(normalized, normalized.transpose(-2, -1))
 # similarity shape: (batch, seq_len, seq_len)

 # Determine number of tokens to keep
 keep_token_count = int(seq_len * reduction_ratio)

 # Greedy bipartite matching
 merged_states = []
 merged_indices = []
 used = set()

 # Score each token by importance (max similarity to others)
 token_scores = similarity.mean(dim=-1) # Average similarity

 for i in range(seq_len):
 if i in used:
 continue

 if len(merged_states) >= keep_token_count:
 # Reached target number, merge remaining
 merged_states.append(hidden_states[:, i:i+1,:])
 used.add(i)
 continue

 # Find most similar unused token
 best_match = None
 best_similarity = -1

 for j in range(i+1, seq_len):
 if j not in used and similarity[:, i, j].mean() > best_similarity:
 best_similarity = similarity[:, i, j].mean()
 best_match = j

 if best_match is not None and best_similarity > threshold:
 # Merge i and best_match
 merged = (hidden_states[:, i:i+1,:] + 
 hidden_states[:, best_match:best_match+1,:]) / 2
 merged_states.append(merged)
 used.add(i)
 used.add(best_match)
 else:
 # Keep token alone
 merged_states.append(hidden_states[:, i:i+1,:])
 used.add(i)

 # Concatenate merged states
 merged = torch.cat(merged_states, dim=1) # (batch, merged_len, hidden_dim)
 return merged

# Usage in transformer:
class TokenMergingTransformerLayer(nn.Module):
 def __init__(self, hidden_dim, merge_ratio=0.5):
 super().__init__()
 self.merge_ratio = merge_ratio
 self.attention = MultiHeadAttention(hidden_dim)
 self.ffn = FeedForward(hidden_dim)

 def forward(self, x):
 # x: (batch, seq_len, hidden_dim)

 # Merge tokens BEFORE attention (expensive operation)
 x_merged = token_merge(x, self.merge_ratio)

 # Attention on merged sequence (cheaper!)
 attn_out = self.attention(x_merged)

 # Unmerge or interpolate back to original length
 # (implementation depends on tracking which tokens merged)
 x_out = unmerge_tokens(attn_out, original_length=x.shape[1])

 # FFN on original length
 out = self.ffn(x_out)

 return out

Merging Strategy

When to Merge

Option 1: Merge in early layers
 - Merge in layer 1-10
 - Early layers: more redundancy
 - Later layers: more specific representations
 - Strategy: Aggressive early merge (70% reduction)

Option 2: Merge progressively
 - Layer 1: Merge 30% (keep 70% of tokens)
 - Layer 5: Merge 30% (keep 70% of remaining)
 - Layer 10: Merge 30% (keep ~34% of original)
 - Strategy: Gradual reduction of sequence length

Option 3: Merge in middle layers
 - Layers 1-8: No merge (preserve early context)
 - Layers 9-16: Merge (reduce attention cost)
 - Layers 17-32: No merge (final layers process merged)

Recommendation:
 - Progressive merging works best
 - 2-3% quality loss vs 3-5% with aggressive early merge

Similarity Threshold

Threshold determines what counts as "similar enough to merge"

High threshold (0.9):
 - Only merge extremely similar tokens
 - Fewer merges, less speedup
 - Higher quality preservation

Low threshold (0.7):
 - Merge more aggressively
 - More speedup (larger sequence reduction)
 - Higher quality loss

Typical: 0.75-0.85
 - Balances speedup and quality
 - ~2x speedup, <1% quality loss

-

Performance Impact

Speedup Analysis

Scenario: 4096 token sequence, LLaMA 7B

Standard inference:
 - Attention computation: 50ms per layer × 32 = 1600ms
 - FFN computation: 30ms per layer × 32 = 960ms
 - Total: ~2.5 seconds per forward pass
 - Throughput: 400 tokens/sec per GPU

With Token Merging (50% reduction):
 - Merge 4096 → 2048 tokens (cost: ~5ms)
 - Attention: 12ms per layer × 32 = 384ms (4x reduction!)
 - FFN: 30ms per layer × 32 = 960ms (same, operates on 2048 anyway)
 - Unmerge back: ~5ms
 - Total: ~1.4 seconds
 - Speedup: 1.8x (close to 2x expected from O(N²) reduction)

With aggressive merging (70% reduction):
 - Merge 4096 → 1229 tokens
 - Attention: 4.4ms per layer × 32 = 141ms (10x reduction!)
 - Total: ~1.1 seconds
 - Speedup: 2.3x

Quality Impact

Experiment: LLaMA 7B on MMLU benchmark

Reduction Speedup MMLU Score Quality Loss
───────────────────────────────────────────────
None 1x 45.3% 0%
30% (keep 70%) 1.3x 45.2% 0.2%
50% (keep 50%) 1.8x 45.0% 0.7%
70% (keep 30%) 2.3x 44.5% 1.8%

Observation:
 - <1% quality loss for 2x speedup (excellent!)
 - Progressive merging better than aggressive
 - Works best with well-trained models

Different domains:
 - Code: Slightly higher quality loss (1.0% at 2x)
 - Creative: Lower quality loss (0.5% at 2x)
 - Reasoning: Higher quality loss (1.5% at 2x)
 - Merging removes information needed for reasoning

-

Unmerging Strategies

Simple Average

Merged output: [a', b', c'] (where a'=(a+b)/2)
Original: [a, b, c, d,...]

Unmerge by duplication:
 - a' → propagate to both a and b positions
 - b' → propagate to both c position
 - Result: Fill in original sequence length

Attention-Based Unmerging

Instead of simple duplication, use attention:
 - Merged token a' represents both a and b
 - Query: What should a and b be individually?
 - Use attention to distribute a' back to a and b positions
 - Smoother than simple duplication

Advantages and Limitations

Advantages

2-4x speedup without retraining
Works during inference only
Minimal quality loss (<1-2%)
Simple to implement
Compatible with other optimizations
Effective for long sequences

Limitations

Information loss from merging
Quality degrades for complex reasoning
Unmerging can be expensive if not careful
Less effective on short sequences
Sensitive to token similarity threshold

When to Use Token Merging

Use ToMe when:
Long sequences (>2K tokens)
Inference speed critical
Slightly lower quality acceptable
Attention cost is bottleneck

Avoid ToMe when:
Short sequences (<512 tokens)
Quality is paramount
Complex reasoning (merging breaks it)

Key Takeaways

Merge similar tokens to reduce sequence length O(N²) attention becomes O(M²) where M << N 2x speedup with <1% quality loss typical Progressive merging better than aggressive Most effective for long context inference

-