File size: 3,806 Bytes
2db32a1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
# 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))