kvpress 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
kvpress/__init__.py ADDED
@@ -0,0 +1,36 @@
1
+ # SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+
16
+ from kvpress.per_layer_compression_wrapper import apply_per_layer_compression
17
+ from kvpress.pipeline import KVPressTextGenerationPipeline
18
+ from kvpress.presses.base_press import BasePress
19
+ from kvpress.presses.expected_attention_press import ExpectedAttentionPress
20
+ from kvpress.presses.knorm_press import KnormPress
21
+ from kvpress.presses.observed_attention_press import ObservedAttentionPress
22
+ from kvpress.presses.random_press import RandomPress
23
+ from kvpress.presses.snapkv_press import SnapKVPress
24
+ from kvpress.presses.streaming_llm_press import StreamingLLMPress
25
+
26
+ __all__ = [
27
+ "BasePress",
28
+ "ExpectedAttentionPress",
29
+ "KnormPress",
30
+ "ObservedAttentionPress",
31
+ "RandomPress",
32
+ "SnapKVPress",
33
+ "StreamingLLMPress",
34
+ "KVPressTextGenerationPipeline",
35
+ "apply_per_layer_compression",
36
+ ]
@@ -0,0 +1,59 @@
1
+ # SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+
16
+ import logging
17
+ from typing import List
18
+
19
+ import torch
20
+
21
+ from kvpress.presses.base_press import BasePress
22
+
23
+ logger = logging.getLogger(__name__)
24
+
25
+
26
+ def apply_per_layer_compression(press: BasePress, compression_ratios: List[float]) -> BasePress:
27
+ """
28
+ Apply per-layer compression to a given press object.
29
+ This function wraps the forward hook of the press object to apply per-layer compression.
30
+
31
+ Parameters
32
+ ----------
33
+ press : BasePress
34
+ The press object to apply per-layer compression to.
35
+ compression_ratios : Dict[int, float]
36
+
37
+ Returns
38
+ -------
39
+ BasePress
40
+ The press object with per-layer compression applied.
41
+ """
42
+ press.compression_ratios = compression_ratios # type: ignore[attr-defined]
43
+ press.compression_ratio = None
44
+
45
+ logger.warning(
46
+ "Per layer compression wrapper is an experimental feature and only works with flash attention. "
47
+ "Please make sure that the model uses flash attention."
48
+ )
49
+
50
+ original_forward_hook = press.forward_hook
51
+
52
+ def _forward_hook(module: torch.nn.Module, input: list[torch.Tensor], kwargs: dict, output: list):
53
+ press.compression_ratio = press.compression_ratios[module.layer_idx] # type: ignore[attr-defined]
54
+ output = original_forward_hook(module, input, kwargs, output)
55
+ press.compression_ratio = None
56
+ return output
57
+
58
+ press.forward_hook = _forward_hook # type: ignore[method-assign]
59
+ return press
kvpress/pipeline.py ADDED
@@ -0,0 +1,266 @@
1
+ # SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+
16
+ import contextlib
17
+ import logging
18
+ from typing import Dict, List, Optional
19
+
20
+ import torch
21
+ from transformers import AutoModelForCausalLM, DynamicCache, Pipeline
22
+ from transformers.pipelines import PIPELINE_REGISTRY
23
+ from transformers.pipelines.base import GenericTensor
24
+
25
+ from kvpress.presses.base_press import BasePress
26
+ from kvpress.presses.observed_attention_press import ObservedAttentionPress
27
+
28
+ logger = logging.getLogger(__name__)
29
+
30
+
31
+ class KVPressTextGenerationPipeline(Pipeline):
32
+ """
33
+ Pipeline for key-value compression in causal language models.
34
+ This pipeline allows you to compress a long prompt using a key-value press
35
+ and then generate answers using greedy decoding.
36
+ """
37
+
38
+ def _sanitize_parameters(
39
+ self,
40
+ question: Optional[str] = None,
41
+ questions: Optional[List[str]] = None,
42
+ answer_prefix: Optional[str] = None,
43
+ press: Optional[BasePress] = None,
44
+ max_new_tokens: int = 50,
45
+ max_context_length: Optional[int] = None,
46
+ **kwargs,
47
+ ):
48
+ """
49
+ Sanitize the input parameters for the pipeline.
50
+ The user can either provide a single question or a list of questions to be asked about the context.
51
+
52
+ Parameters
53
+ ----------
54
+ question : str, optional
55
+ The question to be asked about the context. Exclusive with `questions`.
56
+ questions : List[str], optional
57
+ A list of questions to be asked about the context. Exclusive with `question`.
58
+ answer_prefix : str, optional
59
+ The prefix to be added to the generated answer.
60
+ press : BasePress, optional
61
+ The key-value press to use for compression.
62
+ max_new_tokens : int, optional
63
+ The maximum number of new tokens to generate for each answer.
64
+ max_context_length : int, optional
65
+ The maximum number of tokens in the context. By default will use the maximum length supported by the model.
66
+ **kwargs : dict
67
+ Additional keyword arguments, currently ignored.
68
+
69
+ Returns
70
+ -------
71
+ Tuple[Dict, Dict, Dict]
72
+ A tuple containing three dictionaries:
73
+ - preprocess_kwargs: The keyword arguments for the preprocess function.
74
+ - forward_kwargs: The keyword arguments for the forward function.
75
+ - postprocess_kwargs: The keyword arguments for the postprocess function.
76
+ """
77
+
78
+ answer_prefix = answer_prefix or ""
79
+ postprocess_kwargs = {"single_question": questions is None}
80
+ assert question is None or questions is None, "Either question or questions should be provided, not both."
81
+ questions = questions or ([question] if question else [""])
82
+ if max_context_length is None:
83
+ max_context_length = min(self.tokenizer.model_max_length, int(1e10)) # 1e10 to avoid overflow
84
+ preprocess_kwargs = {
85
+ "questions": questions,
86
+ "answer_prefix": answer_prefix,
87
+ "max_context_length": max_context_length,
88
+ }
89
+ forward_kwargs = {"press": press, "max_new_tokens": max_new_tokens}
90
+ return preprocess_kwargs, forward_kwargs, postprocess_kwargs
91
+
92
+ def preprocess(
93
+ self,
94
+ context: str,
95
+ questions: List[str],
96
+ answer_prefix: str,
97
+ max_context_length: int,
98
+ ):
99
+ """
100
+ Apply the chat template to the triplet (context, questions, answer_prefix) and tokenize it.
101
+
102
+ Returns
103
+ -------
104
+ Dict[str, GenericTensor]
105
+ A dictionary containing the tokenized context (key: "context_ids") and questions (key: "questions_ids").
106
+
107
+ """
108
+
109
+ # Apply chat template if available
110
+ if self.tokenizer.chat_template is None:
111
+ bos_token = getattr(self.tokenizer, "bos_token", "")
112
+ context = bos_token + context
113
+ question_suffix = "\n" # to separate the question from the answer
114
+ else:
115
+ separator = "\n" + "#" * len(context)
116
+ context = self.tokenizer.apply_chat_template(
117
+ [{"role": "user", "content": context + separator}], add_generation_prompt=True, tokenize=False
118
+ )
119
+ context, question_suffix = context.split(separator)
120
+
121
+ # Add question_suffix and answer prefix
122
+ # e.g. for llama3.1, question_suffix="<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n")
123
+ questions = [question + question_suffix + answer_prefix for question in questions]
124
+
125
+ # Tokenize the context and questions
126
+ context_ids = self.tokenizer.encode(context, return_tensors="pt", add_special_tokens=False)
127
+ question_ids = [
128
+ self.tokenizer.encode(question, return_tensors="pt", add_special_tokens=False) for question in questions
129
+ ]
130
+
131
+ # Truncate context
132
+ if context_ids.shape[1] > max_context_length:
133
+ logger.warning(
134
+ f"Context length has been truncated from {context_ids.shape[1]} to {max_context_length} tokens."
135
+ )
136
+ context_ids = context_ids[:, :max_context_length]
137
+
138
+ return {"context_ids": context_ids, "questions_ids": question_ids}
139
+
140
+ def _forward(
141
+ self, input_tensors: Dict[str, GenericTensor], max_new_tokens: int = 50, press: Optional[BasePress] = None
142
+ ):
143
+ """
144
+ Forward pass of the kv-press pipeline.
145
+
146
+ Parameters
147
+ ----------
148
+ input_tensors : Dict[str, GenericTensor]
149
+ A dictionary containing the tokenized context and questions.
150
+ max_new_tokens : int, optional
151
+ The maximum number of new tokens to generate for each answer. Defaults to 50.
152
+ press : BasePress, optional
153
+ The key-value press to use for compression. Defaults to None.
154
+
155
+ Returns
156
+ -------
157
+ List[str]
158
+ A list of generated answers.
159
+ """
160
+
161
+ context_ids = input_tensors["context_ids"].to(self.model.device)
162
+ context_length = context_ids.shape[1]
163
+
164
+ # Prefilling using the press on the context
165
+ with press(self.model) if press is not None else contextlib.nullcontext():
166
+ past_key_values = self.model(
167
+ input_ids=context_ids,
168
+ past_key_values=DynamicCache(),
169
+ output_attentions=isinstance(press, ObservedAttentionPress),
170
+ num_logits_to_keep=1,
171
+ ).past_key_values
172
+
173
+ logger.debug(f"Context Length: {context_length}")
174
+ logger.debug(f"Compressed Context Length: {past_key_values.get_seq_length()}")
175
+
176
+ # Greedy decoding for each question
177
+ answers = []
178
+ for question_ids in input_tensors["questions_ids"]:
179
+ answer = self.generate_answer(
180
+ question_ids=question_ids.to(self.model.device),
181
+ past_key_values=past_key_values,
182
+ context_length=context_length,
183
+ max_new_tokens=max_new_tokens,
184
+ )
185
+ answers.append(answer)
186
+
187
+ return answers
188
+
189
+ def postprocess(self, model_outputs, single_question):
190
+ if single_question:
191
+ return {"answer": model_outputs[0]}
192
+ return {"answers": model_outputs}
193
+
194
+ def generate_answer(
195
+ self, question_ids: torch.Tensor, past_key_values: DynamicCache, context_length: int, max_new_tokens: int
196
+ ) -> str:
197
+ """
198
+ Generate an answer to a question using greedy decoding.
199
+
200
+ Parameters
201
+ ----------
202
+ question_ids : torch.Tensor
203
+ The tokenized question.
204
+ past_key_values : DynamicCache
205
+ The compressed key-value cache.
206
+ context_length : int
207
+ The length of the context.
208
+ max_new_tokens : int
209
+ The maximum number of new tokens to generate.
210
+
211
+ Returns
212
+ -------
213
+ str
214
+ The generated answer.
215
+ """
216
+
217
+ cache_seq_lengths = [
218
+ past_key_values.get_seq_length(layer_idx=layer_idx) for layer_idx in range(len(past_key_values))
219
+ ]
220
+
221
+ position_ids = torch.arange(
222
+ context_length, context_length + question_ids.shape[1], device=self.model.device
223
+ ).unsqueeze(0)
224
+
225
+ # if the user doesn't provide a question, skip forward pass
226
+ outputs = self.model(
227
+ input_ids=question_ids.to(self.model.device),
228
+ past_key_values=past_key_values,
229
+ position_ids=position_ids,
230
+ num_logits_to_keep=1,
231
+ )
232
+
233
+ position_ids = position_ids[:, -1:] + 1
234
+ generated_ids = [outputs.logits[0, -1].argmax()]
235
+
236
+ should_stop_token_ids = self.model.generation_config.eos_token_id
237
+ if not isinstance(should_stop_token_ids, list):
238
+ should_stop_token_ids = [should_stop_token_ids]
239
+
240
+ for i in range(max_new_tokens - 1):
241
+ outputs = self.model(
242
+ input_ids=generated_ids[-1].unsqueeze(0).unsqueeze(0),
243
+ past_key_values=outputs.past_key_values,
244
+ position_ids=position_ids + i,
245
+ )
246
+ new_id = outputs.logits[0, -1].argmax()
247
+ generated_ids.append(new_id)
248
+ if new_id.item() in should_stop_token_ids:
249
+ break
250
+ answer = self.tokenizer.decode(torch.stack(generated_ids), skip_special_tokens=True)
251
+
252
+ # remove the generated tokens from the cache
253
+ past_key_values.key_cache = [
254
+ key[:, :, :cache_seq_len] for key, cache_seq_len in zip(past_key_values.key_cache, cache_seq_lengths)
255
+ ]
256
+ past_key_values.value_cache = [
257
+ value[:, :, :cache_seq_len] for value, cache_seq_len in zip(past_key_values.value_cache, cache_seq_lengths)
258
+ ]
259
+ return answer
260
+
261
+
262
+ PIPELINE_REGISTRY.register_pipeline(
263
+ "kv-press-text-generation",
264
+ pipeline_class=KVPressTextGenerationPipeline,
265
+ pt_model=AutoModelForCausalLM,
266
+ )
@@ -0,0 +1,14 @@
1
+ # SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
@@ -0,0 +1,146 @@
1
+ # SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+
16
+ import logging
17
+ from contextlib import contextmanager
18
+ from typing import Dict, Generator
19
+
20
+ import torch
21
+ from torch import nn
22
+ from transformers import LlamaForCausalLM, MistralForCausalLM, Phi3ForCausalLM, PreTrainedModel, Qwen2ForCausalLM
23
+
24
+ logger = logging.getLogger(__name__)
25
+
26
+
27
+ class BasePress:
28
+ """Base class for pruning methods.
29
+ Each pruning method should implement a `score` method that computes the scores for each KV pair in a layer.
30
+ This score is used to prune the KV pairs with the lowest scores in the `hook` method
31
+ The `hook` method is called after the forward pass of a layer and updates the cache with the pruned KV pairs.
32
+ The press can be applied to a model by calling it with the model as an argument.
33
+ """
34
+
35
+ def __init__(self, compression_ratio: float = 0.0):
36
+ self.compression_ratio = compression_ratio
37
+ assert 0 <= compression_ratio < 1, "Compression ratio must be between 0 and 1"
38
+
39
+ def score(
40
+ self,
41
+ module: nn.Module,
42
+ hidden_states: torch.Tensor,
43
+ keys: torch.Tensor,
44
+ values: torch.Tensor,
45
+ attentions: torch.Tensor,
46
+ kwargs,
47
+ ) -> torch.Tensor:
48
+ """Compute the scores for each KV pair in the layer.
49
+
50
+ Parameters
51
+ ----------
52
+ module :
53
+ Transformer layer, see `hook` method for more details.
54
+ hidden_states :
55
+ Hidden states of the layer.
56
+ keys :
57
+ Keys of the cache. Note keys are after RoPE.
58
+ values :
59
+ Values of the cache.
60
+ attentions :
61
+ Attention weights of the layer.
62
+ kwargs :
63
+ Keyword arguments, as given to the forward pass of the layer.
64
+
65
+ Returns
66
+ -------
67
+ Scores for each KV pair in the layer, shape keys.shape[:-1].
68
+
69
+ """
70
+ raise NotImplementedError
71
+
72
+ def forward_hook(self, module: nn.Module, input: list[torch.Tensor], kwargs: Dict, output: list):
73
+ """Cache compression hook called after the forward pass of a decoder layer.
74
+ The hook is applied only during the pre-filling phase if there is some pruning ratio.
75
+ The current implementation only allows to remove a constant number of KV pairs.
76
+
77
+ Parameters
78
+ ----------
79
+ module :
80
+ Transformer attention layer.
81
+ input :
82
+ Input to the hook. This is the input to the forward pass of the layer.
83
+ kwargs :
84
+ Keyword arguments, as given to the forward pass of the layer.
85
+ output :
86
+ Output of the hook. This is the original output of the forward pass of the layer.
87
+
88
+ Returns
89
+ -------
90
+ Modified output of the forward pass of the layer.
91
+
92
+ """
93
+ # See e.g. LlamaDecoderLayer.forward for the output structure
94
+ if len(output) == 3:
95
+ _, attentions, cache = output
96
+ else:
97
+ attentions, cache = None, output[-1]
98
+
99
+ hidden_states = kwargs["hidden_states"]
100
+ q_len = hidden_states.shape[1]
101
+
102
+ # Don't compress if the compression ratio is 0 or this is not pre-filling
103
+ if (self.compression_ratio == 0) or (cache.seen_tokens > q_len):
104
+ return output
105
+
106
+ keys = cache.key_cache[module.layer_idx]
107
+ values = cache.value_cache[module.layer_idx]
108
+
109
+ with torch.no_grad():
110
+ scores = self.score(module, hidden_states, keys, values, attentions, kwargs)
111
+
112
+ # Prune KV pairs with the lowest scores
113
+ n_kept = int(q_len * (1 - self.compression_ratio))
114
+ indices = scores.topk(n_kept, dim=-1).indices
115
+ indices = indices.unsqueeze(-1).expand(-1, -1, -1, module.head_dim)
116
+
117
+ # Update cache
118
+ cache.key_cache[module.layer_idx] = keys.gather(2, indices)
119
+ cache.value_cache[module.layer_idx] = values.gather(2, indices)
120
+
121
+ return output
122
+
123
+ @contextmanager
124
+ def __call__(self, model: PreTrainedModel) -> Generator:
125
+ """
126
+ Context manager to apply a compression method to a model.
127
+ Apply this context manager during the pre-filling phase to compress the context.
128
+
129
+ Parameters
130
+ ----------
131
+ model : PreTrainedModel
132
+ Model to apply the compression method to
133
+ """
134
+
135
+ if not isinstance(model, (LlamaForCausalLM, MistralForCausalLM, Phi3ForCausalLM, Qwen2ForCausalLM)):
136
+ logger.warning(f"Model {type(model)} not tested")
137
+
138
+ try:
139
+ hooks = []
140
+ for layer in model.model.layers:
141
+ hooks.append(layer.self_attn.register_forward_hook(self.forward_hook, with_kwargs=True))
142
+
143
+ yield
144
+ finally:
145
+ for forward_hook in hooks:
146
+ forward_hook.remove()
@@ -0,0 +1,148 @@
1
+ # SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+
16
+ import inspect
17
+ import math
18
+ from dataclasses import dataclass
19
+
20
+ import torch
21
+ import torch.nn.functional as F
22
+ from torch import nn
23
+ from transformers.models.llama.modeling_llama import repeat_kv
24
+
25
+ from kvpress.presses.base_press import BasePress
26
+
27
+
28
+ @dataclass
29
+ class ExpectedAttentionPress(BasePress):
30
+ """
31
+ Compute scores based on the expected attention on next positions. To do so
32
+ 1. Compute the mean and covariance matrix of the queries before RoPE.
33
+ 2. Compute the RoPE rotation matrix R on next n_future_positions and average it
34
+ 3. Apply R to the mean and covariance matrice of the queries.
35
+ 4. As attention A = exp(Q @ K / sqrt(d)), we compute the expected attention
36
+ E(A) = exp(K @ mean.T / sqrt(d) + 1/2 K @ cov @ K.T / d)
37
+ 5. Rescale the scores by the norm of the values
38
+ The first n_sink tokens are removed from calculations (sink attention phenomenon).
39
+ """
40
+
41
+ compression_ratio: float = 0.0
42
+ n_future_positions: int = 512
43
+ n_sink: int = 4
44
+ use_covariance: bool = True
45
+ use_vnorm: bool = True
46
+
47
+ def get_query_statistics(self, module: nn.Module, hidden_states: torch.Tensor):
48
+ """
49
+ Compute the mean and covariance matrix of the queries
50
+ """
51
+
52
+ bsz, q_len, _ = hidden_states.shape
53
+ n, d = module.num_heads, module.head_dim
54
+
55
+ # Remove first hidden_states that likely contain outliers
56
+ h = hidden_states[:, self.n_sink :]
57
+
58
+ if hasattr(module, "q_proj"):
59
+ Wq = module.q_proj.weight
60
+ elif hasattr(module, "qkv_proj"):
61
+ Wq = module.qkv_proj.weight[: n * d]
62
+ else:
63
+ raise NotImplementedError(f"ExpectedAttentionPress not yet implemented for {module.__class__}.")
64
+
65
+ # Query mean
66
+ mean_h = torch.mean(h, dim=1, keepdim=True)
67
+ mu = torch.matmul(mean_h, Wq.T).squeeze(1)
68
+ mu = mu.view(bsz, n, d)
69
+
70
+ # Query covariance
71
+ cov = None
72
+ if self.use_covariance:
73
+ h = h - mean_h
74
+ cov = torch.matmul(h.transpose(1, 2), h) / h.shape[1]
75
+ cov = torch.matmul(Wq, torch.matmul(cov, Wq.T)) # TODO: not optimal
76
+ cov = cov.view(bsz, n, d, n, d).diagonal(dim1=1, dim2=3)
77
+ cov = cov.permute(0, 3, 1, 2)
78
+
79
+ # RoPE rotation matrix on next n_future_positions
80
+ if "position_ids" in inspect.signature(module.rotary_emb.forward).parameters:
81
+ position_ids = torch.arange(q_len, q_len + self.n_future_positions).unsqueeze(0).to(mu.device)
82
+ cos, sin = module.rotary_emb(mu, position_ids)
83
+ cos, sin = cos[0], sin[0]
84
+ else:
85
+ cos, sin = module.rotary_emb(mu, q_len + self.n_future_positions)
86
+ cos, sin = cos[q_len:], sin[q_len:]
87
+
88
+ Id = torch.eye(d, device=cos.device, dtype=cos.dtype)
89
+ P = torch.zeros((d, d), device=cos.device, dtype=cos.dtype)
90
+ P[d // 2 :, : d // 2], P[: d // 2, d // 2 :] = torch.eye(d // 2), -torch.eye(d // 2)
91
+ R = cos.unsqueeze(1) * Id + sin.unsqueeze(1) * P
92
+
93
+ # Apply average rotation to the mean and covariance
94
+ R = R.mean(dim=0)
95
+ mu = torch.matmul(mu, R.T)
96
+ if self.use_covariance:
97
+ cov = torch.matmul(R, torch.matmul(cov, R.T))
98
+
99
+ # Instead of using the average rotation matrix, we could use a mixture of gaussian statistics to
100
+ # estimate mean and covariance. Estimation is better, but end-to-end performance was lower.
101
+ # mu = torch.einsum("bhj, fij -> bhfi", mu, R)
102
+ # mean_mu = mu.mean(dim=2, keepdim=True)
103
+ # if self.use_covariance:
104
+ # cov = torch.einsum("fki, bhkl, fjl -> bhfij", R, cov, R)
105
+ # cov = cov.mean(dim=2)
106
+ # cov += torch.einsum("bhfi, bhfj -> bhji", mu - mean_mu, mu - mean_mu) / self.n_future_positions
107
+ # mu = mean_mu.squeeze(2)
108
+
109
+ return mu, cov
110
+
111
+ def score(
112
+ self,
113
+ module: nn.Module,
114
+ hidden_states: torch.Tensor,
115
+ keys: torch.Tensor,
116
+ values: torch.Tensor,
117
+ attentions: torch.Tensor,
118
+ kwargs,
119
+ ) -> torch.Tensor:
120
+
121
+ # Remove sink tokens
122
+ assert keys.size(2) > self.n_sink, f"Input should contain more tokens than n_sink={self.n_sink}"
123
+ keys = keys[:, :, self.n_sink :]
124
+ values = values[:, :, self.n_sink :]
125
+
126
+ # Compute query statistics
127
+ mean_query, cov_query = self.get_query_statistics(module, hidden_states)
128
+
129
+ # Compute scores
130
+ bsz, num_key_value_heads, q_len, d = keys.shape
131
+ keys = repeat_kv(keys, module.num_key_value_groups).transpose(2, 3)
132
+ scores = torch.matmul(mean_query.unsqueeze(2), keys).squeeze(2) / math.sqrt(d)
133
+ if self.use_covariance:
134
+ scores += torch.einsum("bhin, bhij, bhjn->bhn", keys, cov_query, keys) / d / 2
135
+ scores = F.softmax(scores, dim=-1)
136
+
137
+ # Average scores across groups
138
+ scores = scores.view(bsz, num_key_value_heads, module.num_key_value_groups, q_len)
139
+ scores = scores.mean(dim=2)
140
+
141
+ # Rescale scores by the norm of the values
142
+ if self.use_vnorm:
143
+ scores = scores * values.norm(dim=-1)
144
+
145
+ # Add back the sink tokens. Use max score to make sure they are not pruned.
146
+ scores = F.pad(scores, (self.n_sink, 0), value=scores.max().item())
147
+
148
+ return scores