aryuemaan's picture
Upload folder using huggingface_hub
2db32a1 verified
Raw
History Blame Contribute Delete
3.81 kB
# coding=utf-8
"""Adaptive Token Merging (ATM).
Implements Section VIII of the Wiola paper. Adjacent tokens whose cosine
similarity exceeds ``tau`` are greedily merged (averaged) in a left-to-right,
non-overlapping scan. A merge map records the source positions so the original
sequence length can be restored exactly after attention.
ATM is **training-only**: it is disabled at inference to keep the KV-cache
length consistent with the un-merged sequence.
Because sequences in a batch may merge by different amounts, merged sequences
are right-padded to the batch maximum and a boolean ``keep_mask`` marks the real
(non-padded) merged positions. Unmerge ignores padded slots.
"""
from typing import List, Tuple
import torch
def _greedy_merge_one(sim: torch.Tensor, threshold: float) -> List[Tuple[int, ...]]:
"""Greedy non-overlapping merge decisions for a single sequence.
Args:
sim: cosine similarities between adjacent tokens, shape [T-1].
threshold: tau.
Returns:
A list of groups; each group is a tuple of 1 or 2 source indices.
"""
seq_len = sim.shape[0] + 1
groups: List[Tuple[int, ...]] = []
i = 0
while i < seq_len:
if i < seq_len - 1 and sim[i].item() > threshold:
groups.append((i, i + 1))
i += 2
else:
groups.append((i,))
i += 1
return groups
def merge_tokens(hidden_states: torch.Tensor, threshold: float):
"""Merge adjacent redundant tokens.
Args:
hidden_states: [B, T, d].
threshold: tau (cosine similarity merge threshold).
Returns:
merged: [B, T_prime_max, d] (right padded with zeros).
keep_mask: [B, T_prime_max] bool, True for real merged tokens.
merge_maps: list (len B) of lists of source-index tuples.
"""
bsz, seq_len, dim = hidden_states.shape
if seq_len < 2:
keep_mask = torch.ones(bsz, seq_len, dtype=torch.bool, device=hidden_states.device)
merge_maps = [[(t,) for t in range(seq_len)] for _ in range(bsz)]
return hidden_states, keep_mask, merge_maps
normed = torch.nn.functional.normalize(hidden_states, dim=-1, eps=1e-8)
# rho_t = <x_hat_t, x_hat_{t+1}> -> [B, T-1]
sim = (normed[:, :-1] * normed[:, 1:]).sum(-1)
merge_maps = [_greedy_merge_one(sim[b], threshold) for b in range(bsz)]
new_len = max(len(g) for g in merge_maps)
merged = hidden_states.new_zeros(bsz, new_len, dim)
keep_mask = torch.zeros(bsz, new_len, dtype=torch.bool, device=hidden_states.device)
for b, groups in enumerate(merge_maps):
for k, grp in enumerate(groups):
if len(grp) == 2:
merged[b, k] = 0.5 * (hidden_states[b, grp[0]] + hidden_states[b, grp[1]])
else:
merged[b, k] = hidden_states[b, grp[0]]
keep_mask[b, k] = True
return merged, keep_mask, merge_maps
def unmerge_tokens(
merged: torch.Tensor, merge_maps: List[List[Tuple[int, ...]]], original_len: int
) -> torch.Tensor:
"""Restore original sequence length by broadcasting each merged token back
to its source positions (Eq. 31)."""
bsz, _, dim = merged.shape
out = merged.new_zeros(bsz, original_len, dim)
for b, groups in enumerate(merge_maps):
for k, grp in enumerate(groups):
for src in grp:
out[b, src] = merged[b, k]
return out
def merge_ratio(merge_maps: List[List[Tuple[int, ...]]], original_len: int) -> float:
"""Average merge ratio mu = 1 - T'/T across the batch."""
if original_len == 0:
return 0.0
ratios = [1.0 - len(g) / original_len for g in merge_maps]
return float(sum(ratios) / len(ratios))