File size: 5,307 Bytes
f55039b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
"""Request-scoped, differentiable prefix sharing for the Qwen hybrid cache."""

import torch
from transformers.cache_utils import DynamicCache


def plan_prefix(inputs, positions):
    """Plan on CPU, before asynchronous transfer; candidates stay in each branch."""
    ids = inputs["input_ids"]
    cut = 0
    images = inputs.get("image_keys", ())
    if len(ids) > 1 and len(set(images)) <= 1:
        common = (ids == ids[:1]).all(0) & inputs["attention_mask"].bool().all(0)
        different = (~common).nonzero()
        length = int(different[0]) if len(different) else common.numel()
        cut = min(length, int(positions[:, 0].min()), ids.shape[1] - 2) // 64 * 64
    inputs["prefix_length"] = cut


class PrefixCache(DynamicCache):
    """Replace state tensors instead of overwriting values needed by backward."""

    def update_conv_state(
        self, conv_states, layer_idx, state_idx=0, conv_kernel_size=None, **kwargs
    ):
        if conv_kernel_size is None:
            raise ValueError("PrefixCache requires conv_kernel_size")
        layer = self.layers[layer_idx]
        layer.device, layer.dtype = conv_states.device, conv_states.dtype
        if layer.has_previous_state[state_idx]:
            conv_states = torch.cat([layer.conv_states[state_idx], conv_states], dim=-1)
        layer.conv_kernel_size[state_idx] = conv_kernel_size
        layer.conv_states[state_idx] = conv_states[..., -conv_kernel_size:]
        layer.has_previous_state[state_idx] = True
        layer.is_conv_states_initialized[state_idx] = True
        return conv_states

    def update_recurrent_state(self, recurrent_states, layer_idx, state_idx=0, **kwargs):
        layer = self.layers[layer_idx]
        layer.recurrent_states[state_idx] = recurrent_states
        layer.is_recurrent_states_initialized[state_idx] = True
        return recurrent_states


def language_forward(
    language,
    embeddings,
    attention_mask,
    position_ids,
    cut,
    capture_layers=(),
    capture_raw_layers=None,
):
    requested = list(capture_layers) if capture_layers else False
    capture_raw_layers = capture_raw_layers or {}

    def forward(
        inputs_embeds,
        attention_mask,
        position_ids,
        *,
        past_key_values=None,
        use_cache=False,
        capture=True,
    ):
        return language(
            inputs_embeds=inputs_embeds,
            attention_mask=attention_mask,
            position_ids=position_ids,
            past_key_values=past_key_values,
            use_cache=use_cache,
            output_hidden_states=requested if capture else False,
            return_dict=True,
        )

    def result(outputs, raw=None):
        if not capture_layers and not capture_raw_layers:
            return outputs.last_hidden_state
        raw = raw or {}
        return (
            outputs.last_hidden_state,
            *(outputs.hidden_states[index] for index in capture_layers),
            *(raw[index] for index in capture_raw_layers),
        )

    def captured_forward(inputs_embeds, attention_mask, position_ids, **kwargs):
        raw = {}
        handles = []
        if kwargs.pop("capture", True):
            for index, layer in capture_raw_layers.items():
                def save_raw(_module, _args, output, layer_index=index):
                    raw[layer_index] = output[0] if isinstance(output, tuple) else output

                handles.append(layer.register_forward_hook(save_raw))
        try:
            outputs = forward(inputs_embeds, attention_mask, position_ids, **kwargs)
        finally:
            for handle in handles:
                handle.remove()
        return result(outputs, raw)

    if not cut:
        return captured_forward(embeddings, attention_mask, position_ids)

    def shared_forward():
        # A mutable cache cannot survive a decoder-layer checkpoint replay.
        # Replay the complete request instead, constructing a fresh cache each time.
        layers = [m for m in language.modules() if getattr(m, "gradient_checkpointing", False)]
        for layer in layers:
            layer.gradient_checkpointing = False
        try:
            cache = PrefixCache(config=language.config)
            forward(
                embeddings[:1, :cut],
                attention_mask[:1, :cut],
                None if position_ids is None else position_ids[:, :1, :cut],
                past_key_values=cache,
                use_cache=True,
                capture=False,
            )
            # Transformers expects LongTensor; factories return Tensor with int64 dtype.
            indices = torch.zeros(len(embeddings), dtype=torch.long, device=embeddings.device)
            cache.reorder_cache(indices)  # ty: ignore[invalid-argument-type]
            return captured_forward(
                embeddings[:, cut:],
                attention_mask,
                None if position_ids is None else position_ids[:, :, cut:],
                past_key_values=cache,
                use_cache=True,
            )
        finally:
            for layer in layers:
                layer.gradient_checkpointing = True

    if torch.is_grad_enabled():
        from torch.utils.checkpoint import checkpoint

        return checkpoint(shared_forward, use_reentrant=False)
    return shared_forward()