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 +36 -0
- kvpress/per_layer_compression_wrapper.py +59 -0
- kvpress/pipeline.py +266 -0
- kvpress/presses/__init__.py +14 -0
- kvpress/presses/base_press.py +146 -0
- kvpress/presses/expected_attention_press.py +148 -0
- kvpress/presses/knorm_press.py +34 -0
- kvpress/presses/observed_attention_press.py +69 -0
- kvpress/presses/random_press.py +34 -0
- kvpress/presses/snapkv_press.py +107 -0
- kvpress/presses/streaming_llm_press.py +51 -0
- kvpress-0.0.1.dist-info/LICENSE +201 -0
- kvpress-0.0.1.dist-info/METADATA +203 -0
- kvpress-0.0.1.dist-info/RECORD +15 -0
- kvpress-0.0.1.dist-info/WHEEL +4 -0
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
|