Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
past_helper.py150 linesDownload Raw Back to transformers
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 
codekingpro/portable-devtools · Team Ai