codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License. See License.txt in the project root for
4# license information.
5# --------------------------------------------------------------------------
6
7import logging
8
9import torch
10
11logger = logging.getLogger(__name__)
12
13
14class PastKeyValuesHelper:
15 """Helper functions to process past key values for encoder-decoder model"""
16
17 @staticmethod
18 def get_past_names(num_layers, present: bool = False):
19 past_self_names = []
20 past_cross_names = []
21 for i in range(num_layers):
22 past_self_names.extend(
23 [f"present_key_self_{i}", f"present_value_self_{i}"]
24 if present
25 else [f"past_key_self_{i}", f"past_value_self_{i}"]
26 )
27 past_cross_names.extend(
28 [f"present_key_cross_{i}", f"present_value_cross_{i}"]
29 if present
30 else [f"past_key_cross_{i}", f"past_value_cross_{i}"]
31 )
32 return past_self_names + past_cross_names
33
34 @staticmethod
35 def group_by_self_or_cross(present_key_values):
36 """Split present state from grouped by layer to grouped by self/cross attention.
37 Before: (past_key_self_0, past_value_self_0, past_key_cross_0, past_value_cross_0), (past_key_self_1, past_value_self_1, past_key_cross_1, past_value_cross_1), ...
38 After: (past_key_self_0, past_value_self_0, past_key_self_1, past_value_self_1, ...), (past_key_cross_0, past_value_cross_0, past_key_cross_1, past_value_cross_1, ...)
39
40 """
41 present_self = []
42 present_cross = []
43 for _i, present_layer_i in enumerate(present_key_values):
44 assert len(present_layer_i) == 4, f"Expected to have four items. Got {len(present_layer_i)}"
45 (
46 present_key_self,
47 present_value_self,
48 present_key_cross,
49 present_value_cross,
50 ) = present_layer_i
51 present_self.extend([present_key_self, present_value_self])
52 present_cross.extend([present_key_cross, present_value_cross])
53 return present_self, present_cross
54
55 @staticmethod
56 def group_by_layer(past, num_layers):
57 """Reorder past state from grouped by self/cross attention to grouped by layer.
58 Before: past_key_self_0, past_value_self_0, past_key_self_1, past_value_self_1, ..., past_key_cross_0, past_value_cross_0, past_key_cross_1, past_value_cross_1, ...
59 After: (past_key_self_0, past_value_self_0, past_key_cross_0, past_value_cross_0), (past_key_self_1, past_value_self_1, past_key_cross_1, past_value_cross_1),
60 """
61 assert len(past) == 4 * num_layers
62 return tuple(
63 [
64 past[2 * i],
65 past[2 * i + 1],
66 past[2 * num_layers + 2 * i],
67 past[2 * num_layers + 2 * i + 1],
68 ]
69 for i in range(num_layers)
70 )
71
72 @staticmethod
73 def back_group_by_layer(past_key_values: tuple[tuple[torch.Tensor]]):
74 """Categorize present_key_values from self and cross attention to layer by layer.
75
76 Reorder past state from grouped by self/cross attention to grouped by layer.
77 Before: past_key_self_0, past_value_self_0, past_key_self_1, past_value_self_1, ...,
78 past_key_cross_0, past_value_cross_0, past_key_cross_1, past_value_cross_1, ...
79 After: (past_key_self_0, past_value_self_0, past_key_cross_0, past_value_cross_0),
80 (past_key_self_1, past_value_self_1, past_key_cross_1, past_value_cross_1),
81
82 Args:
83 present_key_values: From past_key_values of a model (group by self and cross attention)
84
85 Returns:
86 past_tuples: present key and values grouped by layer.
87 """
88 past_tuples = ()
89 half_idx = len(past_key_values) // 2
90 for i in range(len(past_key_values) // 4):
91 idx = 2 * i
92 past_tuples += (
93 (
94 past_key_values[idx],
95 past_key_values[idx + 1],
96 past_key_values[half_idx + idx],
97 past_key_values[half_idx + idx + 1],
98 ),
99 )
100 return past_tuples
101
102 @staticmethod
103 def group_by_self_and_cross(present_key_values: tuple[torch.Tensor], concat: bool = False):
104 """Categorize present_key_values into self and cross attention.
105
106 Split present state from grouped by layer to grouped by self/cross attention.
107 Before: (past_key_self_0, past_value_self_0, past_key_cross_0, past_value_cross_0),
108 (past_key_self_1, past_value_self_1, past_key_cross_1, past_value_cross_1), ...
109 After: (past_key_self_0, past_value_self_0, past_key_self_1, past_value_self_1, ...),
110 (past_key_cross_0, past_value_cross_0, past_key_cross_1, past_value_cross_1, ...)
111
112 Args:
113 present_key_values: From past_key_values of a model (group by layer)
114 concat: If concat self attention with cross attention key/value to return
115
116 Returns:
117 present_self (Tuple[torch.Tensor]): present key and values from self attention
118 present_cross (Tuple[torch.Tensor]): present key and values from cross attention
119 """
120 present_self: list[torch.Tensor] = []
121 present_cross: list[torch.Tensor] = []
122 for _, present_layer_i in enumerate(present_key_values):
123 assert len(present_layer_i) == 4, f"Expected to have four items. Got {len(present_layer_i)}"
124 present_key_self, present_value_self, present_key_cross, present_value_cross = present_layer_i
125 present_self.extend([present_key_self, present_value_self])
126 present_cross.extend([present_key_cross, present_value_cross])
127 if concat:
128 return present_self + present_cross
129 else:
130 return present_self, present_cross
131
132 @staticmethod
133 def get_input_names(past_key_values: tuple[tuple[torch.Tensor]], encoder=True):
134 """Process input names of model wrapper.
135
136 Args:
137 past_key_values: Consider `self` and `cross` past_key_values
138
139 Returns:
140 names (List[string]): input names
141 """
142 names = []
143 num_layers = len(past_key_values) // 4 if encoder else len(past_key_values)
144 prefix = "past_" if not encoder else "present_"
145 for i in range(num_layers):
146 names.extend([prefix + s for s in [f"key_self_{i}", f"value_self_{i}"]])
147 for i in range(num_layers):
148 names.extend([prefix + s for s in [f"key_cross_{i}", f"value_cross_{i}"]])
149 return names
150 