A Sequential Parallel Method, Device, Equipment and Medium for Linear Attention
By distributing exception long sequences to multiple processing devices and forward and backpropagating on each device, the memory constraint problem is solved, and the processing efficiency and expansion ability of sequence length is improved.
Patent Information
- Application Number
- CN202410293061.5
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2024-03-14
- Publication Date
- 2025-06-17
- Estimated Expiration
- 2044-03-14
AI Technical Summary
The prior art, when dealing with exceptional long sequences, has poor parallel efficiency and the availability of linear attention-based language models due to memory constraints.
The first processing device distributes multiple subsequences corresponding to the original sequence on the linear transformer in the distributed environment to the corresponding second processing device according to the pre-configured data distribution strategy, and the forward propagation and backpropagation methods of the subsequence are determined on the second processing device to update the parameters of the subsequence.
Distributed processing is implemented, and the processing efficiency of the sequence is improved, making the implementation on the cluster of the second processing device more hardware-friendly, and can expand the sequence length by 8 times and be faster.
Smart Images

Figure CN118132155B_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the field of artificial intelligence technology, and particularly to a sequence parallel method, device, equipment and medium for linear attention. Background Art
[0002] Sequence parallelism (SP) is a common strategy for processing long sequences that exceed the memory limit of a single GPU. However, existing SP methods do not fully utilize the characteristics of linear attention, resulting in poor parallel efficiency and usability of language models based on linear attention. Summary of the Invention
[0003] The present invention provides a sequence parallel method, device, equipment and medium for linear attention to solve the technical problem of memory limitation when using a single device to calculate an extremely long sequence in the prior art.
[0004] According to one aspect of the present invention, there is provided a sequence parallel method for linear attention, including:
[0005] Distributing, by a first processing device, a plurality of subsequences corresponding to an original sequence on a linear transformer in a distributed environment to corresponding second processing devices according to a pre-configured data distribution strategy;
[0006] Determining, by the second processing device, a forward total output matrix corresponding to the subsequence by using a pre-configured forward propagation method;
[0007] Determining, by the second processing device, a parameter gradient corresponding to the subsequence by using a pre-configured backward propagation method and the forward total output matrix, so as to update the parameters of the corresponding subsequence according to the parameter gradient.
[0008] According to another aspect of the present invention, there is provided a sequence parallel device for linear attention, including:
[0009] A distribution module, configured to distribute, by a first processing device, a plurality of subsequences corresponding to an original sequence on a linear transformer in a distributed environment to corresponding second processing devices according to a pre-configured data distribution strategy;
[0010] A first determination module, configured to determine, by the second processing device, a forward total output matrix corresponding to the subsequence by using a pre-configured forward propagation method;
[0011] A second determination module, configured to determine, by the second processing device, a parameter gradient corresponding to the subsequence by using a pre-configured backward propagation method and the forward total output matrix, so as to update the parameters of the corresponding subsequence according to the parameter gradient.
[0012] According to another aspect of the present invention, there is provided an electronic device, which includes:
[0013] at least one processor; and a memory communicatively connected to the at least one processor; wherein, the memory stores a computer program executable by the at least one processor, and when the computer program is executed by the at least one processor, the at least one processor is enabled to execute the sequence parallel method for linear attention according to any embodiment of the present invention.
[0014] According to another aspect of the present invention, there is provided a computer-readable storage medium storing computer instructions for causing a processor to implement the sequence parallel method for linear attention according to any embodiment of the present invention when executed.
[0015] In the technical solution of the embodiment of the present invention, a first processing device distributes a plurality of subsequences corresponding to an original sequence on a linear transformer in a distributed environment to corresponding second processing devices according to a pre-configured data distribution strategy; a second processing device determines a forward total output matrix corresponding to the subsequence by using a pre-configured forward propagation method; and the second processing device determines a parameter gradient corresponding to the subsequence by using a pre-configured backward propagation method and the forward total output matrix, so as to update the parameters of the corresponding subsequence according to the parameter gradient. In the technical solution of the present invention, a long original sequence is split into a plurality of subsequences according to a pre-configured data distribution strategy and distributed to corresponding second processing devices to implement a distributed processing process, and intermediate states are exchanged during the forward propagation and backward propagation within or between a plurality of second processing devices, improving the processing efficiency of the sequence, and thus making the implementation on a cluster of second processing devices more hardware-friendly.
[0016] It should be understood that the content described in this part is not intended to identify the key or important features of the embodiments of the present invention, nor is it used to limit the scope of the present invention. Other features of the present invention will become easily understood through the following description. BRIEF DESCRIPTION OF THE DRAWINGS
[0017] In order to more clearly illustrate the technical solutions in the embodiments of the present invention, the drawings required for the description of the embodiments will be briefly introduced below. Obviously, the drawings in the following description are only some embodiments of the present invention, and those of ordinary skill in the art can also obtain other drawings without creative efforts based on these drawings.
[0018] Figure 1 is a flowchart of a sequence parallel method for linear attention provided by an embodiment of the present invention;
[0019] Figure 2 It is a visualization block diagram of LASP provided by an embodiment of the present invention;
[0020] Figure 3 It is a flowchart of another sequence parallel method for linear attention provided by an embodiment of the present invention;
[0021] Figure 4 It is a flowchart of the implementation of data distribution provided by an embodiment of the present invention;
[0022] Figure 5 It is an example diagram of data distribution in LASP provided by an embodiment of the present invention;
[0023] Figure 6 It is a flowchart of the implementation of a forward propagation method provided by an embodiment of the present invention;
[0024] Figure 7 It is a flowchart of the implementation of a backward propagation method provided by an embodiment of the present invention;
[0025] Figure 8 It is a schematic structural diagram of a sequence parallel device for linear attention provided by an embodiment of the present invention;
[0026] Figure 9 It is a structural block diagram of an electronic device provided by an embodiment of the present invention. Detailed implementation manners
[0027] In order to enable those skilled in the art to better understand the solution of the present invention, the technical solutions in the embodiments of the present invention will be clearly and completely described below with reference to the accompanying drawings in the embodiments of the present invention. Obviously, the described embodiments are only a part of the embodiments of the present invention, rather than all the embodiments. Based on the embodiments of the present invention, all other embodiments obtained by those of ordinary skill in the art without creative efforts shall fall within the protection scope of the present invention.
[0028] It should be noted that the terms "first", "second", etc. in the specification and claims of the present invention and the above-mentioned drawings are used to distinguish similar objects, and do not necessarily need to be used to describe a specific order or sequence. It should be understood that such used data can be interchanged under appropriate circumstances so that the embodiments of the present invention described here can be implemented in an order different from those illustrated or described here. In addition, the terms "comprising" and "having" and any variations thereof are intended to cover non-exclusive inclusion. For example, a process, method, system, product or device including a series of steps or units does not necessarily have to be limited to those steps or units clearly listed, but may include other steps or units not clearly listed or inherent to these processes, methods, products or devices.
[0029] The present invention designs an efficient point-to-point communication mechanism, which uses the right multiplication kernel technique of linear attention to greatly reduce the communication overhead of SP. It also improves the actual efficiency of LASP by performing kernel fusion and intermediate state caching, making the implementation of LASP on GPU clusters more hardware-friendly. In addition, the compatibility of sequence-level LASP with all types of batch-level data parallel methods is carefully ensured, which is crucial for distributed training on large clusters with long sequences and large batches. A large number of experiments are conducted on two linear attention-based models, covering different sequence lengths and GPU cluster scales. For the 1B model using 128 A100 80G GPUs, LASP can extend the sequence length to 4096K, which is 8 times longer than the existing SP method and faster at the same time.
[0030] In the present invention, a linear attention sequence parallel (LASP) technique applicable to linear transformers is proposed for achieving efficient sequence parallelism. The method includes a complex communication mechanism based on point-to-point communication for exchanging intermediate states during the forward and backward passes within a node or across multiple nodes. This design maximizes the utilization of the right multiplication kernel technique in linear attention. Notably, the technique does not rely on the partitioning of attention heads, which enables it to be applied to models with different numbers or styles of attention heads, such as multi-head attention, multi-query attention, and grouped query attention. This flexibility exceeds the capabilities of existing SP methods in Megatron-LM or DeepSpeed.
[0031] Moreover, the LASP implementation adopts system engineering optimizations such as kernel fusion and KV state caching, thus significantly improving the execution efficiency. In addition, during the implementation process, great attention is paid to ensuring the compatibility of LASP with various (sharded) distributed data parallel (DDP) training methods, which is referred to as data-sequence hybrid parallelism. Through a large number of experiments on linear transformer models with different numbers of parameters, cluster scales, and sequence lengths, the excellent performance and efficiency of LASP when used in conjunction with these DDP instances are demonstrated. Specifically, LASP is much faster than the existing SP method and can extend the sequence length by 8 times under the same hardware constraints.
[0032] The embodiment of the present invention is a novel sequence parallel strategy for linear attention. It enables linear attention-based models to scale on long sequences without being limited by a single GPU.
[0033] Communication overhead independent of sequence length. The elegant communication mechanism uses the right multiplication kernel technique of linear attention to ensure that the exchange of linear attention intermediate states is independent of the sequence length.
[0034] GPU-friendly implementation. Through meticulous system engineering optimizations, including kernel fusion and KV state caching, the execution efficiency of LASP on GPUs is optimized.
[0035] Compatibility with data parallelism. LASP is compatible with all batch-level DDP methods, such as PyTorch / LegacyDDP, FSDP, and ZeRO series optimizers.
[0036] In one embodiment, Figure 1 is a flowchart of a sequence parallel method for linear attention provided by an embodiment of the present invention. This embodiment is applicable to the situation of processing extremely long sequences. The method can be executed by a sequence parallel device for linear attention, which can be implemented in the form of hardware and / or software, and the sequence parallel device for linear attention can be configured in an electronic device. Exemplarily, the electronic device can include: terminal devices with data processing functions such as computers, iPads, and tablet computers. In this embodiment, LASP slices the sequence on the cluster. Following the slicing idea, LASP divides the input sequence into multiple subsequence blocks and distributes these blocks to different GPUs respectively. For the application of linear attention in an informal environment, in order to fully utilize the advantage of right multiplication in linear attention, the attention calculation of the subsequence can be divided into two different types: internal blocks and cross blocks. Internal blocks involve conventional attention calculations, while cross blocks utilize kernel tricks related to the right multiplication of linear attention. Figure 2 is a visualization block diagram of LASP provided by an embodiment of the present invention. To provide more detailed information, the complex mechanisms of LASP in terms of data distribution, forward propagation, and backward propagation are described. Figure 2 Shows the visualization effect of LASP, further deepening the understanding of LASP. As Figure 2 shown, it can include: a first processing device and a second processing device (which can be called discrete devices). For example, the first processing device can be a CPU; the second processing device can be a GPU (which can be abbreviated as Device), as Figure 2 shown, two adjacent GPUs are respectively Device i and Device i + 1, and each Device contains linear attention, activation function (GLU), and two normalization layers (Norm). Two subsequences Xi and Xi + 1 are respectively distributed to Device i and Device i + 1.
[0037] Taking a typical linear transformer layer as an example, the mechanism of LASP is illustrated. Suppose the original sequence X as the input is divided into multiple subsequence blocks Xi, and then fed into different model replicas on different second processing devices. g represents the conjugate communication operations in both the forward and backward propagations. In the forward propagation, g is the Send and Recv operations sent from device i to (i + 1); in the backward propagation, g is the Send and Recv operations sent from device (i + 1) to i. The communication operations exchange the forward intermediate states KV and the backward intermediate state dKV during the forward and backward propagations to ensure the performance of sequence parallelism. As Figure 1 shown, the method includes:
[0038] S110. Distribute multiple subsequences corresponding to the original sequence on the linear transformer in the distributed environment to the corresponding second processing devices according to a pre-configured data distribution strategy by the first processing device.
[0039] Among them, the data distribution strategy is used to split the original sequence into multiple subsequences and distribute the multiple subsequences to the corresponding second processing devices. In an embodiment, a distributed environment can be understood as a communication group, which can include a first processing device and multiple second processing devices. Moreover, the number of second processing devices included in this communication group is the same as the number of subsequences. It can also be understood that the original sequence is split according to the number of second processing devices included in the communication group to obtain multiple subsequences equal to the number of second processing devices. The original sequence can be understood as a sequence with an extremely long length, and it is prone to memory limitation when processed by a single second processing device.
[0040] S120. Determine the forward total output matrix corresponding to the subsequence by the second processing device using a pre-configured forward propagation method.
[0041] Among them, the forward total output matrix can be understood as a matrix generated by using a pre-configured forward propagation method, and this forward total output matrix is used to evaluate the loss function of the linear attention model in the second processing device. In an embodiment, the forward total output matrix is related to the parameters corresponding to each subsequence. For example, the parameters can include: the total query matrix, the total key matrix, and the total value matrix. The corresponding forward total output matrix can be obtained by using a pre-configured forward propagation method and calculating multiple parameters.
[0042] S130. Determine the parameter gradient corresponding to the subsequence by the second processing device using a pre-configured backward propagation method and the forward total output matrix, and update the parameters of the corresponding subsequence according to the parameter gradient.
[0043] Among them, the parameter gradient refers to the gradients of different parameters corresponding to each subsequence. When the parameters include the total query matrix, the total keyword matrix, and the total value matrix, correspondingly, the parameter gradients can include: the gradient of the total query matrix, the gradient of the total keyword matrix, and the gradient of the total value matrix. In an embodiment, a pre-configured backpropagation method can be adopted, and the gradients of multiple parameters and the forward total output matrix are calculated to obtain the corresponding parameter gradients, and the parameters corresponding to the subsequence are updated according to the parameter gradients.
[0044] The technical solution of this embodiment splits a long original sequence into multiple subsequences through a pre-configured data distribution strategy and distributes them to the corresponding second processing devices to implement the process of distributed processing, and exchanges intermediate states during the forward propagation and backpropagation processes within or between multiple second processing devices, improving the processing efficiency of the sequence, thereby making the implementation on the cluster of second processing devices more hardware-friendly.
[0045] In one embodiment, Figure 3 is a flowchart of another sequence parallel method for linear attention provided by an embodiment of the present invention. This embodiment explains the implementation processes of data distribution, forward propagation, and backpropagation on the basis of the above embodiment. As Figure 3 shown, the method includes:
[0046] S210. Determine the corresponding number of sequence parallel groups according to the pre-configured total number of distributed cards and the sequence parallel scale.
[0047] Among them, the total number of distributed cards can also be referred to as the distributed world size, which is used to represent the number of second processing devices included in a communication group; the sequence parallel scale refers to the number of subsequences into which an original sequence is split, and the value of the sequence parallel scale needs to be divisible by the total number of distributed cards; the number of sequence parallel groups refers to the total number of sequence parallel groups included in a communication group. In an embodiment, the ratio between the total number of distributed cards and the sequence parallel scale can be used as the corresponding number of sequence parallel groups.
[0048] S220. Determine the corresponding subsequence length according to the total length of the original sequence and the sequence parallel scale.
[0049] Among them, the total sequence length is used to represent the total length corresponding to an original sequence; the subsequence length is used to represent the total length corresponding to a subsequence. In an embodiment, the ratio between the total sequence length and the sequence parallel scale can be used as the corresponding subsequence length. Generally, the lengths of each subsequence in a communication group are the same.
[0050] S230. Determine a sequence parallel start device index list according to a pre-acquired global device index list and a sequence parallel scale.
[0051] Among them, the global device index list is used to represent a set of device indexes corresponding to each second processing device included in a communication group; the sequence parallel start device index list is used to represent a set of device indexes corresponding to the first second processing device in each sequence parallel group. The number of device indexes included in the sequence parallel start device index list is equal to the number of sequence parallel groups included in the communication group. In an embodiment, the global device index list can be obtained by using the get_global_rank() function, and the get_global_rank() function is a packaged function, similar to an API interface. In an embodiment, first determine the floor value of the ratio between the global device index list and the sequence parallel scale, then obtain the corresponding sequence parallel start device index according to the product value between the floor value and the sequence parallel scale, and then form all the sequence parallel start device indexes into a corresponding sequence parallel start device index list.
[0052] S240. Split the original sequence into corresponding multiple subsequences according to the subsequence length.
[0053] In an embodiment, the first processing device splits the original sequence into multiple subsequences according to the subsequence length and the total sequence length corresponding to the original sequence, that is, takes the ratio between the total sequence length and the subsequence length as the number of subsequences; alternatively, the first processing device can directly split the original sequence into the same number of subsequences as the sequence parallel scale, that is, each original sequence contains the same number of subsequences as the sequence parallel scale.
[0054] S250. Transmit the subsequences to the corresponding second processing device indexes in the parallel start device index list.
[0055] In an embodiment, the first processing device can split and distribute the corresponding subsequences of multiple original sequences, and transmit all the subsequences corresponding to one original sequence to the corresponding second processing device indexes in the parallel start device index list.
[0056] S260. Scatter and send the subsequences from the parallel start device index list to the second processing devices corresponding to the second processing device indexes in each sequence parallel communication group.
[0057] In an embodiment, the first processing device can scatter and send the subsequences from the parallel start device index list to the second processing devices corresponding to the second processing device indexes in each sequence parallel communication group.
[0058] S270. Determine the total query matrix, total keyword matrix, and total numerical matrix corresponding to the subsequence according to the subsequence, as well as the corresponding query weight coefficient, keyword weight coefficient, and numerical weight coefficient.
[0059] In the embodiment, the product value between each subsequence and its corresponding query weight coefficient can be used as the total query matrix corresponding to the subsequence; the product value between each subsequence and its corresponding keyword weight coefficient can be used as the total keyword matrix corresponding to the subsequence; the product value between each subsequence and its corresponding numerical weight coefficient can be used as the total numerical matrix corresponding to the subsequence. The number of the total query matrix, total keyword matrix, and total numerical matrix calculated by the second processing device is the same as the number of subsequences included in the original sequence. For example, if an original sequence is split into T subsequences, the number of the corresponding total query matrix, total keyword matrix, and total numerical matrix is all T.
[0060] S280. Determine the in-card output matrix corresponding to the subsequence according to the product value between the total query matrix and the transposed matrix of the total keyword matrix, as well as the pre-configured mask matrix and the total numerical matrix.
[0061] In the embodiment, first calculate the product value between the total query matrix and the transposed matrix of the total keyword matrix, then perform an exclusive NOR calculation on the product value and the pre-configured mask matrix to obtain the corresponding exclusive NOR calculation result, and multiply the synchronous calculation result by the total numerical matrix to obtain the in-card output matrix corresponding to the subsequence.
[0062] S290. Determine the inter-card output matrix corresponding to the subsequence according to the total query matrix, the pre-configured set of decay rates, and the forward intermediate state of the previous subsequence.
[0063] Among them, the set of decay rates refers to a decay rate diagonal matrix composed of the decay rates corresponding to each subsequence. It should be noted that the number of decay rates included in the decay rate diagonal matrix is related to the length of the subsequence. For example, if the length of the subsequence is C, the decay rate diagonal matrix contains C decay rates, and the first element on the decay rate diagonal matrix is λ, the second element is λ 2 , and so on, the Cth element is λ C . In the embodiment, first, the second processing device can receive the forward intermediate state from the second processing device of the previous subsequence corresponding to the current subsequence, and save the forward intermediate state of the previous subsequence on the second processing device of the previous subsequence as the forward intermediate state of the current subsequence for reverse calculation; then use the product value between the total query matrix, the pre-configured set of decay rates, and the forward intermediate state of the previous subsequence as the inter-card output matrix of the current subsequence.
[0064] S2100. Determine the forward total output matrix corresponding to the subsequence according to the in-card output matrix and the inter-card output matrix.
[0065] In the embodiment, add the in-card output matrix and the inter-card output matrix corresponding to each subsequence to obtain the forward total output matrix corresponding to the subsequence.
[0066] S2110. Determine the gradient of the in-card query matrix corresponding to the subsequence according to the total output matrix gradient, the transpose matrix of the total value matrix, the pre-configured mask matrix, and the total keyword matrix.
[0067] In the embodiment, first determine the product value between the total output matrix gradient and the transpose matrix corresponding to the total value matrix, then perform an exclusive NOR calculation on the product value and the pre-configured mask matrix, and multiply the result of the exclusive NOR calculation by the total keyword matrix to obtain the gradient of the in-card query matrix corresponding to the subsequence.
[0068] S2120. Determine the gradient of the inter-card query matrix corresponding to the subsequence according to the pre-configured decay rate set, the total output matrix gradient, and the forward intermediate state of the previous subsequence.
[0069] Among them, the decay rate set refers to a decay rate diagonal matrix composed of the decay rates corresponding to each subsequence. In the embodiment, first determine the product value between the pre-configured decay rate set, the total output matrix gradient, and the forward intermediate state of the previous subsequence to obtain the gradient of the inter-card query matrix corresponding to the subsequence.
[0070] S2130. Determine the gradient of the in-card keyword matrix corresponding to the subsequence according to the total output matrix gradient, the transpose matrix of the total value matrix, the pre-configured mask matrix, and the total query matrix.
[0071] In the embodiment, first multiply the total output matrix gradient and the transpose matrix of the total value matrix to obtain the corresponding matrix product value, then perform an exclusive NOR calculation on the matrix product value and the pre-configured mask matrix, and multiply the transpose matrix of the result of the exclusive NOR calculation by the total query matrix to obtain the gradient of the in-card keyword matrix corresponding to the subsequence.
[0072] S2140. Determine the gradient of the in-card value matrix corresponding to the subsequence according to the total query matrix, the transpose matrix of the total keyword matrix, the pre-configured mask matrix, and the total output matrix gradient.
[0073] In the embodiment, first multiply the total query matrix and the transpose matrix of the total keyword matrix to obtain the corresponding matrix product value, then perform an exclusive NOR calculation on the matrix product value and the pre-configured mask matrix, and multiply the transpose matrix of the result of the exclusive NOR calculation by the total output matrix gradient to obtain the gradient of the in-card value matrix corresponding to the subsequence.
[0074] S2150. Determine the gradient of the inter-card keyword matrix for the corresponding subsequence based on the reverse intermediate state of the next received subsequence, the inverse matrix of the pre-configured decay rate set, the decay rate corresponding to the subsequence, and the total value matrix.
[0075] Among them, the decay rate corresponding to the subsequence refers to the last element in the decay rate set, that is, λ. C . In the embodiment, first, the second processing device can receive the reverse intermediate state from the second processing device of the next subsequence corresponding to the current subsequence; then calculate the product value between the inverse matrix of the decay rate set, λ C and the total value matrix; then multiply the product value by the reverse intermediate state of the next subsequence to obtain the gradient of the inter-card keyword matrix for the corresponding subsequence.
[0076] S2160. Determine the gradient of the inter-card value matrix for the corresponding subsequence based on the reverse intermediate state of the next received subsequence, the inverse matrix of the pre-configured decay rate set, the decay rate corresponding to the subsequence, and the total keyword matrix.
[0077] In the embodiment, among them, the decay rate corresponding to the subsequence refers to the last element in the decay rate set, that is, λ. C . In the embodiment, first, the second processing device can receive the reverse intermediate state from the second processing device of the next subsequence corresponding to the current subsequence; then calculate the product value between the inverse matrix of the decay rate set, λ C and the total keyword matrix; then multiply the product value by the reverse intermediate state of the next subsequence to obtain the gradient of the inter-card value matrix for the corresponding subsequence.
[0078] S2170. Determine the gradient of the corresponding total query matrix based on the gradient of the intra-card query matrix and the gradient of the inter-card query matrix.
[0079] In the embodiment, add the gradient of the intra-card query matrix and the gradient of the inter-card query matrix to obtain the gradient of the corresponding total query matrix.
[0080] S2180. Determine the gradient of the corresponding total keyword matrix based on the gradient of the intra-card keyword matrix and the gradient of the inter-card keyword matrix.
[0081] In the embodiment, add the gradient of the intra-card keyword matrix and the gradient of the inter-card keyword matrix to obtain the gradient of the corresponding total keyword matrix.
[0082] S2190. Determine the gradient of the corresponding total value matrix based on the gradient of the intra-card value matrix and the gradient of the inter-card value matrix.
[0083] In an embodiment, the gradients of the in-card numerical matrix and the inter-card numerical matrix are added to obtain the gradient of the corresponding total numerical matrix.
[0084] In one embodiment, to determine the forward total output matrix corresponding to a subsequence using a pre-configured forward propagation method, it further includes: updating the forward intermediate state of the subsequence according to the product value of the forward intermediate state of the previous subsequence and the decay rate corresponding to the subsequence, and the product value of the inverse matrix of the pre-configured decay rate set, the decay rate corresponding to the subsequence, and the transpose matrix of the total keyword matrix and the total numerical matrix. In the embodiment, the decay rate corresponding to the subsequence refers to the last element in the decay rate set, i.e., λ. C First, calculate the product value of the forward intermediate state of the previous subsequence and the decay rate corresponding to the subsequence as the first product value; then calculate the product value of the inverse matrix of the pre-configured decay rate set, the decay rate corresponding to the subsequence, and the total keyword matrix as the second product value; then multiply the transpose matrix of the second product value by the total numerical matrix to obtain the third product value; finally, update the forward intermediate state of the current subsequence with the sum of the first product value and the third product value.
[0085] In one embodiment, for the sequence parallel method of linear attention, it further includes: storing the forward intermediate state corresponding to each subsequence in the preset memory space of the second processing device. In the embodiment, to avoid recalculating the forward intermediate state KV during the backpropagation process, it can be chosen to store it immediately in the high bandwidth memory (HBM) of the GPU after the forward propagation calculation. During the subsequent backpropagation process, LASP directly accesses KV for use. It should be noted that the activation size of KV stored in HBM is d×d and is not affected by the total sequence length N of the original sequence. When the total sequence length N of the input original sequence is extremely large, the memory usage of KV becomes negligible.
[0086] In one embodiment, to determine the parameter gradient corresponding to a subsequence using a pre-configured backpropagation method and the forward total output matrix, it further includes: updating the backward intermediate state of the subsequence according to the backward intermediate state of the next subsequence, the pre-configured decay rate and decay rate set, the total query matrix, and the total output matrix gradient. In the embodiment, the decay rate corresponding to the subsequence refers to the last element in the decay rate set, i.e., λ. CFirst, calculate the product value between the reverse intermediate state of the next subsequence and the decay rate corresponding to the subsequence as the fourth product value; then calculate the product value between the pre-configured decay rate set and the total query matrix as the fifth product value; then multiply the transposed matrix of the fifth product value by the total output matrix gradient to obtain the sixth product value; finally, update the reverse intermediate state of the current subsequence with the sum between the fourth product value and the sixth product value.
[0087] In an embodiment, LASP aims to train long sequences on a linear transformer in a distributed environment by partitioning the input data along its sequence dimension. In this case, each GPU in the distributed environment undertakes the training of the subsequences, thus reducing the large memory footprint associated with activations when training long sequences. Communication operations between GPUs are introduced to transfer intermediate states. The finally trained model absorbs the knowledge obtained from the entire long sequence.
[0088] In one embodiment, Figure 4 is a flowchart of the implementation of data distribution provided by an embodiment of the present invention. It should be noted that the process of data distribution is executed by the first processing device, that is, the subsequences are distributed to the corresponding second processing devices by the first processing device.
[0089] For an input sequence of length N, its embedding space representation is established, denoted as the original sequence where the feature dimension is d. In the LASP framework, the original sequence X is evenly divided into T blocks, where T is called the sequence parallel size and must be divisible by the distributed data size W. Then these divided data blocks are assigned to the corresponding GPUs. It should be noted that different sequence parallel groups receive different data batches. However, within the same group, all data blocks come from the same data batch. In LASP, a detailed description of the data distribution process is shown in Algorithm 1. In addition, Figure 5 is an example diagram of data distribution in LASP provided by an embodiment of the present invention, considering a node with 8 GPUs and dividing 2 sequences into 4 subsequence blocks.
[0090] In this example, the distributed world size is W = 8, the sequence parallel size is T = 4, the number of sequence parallel groups is G = 2, and the list of sequence parallel starting device indices is R src= [0, 4] (i.e., the first second processing device index contained in each of the two sequence parallel groups). For the first batch Seq0, the input sequence X is divided into T chunks X_1, X_2, ... X_T along the sequence dimension, and then transmitted to the first rank in the first SP group, denoted as SP-Group0, corresponding to the global rank 0. The data chunks at the global rank 0 are then scattered to the global ranks 0, 1, 2, 3 within SP-Group0, where only one chunk is retained for each rank. The subsequent batch Seq1 follows a similar distribution process and is assigned to the global ranks 4, 5, 6, 7 within SP-Group1.
[0091] As Figure 4 shown, the data distribution process includes the following steps:
[0092] S410. Input the original sequence X expressed in the embedding space, the total sequence length N, the hidden dimension D, the total number of distributed cards W, and the sequence parallel scale T into the CPU.
[0093] S420. Calculate the number of sequence parallel communication groups: G = W / T.
[0094] S430. Calculate the subsequence length (or chunk length): C = N / T.
[0095] S440. Obtain the global device index list R = get_global_rank().
[0096] S450. Calculate the sequence parallel start device index list R_src = [R / T]*T.
[0097] S460. Along the sequence dimension, split the input X into T subsequences: {X_1, X_2, ..., X_T}.
[0098] S470. Transmit the data subsequences {X_1, X_2, ..., X_T} to the corresponding second processing device index in R_src.
[0099] In one embodiment, Figure 6 is the implementation flowchart of a forward propagation method provided by an embodiment of the present invention. It should be noted that the process of the forward propagation method is executed by the second processing device. As Figure 6 shown, the implementation process of the forward propagation method includes the following steps:
[0100] S610. Input the original sequence X expressed in the embedding space, the total sequence length N, the hidden dimension D, the total number of distributed cards W, the sequence parallel scale T = W, and the decay rate λ into the GPU.
[0101] S620. Distribute the input original sequence X according to the data distribution strategy.
[0102] S630. Calculate the subsequence length C = N / T.
[0103] S640. Initialize the mask matrix M.
[0104] In the embodiment, initialize the mask matrix When i ≥ j, M ij = λ i-j , otherwise; M ij = 0.
[0105] S650. Initialize the decay rate λ.
[0106] In the embodiment, initialize the decay rate set where C represents the subsequence length; Λ represents the decay rate set; λ represents the decay rate corresponding to each subsequence.
[0107] S660. Initialize the activation state KV = 0.
[0108] In the embodiment, during the process of implementing the forward propagation mode, the activation state KV can be understood as the forward intermediate state; initialize where d represents the hidden dimension of.
[0109] S670. Determine whether the parallel calculation of the subsequence t = {1,..., T} on the GPU i = {1,..., W} is completed. If so, end the loop; if not, execute S680.
[0110] S680. Calculate Q t , K t and V t .
[0111] where Q t = X t W Q , K t = X t W K , V t = X t W V ; where Q t represents the total query matrix corresponding to the t-th subsequence; K t represents the total keyword matrix corresponding to the t-th subsequence; V t represents the total value matrix corresponding to the t-th subsequence; X t represents the t-th subsequence; W Q , W K and WV respectively represent the query weight coefficient, the keyword weight coefficient, and the numerical value weight coefficient.
[0112] S690, calculate O t,intra .
[0113] Among them, Among them, O t,intra represents the in-card output matrix corresponding to the t-th subsequence; Q t represents the total query matrix corresponding to the t-th subsequence; represents the transposed matrix of the total keyword matrix corresponding to the t-th subsequence; M represents the pre-configured mask matrix; V t represents the total numerical value matrix corresponding to the t-th subsequence.
[0114] S6100. Determine whether the serial calculation of the subsequence t = {1,..., T} on the GPU i = {1,..., W} is completed. If so, end the loop; if not, execute S6110.
[0115] S6110. Receive the activation state KV from the (i - 1)-th GPU t-1 .
[0116] S6120. Save KV on the i-th GPU for reverse calculation t-1 for KV i .
[0117] S6130. Calculate O t,inter .
[0118] In the embodiment, O t,inter = ΛQ t KV t-1 ; among them, O t,inter represents the inter-card output matrix corresponding to the t-th subsequence; KV t-1 represents the forward intermediate state corresponding to the (t - 1)-th subsequence; Q t represents the total query matrix corresponding to the t-th subsequence.
[0119] S6140. Calculate O t = O t,intra + O t,inter .
[0120] Among them, Q t represents the total output matrix of the t-th subsequence; O t,intra represents the in-card output matrix of the t-th subsequence; O t,inter represents the in-card output matrix of the t-th subsequence.
[0121] S6150. Update KV t .
[0122] Among them, update KV t = λ C KV t-1 +(λ c Λ -1 KV t )V t ; Among them, KV t represents the forward intermediate state corresponding to the t-th subsequence; KV t-1 represents the forward intermediate state corresponding to the (t - 1)-th subsequence; λ C represents the last element in the decay rate set; Λ -1 represents the inverse matrix of the decay rate set; V_t represents the total value matrix corresponding to the t-th subsequence.
[0123] S6160, send the activation state KV t to the (i + 1)-th GPU.
[0124] S6170, return O = O t , where t = {1,..., T}.
[0125] Among them, O represents the forward total output matrix; O t represents the total output matrix of the t-th subsequence.
[0126] In one embodiment, Figure 7 is the implementation flowchart of a backpropagation method provided by an embodiment of the present invention. It should be noted that the process of the backpropagation method is executed by a second processing device. As Figure 7 shown, the implementation process of the backpropagation method includes the following steps:
[0127] S710, input the sequence total length N, hidden dimension D, distributed total number of cards W, sequence parallel scale T = W, decay rate lambda, and Q t , K t , V t , O t , dO t into the GPU.
[0128] Among them, where t ∈ {1, 2,..., T}.
[0129] S720, calculate the subsequence length C = N / T.
[0130] S730, initialize the mask matrix M.
[0131] In the embodiment, initialize the mask matrix when i ≥ j, M ij = λ i-j , otherwise; M ij= 0.
[0132] S740. Initialize the decay rate λ.
[0133] In the embodiment, initialize the decay rate set where C represents the subsequence length; Λ represents the decay rate set; λ represents the decay rate corresponding to each subsequence.
[0134] S750. Initialize the activation state dKV = 0.
[0135] In the embodiment, during the process of implementing the backpropagation method, the activation state dKV can be understood as the reverse intermediate state; initialize where d represents the hidden dimension of.
[0136] S760. Determine whether the parallel computation of the subsequence t = {1,..., T} on the GPU i = {1,..., W} is completed. If so, end the loop; if not, execute S770.
[0137] S770. Calculate dQ t,intra , dQ t,inter , dK t,intra and dV t,intra .
[0138] where where dQ t,intra represents the gradient of the in - card query matrix corresponding to the t - th subsequence; dO t represents the gradient of the total output matrix corresponding to the t - th subsequence; Q t represents the total query matrix corresponding to the t - th subsequence; represents the transposed matrix of the total keyword matrix corresponding to the (t - 1) - th subsequence; represents the transposed matrix of the total keyword matrix corresponding to the t - th subsequence; K t represents the total keyword matrix corresponding to the t - th subsequence; M represents the pre - configured mask matrix; represents the transposed matrix of the total value matrix corresponding to the t - th subsequence; dQ t,inter represents the gradient of the inter - card query matrix corresponding to the t - th subsequence; dK t,intra represents the gradient of the in - card keyword matrix corresponding to the t - th subsequence; dV t,intra represents the gradient of the in - card value matrix corresponding to the t - th subsequence; Λ represents the decay rate set.
[0139] S780. Determine whether the serial computation of the subsequence t = {1,..., T} on the GPU i = {1,..., W} is completed. If so, end the loop; if not, execute S790.
[0140] S790, Receive the activation state dKV from the (i + 1)-th GPU t+1 。
[0141] S7100, Calculate dK t,inter and dV t,inter 。
[0142] Where dV t,inter =(λ C Λ -1 K t )dKV t+1 ; represents the transpose matrix of the reverse intermediate state corresponding to the (t + 1)-th subsequence; dKV t+1 represents the reverse intermediate state corresponding to the (t + 1)-th subsequence; Λ -1 represents the inverse matrix of the decay rate set; K t represents the total keyword matrix corresponding to the t-th subsequence; λ C represents the last element in the decay rate set; V t represents the total value matrix corresponding to the t-th subsequence.
[0143] S7110, Load KV on the i-th GPU i as KV t 。
[0144] S7120, Calculate the sum of intra and inter results: dQ t =dQ t,intra +dQ t,inter, dK t =dK t,intra +dK t,inter ,dV t =dV t,intra +dV t,inter 。
[0145] Where, dQ t represents the gradient of the total query matrix corresponding to the t-th subsequence; dK t represents the gradient of the total keyword matrix corresponding to the t-th subsequence; dV t represents the gradient of the total value matrix corresponding to the t-th subsequence.
[0146] S7130, Update dKV t 。
[0147] Where, dKV t =λ C dKV t+1 +(ΛQ t ) TdO t ; dKV t represents the reverse intermediate state corresponding to the t-th subsequence; dKV t+1 represents the reverse intermediate state corresponding to the (t + 1)-th subsequence; dO t represents the gradient of the total output matrix corresponding to the t-th subsequence; λ C represents the last element in the set of decay rates; Q t represents the total query matrix corresponding to the t-th subsequence; Λ represents the set of decay rates.
[0148] S7140, Send the activation state dKV t to the i-th GPU.
[0149] S7150, Return dQ = [dQ t , dK = [dK t , dV = [dV t , where t ∈ {1, 2,..., T}.
[0150] When examining the LASP algorithm, it should be noted that the forward propagation requires communication for KV activation at each linear attention module layer. The communication volume is determined by Bd 2 / h, where B is the batch size and h is the number of heads. In contrast, the sequence parallelization in Megatron-LM uses two all-gather operations after the two layer normalization layers in each Transformer layer and one reduce-scatter operation after the attention and feed-forward neural network (FFN) layers, resulting in a communication volume of 2BNd + 4BNd / T. DeepSpeed uses an all-to-all collective communication operation to process the input Q, K, V and output O of each attention module layer, resulting in a communication volume of 4BNd / T.
[0151] Table 1 shows the comparison results of the communication volumes of the three frameworks. Among them, d / h is the head dimension, usually set to 128. In practical applications, when N / T ≥ 32, LASP can achieve the lowest theoretical communication volume. In addition, the communication volume of LASP is not affected by the changes in the sequence length N or the subsequence length C, which is a great advantage for extremely long sequence parallelization across large GPU clusters.
[0152] Table 1 Comparison table of communication volumes obtained by different implementation frameworks
[0153]
[0154]
[0155] As shown in Table 1, the simplified formula in the last column represents the calculation result after eliminating Bd. Among them, Megatron-SP refers to Megatron-LM Sequence Parallelism.
[0156] Fusion: To improve the efficiency of LASP on GPUs, fusion is performed on both intra-block and cross-block computations, and the updates of KV and dKV are fused into intra-block and cross-block computations.
[0157] KV State Caching: To avoid recomputing the activated KV during backpropagation, it is chosen to store it immediately in the HBM of the GPU after forward propagation computation. During the subsequent backpropagation process, LASP directly accesses KV for use. Note that the size of the KV activation stored in the HBM is d×d and is not affected by the sequence length N. When the input sequence length N is extremely large, the memory usage of KV becomes negligible.
[0158] Data parallelism techniques are commonly used to split input data along the batch dimension in large-scale distributed deep learning. However, LASP adopts a different approach, dividing data along the sequence dimension, which makes it easier to integrate with data parallelism techniques. As described in the data distribution section and illustrated in Figure 2 , LASP allows specifying a smaller sequence parallel size that is divisible by the distributed world size. This configuration results in the input data being split along both the batch and sequence dimensions, which is a type of hybrid parallelism called data-sequence hybrid parallelism.
[0159] As an important distributed training technique, the sharded data parallel method aims to reduce GPU memory usage during large model training. The ZeRO family of optimizers in DeepSpeed and FSDP in PyTorch propose methods to distribute the model state (including optimizer state, gradients, and model parameters) across all GPUs in a distributed environment. This strategic distribution significantly reduces the memory utilization on a single GPU. As variants of data parallelism techniques, these techniques fit perfectly with LASP. However, this greater focus on minimizing the memory footprint of the model state complements LASP's goal of reducing the activation memory on each GPU. By combining these methods, training large models with long sequence lengths becomes more feasible.
[0160] The LASP proposed in the embodiments of the present invention effectively solves the limitations of existing SP methods on linear transformers by making full use of the specific characteristics of linear attention, thereby significantly improving the parallel efficiency and usability of linear attention models. By implementing an efficient point-to-point communication mechanism and engineering optimizations such as kernel fusion and KV state caching, LASP achieves a significant reduction in communication traffic and improves the hardware utilization of GPU clusters. Compatibility with various types of batch-level DDP methods ensures the practicality of LASP in large-scale distributed training. Moreover, the experimental results in Table 1 highlight the advantages of LASP in terms of scalability, speed, memory usage, and convergence performance on linear transformers, compared with existing SP methods in an out-of-the-box framework.
[0161] In one embodiment, Figure 8 is a schematic structural diagram of a sequence parallel device for linear attention provided by an embodiment of the present invention. As Figure 8 shown, the device includes: a distribution module 810, a first determination module 820, and a second determination module 830.
[0162] Among them, the distribution module 810 is configured to distribute multiple subsequences corresponding to the original sequence on the linear transformer in the distributed environment to the corresponding second processing devices according to a pre-configured data distribution strategy through the first processing device;
[0163] The first determination module 820 is configured to determine the forward total output matrix corresponding to the subsequence by the second processing device using a pre-configured forward propagation method;
[0164] The second determination module 830 is configured to determine the parameter gradient corresponding to the subsequence by the second processing device using a pre-configured backward propagation method and the forward total output matrix, so as to update the parameters of the corresponding subsequence according to the parameter gradient.
[0165] In one embodiment, the distribution module 810 includes:
[0166] The first determination unit is configured to determine the corresponding number of sequence parallel groups according to the pre-configured total number of distributed cards and the sequence parallel scale;
[0167] The second determination unit is configured to determine the corresponding subsequence length according to the total sequence length of the original sequence and the sequence parallel scale;
[0168] The third determination unit is configured to determine the sequence parallel start device index list according to the pre-obtained global device index list and the sequence parallel scale;
[0169] The splitting unit is configured to split the original sequence into corresponding multiple subsequences according to the subsequence length;
[0170] A transmission unit, configured to transmit a subsequence to a corresponding second processing device index in a parallel starting device index list;
[0171] A distribution unit, configured to scatter and send the subsequence from the parallel starting device index list to a second processing device corresponding to the second processing device index in each sequence parallel communication group.
[0172] In one embodiment, the first determination module includes:
[0173] A first determination unit, configured to determine a total query matrix, a total keyword matrix, and a total numerical matrix corresponding to the subsequence according to the subsequence, and corresponding query weight coefficients, keyword weight coefficients, and numerical weight coefficients;
[0174] A second determination unit, configured to determine an in-card output matrix corresponding to the subsequence according to a product value between the total query matrix and a transposed matrix of the total keyword matrix, and a pre-configured mask matrix and the total numerical matrix;
[0175] A third determination unit, configured to determine an inter-card output matrix corresponding to the subsequence according to the total query matrix, a pre-configured decay rate set, and a forward intermediate state of the previous subsequence;
[0176] A fourth determination unit, configured to determine a forward total output matrix corresponding to the subsequence according to the in-card output matrix and the inter-card output matrix.
[0177] In one embodiment, the first determination module further includes:
[0178] An update unit, configured to update the forward intermediate state of the subsequence according to a product value between the forward intermediate state of the previous subsequence and the decay rate corresponding to the subsequence, and a product value between an inverse matrix of the pre-configured decay rate set, the decay rate corresponding to the subsequence, and a transposed matrix of the total keyword matrix and the total numerical matrix.
[0179] In one embodiment, for the sequence parallel device with linear attention, it further includes:
[0180] A storage module, configured to store the forward intermediate state corresponding to each subsequence in a preset memory space of the second processing device.
[0181] In one embodiment, the second determination module includes:
[0182] A first determination unit, configured to determine a gradient of an in-card query matrix corresponding to the subsequence according to a total output matrix gradient, a transposed matrix of the total numerical matrix, a pre-configured mask matrix, and the total keyword matrix;
[0183] A second determination unit, configured to determine the gradient of the inter-card query matrix for the corresponding subsequence according to a pre-configured set of decay rates, the total output matrix gradient, and the forward intermediate state of the previous subsequence;
[0184] A third determination unit, configured to determine the gradient of the intra-card keyword matrix for the corresponding subsequence according to the total output matrix gradient, the transposed matrix of the total value matrix, a pre-configured mask matrix, and the total query matrix;
[0185] A fourth determination unit, configured to determine the gradient of the intra-card value matrix for the corresponding subsequence according to the total query matrix, the transposed matrix of the total keyword matrix, a pre-configured mask matrix, and the total output matrix gradient;
[0186] A fifth determination unit, configured to determine the gradient of the inter-card keyword matrix for the corresponding subsequence according to the received reverse intermediate state of the next subsequence, the inverse matrix of the pre-configured set of decay rates, the decay rate corresponding to the subsequence, and the total value matrix;
[0187] A sixth determination unit, configured to determine the gradient of the inter-card value matrix for the corresponding subsequence according to the received reverse intermediate state of the next subsequence, the inverse matrix of the pre-configured set of decay rates, the decay rate corresponding to the subsequence, and the total keyword matrix;
[0188] A seventh determination unit, configured to determine the gradient of the corresponding total query matrix according to the gradient of the intra-card query matrix and the gradient of the inter-card query matrix;
[0189] An eighth determination unit, configured to determine the gradient of the corresponding total keyword matrix according to the gradient of the intra-card keyword matrix and the gradient of the inter-card keyword matrix;
[0190] A ninth determination unit, configured to determine the gradient of the corresponding total value matrix according to the gradient of the intra-card value matrix and the gradient of the inter-card value matrix.
[0191] In one embodiment, it further includes:
[0192] An update unit, configured to update the reverse intermediate state of the subsequence according to the reverse intermediate state of the next subsequence, the pre-configured decay rate and set of decay rates, the total query matrix, and the total output matrix gradient.
[0193] The sequence parallel device for linear attention provided by the embodiments of the present invention can execute the sequence parallel method for linear attention provided by any embodiment of the present invention, and has the corresponding functional modules and beneficial effects for executing the method.
[0194] In one embodiment, Figure 9 is a structural block diagram of an electronic device provided by an embodiment of the present invention, as Figure 9As shown, a schematic structural diagram of an electronic device 10 that can be used to implement an embodiment of the present invention is shown. The electronic device is intended to represent various forms of digital computers, such as laptop computers, desktop computers, workstations, personal digital assistants, servers, blade servers, mainframe computers, and other suitable computers. The electronic device can also represent various forms of mobile devices, such as personal digital processors, cellular phones, smart phones, wearable devices (such as helmets, glasses, watches, etc.) and other similar computing devices. The components shown herein, their connections and relationships, and their functions are merely examples and are not intended to limit the implementation of the present invention described and / or claimed herein.
[0195] As Figure 9 shown, the electronic device 10 includes at least one processor 11, and a memory communicatively connected to the at least one processor 11, such as a read-only memory (ROM) 12, a random access memory (RAM) 13, etc. The memory stores a computer program executable by the at least one processor. The processor 11 can perform various appropriate actions and processes according to the computer program stored in the read-only memory (ROM) 12 or the computer program loaded from the storage unit 18 into the random access memory (RAM) 13. In the RAM 13, various programs and data required for the operation of the electronic device 10 can also be stored. The processor 11, the ROM 12, and the RAM 13 are connected to each other through a bus 14. The input / output (I / O) interface 15 is also connected to the bus 14.
[0196] Multiple components in the electronic device 10 are connected to the I / O interface 15, including: an input unit 16, such as a keyboard, a mouse, etc.; an output unit 17, such as various types of displays, speakers, etc.; a storage unit 18, such as a magnetic disk, an optical disk, etc.; and a communication unit 19, such as a network card, a modem, a wireless communication transceiver, etc. The communication unit 19 allows the electronic device 10 to exchange information / data with other devices through a computer network such as the Internet and / or various telecommunication networks.
[0197] The processor 11 can be various general-purpose and / or special-purpose processing components with processing and computing capabilities. Some examples of the processor 11 include, but are not limited to, a central processing unit (CPU), a graphics processing unit (GPU), various dedicated artificial intelligence (AI) computing chips, various processors running machine learning model algorithms, a digital signal processor (DSP), and any suitable processor, controller, microcontroller, etc. The processor 11 executes the various methods and processes described above, such as the sequential parallel method for linear attention.
[0198] In some embodiments, the sequential parallel method for linear attention can be implemented as a computer program tangibly embodied in a computer-readable storage medium, such as storage unit 18. In some embodiments, part or all of the computer program can be loaded and / or installed onto the electronic device 10 via the ROM 12 and / or the communication unit 19. When the computer program is loaded into the RAM 13 and executed by the processor 11, one or more steps of the sequential parallel method for linear attention described above can be performed. Alternatively, in other embodiments, the processor 11 can be configured to perform the sequential parallel method for linear attention by any other suitable means (e.g., by means of firmware).
[0199] The various embodiments of the systems and techniques described above in this document can be implemented in digital electronic circuitry, integrated circuit systems, field programmable gate arrays (FPGA), application specific integrated circuits (ASIC), application specific standard products (ASSP), systems on a chip (SOC), complex programmable logic devices (CPLD), computer hardware, firmware, software, and / or combinations thereof. These various embodiments can include: being implemented in one or more computer programs that can be executed and / or interpreted on a programmable system including at least one programmable processor, which can be a special-purpose or general-purpose programmable processor that receives data and instructions from a storage system, at least one input device, and at least one output device, and transmits the data and instructions to the storage system, the at least one input device, and the at least one output device.
[0200] The computer programs for implementing the methods of the present invention can be written in any combination of one or more programming languages. These computer programs can be provided to a processor of a general-purpose computer, a special-purpose computer, or other programmable data processing device, such that when the computer programs are executed by the processor, the functions / operations specified in the flowcharts and / or block diagrams are implemented. The computer programs can be executed entirely on the machine, partially on the machine, as a stand-alone software package partially on the machine and partially on a remote machine, or entirely on a remote machine or server.
[0201] In the context of the present invention, a computer-readable storage medium can be a tangible medium that can contain or store a computer program for use by or in connection with an instruction execution system, apparatus, or device. The computer-readable storage medium can include, but is not limited to, electronic, magnetic, optical, electromagnetic, infrared, or semiconductor systems, apparatus, or devices, or any suitable combination of the foregoing. Alternatively, the computer-readable storage medium can be a machine-readable signal medium. More specific examples of the machine-readable storage medium would include an electrical connection based on one or more wires, a portable computer disk, a hard disk, a random access memory (RAM), a read-only memory (ROM), an erasable programmable read-only memory (EPROM or Flash memory), an optical fiber, a portable compact disc read-only memory (CD-ROM), an optical storage device, a magnetic storage device, or any suitable combination of the foregoing.
[0202] In order to provide interaction with a user, the systems and techniques described herein can be implemented on an electronic device having: a display device for displaying information to the user (e.g., a CRT (cathode ray tube) or LCD (liquid crystal display) monitor); and a keyboard and a pointing device (e.g., a mouse or a trackball) by which the user can provide input to the electronic device. Other kinds of devices can also be used to provide interaction with the user; for example, the feedback provided to the user can be any form of sensory feedback (e.g., visual feedback, auditory feedback, or tactile feedback); and input from the user can be received in any form (including acoustic input, voice input, or tactile input).
[0203] The systems and techniques described herein can be implemented in a computing system that includes backend components (e.g., as a data server), or a computing system that includes middleware components (e.g., an application server), or a computing system that includes frontend components (e.g., a user computer having a graphical user interface or a web browser through which the user can interact with an implementation of the systems and techniques described herein), or a computing system that includes any combination of such backend components, middleware components, or frontend components. The components of the system can be interconnected to each other by any form or medium of digital data communication (e.g., a communication network). Examples of the communication network include: a local area network (LAN), a wide area network (WAN), a blockchain network, and the Internet.
[0204] A computing system may include a client and a server. The client and the server are generally far from each other and usually interact via a communication network. The relationship between the client and the server is created by computer programs that run on the respective computers and have a client-server relationship with each other. The server can be a cloud server, also known as a cloud computing server or a cloud host, which is a host product in the cloud computing service system, solving the defects of difficult management and weak business scalability existing in traditional physical hosts and VPS services.
[0205] It should be understood that various forms of the processes shown above can be used, steps can be reordered, added, or deleted. For example, the steps described in the present invention can be executed in parallel, sequentially, or in a different order, as long as the desired results of the technical solution of the present invention can be achieved, and no limitations are imposed herein.
[0206] The above specific embodiments do not constitute a limitation on the protection scope of the present invention. Those skilled in the art should understand that various modifications, combinations, sub-combinations, and substitutions can be made according to design requirements and other factors. Any modifications, equivalent substitutions, and improvements made within the spirit and principles of the present invention shall be included within the protection scope of the present invention.
Claims
1. A sequential parallel approach to linear attention, characterized in that include: Distributing multiple subsequences corresponding to the original sequence on the linear transformer in the distributed environment to the corresponding second processing device according to a preconfigured data distribution strategy by the first processing device; Determine, by the second processing device, a forward total output matrix corresponding to the subsequence in a preconfigured forward propagation manner; Determine, by the second processing device, a parameter gradient corresponding to the subsequence using a preconfigured back propagation method and the forward total output matrix, so as to update the parameters of the corresponding subsequence according to the parameter gradient; The method of using a preconfigured forward propagation method to determine the total forward output matrix corresponding to the subsequence includes: Determine a total query matrix, a total keyword matrix and a total value matrix corresponding to the subsequence according to the subsequence and the corresponding query weight coefficient, keyword weight coefficient and value weight coefficient; Determine the on-card output matrix of the corresponding subsequence according to the product value between the total query matrix and the transposed matrix of the total keyword matrix, the pre-configured mask matrix and the total number matrix; Determine the inter-card output matrix of the corresponding subsequence according to the total query matrix, the pre-configured decay rate diagonal matrix and the forward intermediate state of the previous subsequence; Update the forward intermediate state of the subsequence according to the product value of the forward intermediate state of the previous subsequence and the decay rate corresponding to the subsequence, and the product value of the inverse matrix of the pre-configured decay rate set, the transposed matrix of the product value between the decay rate corresponding to the subsequence and the total keyword matrix, and the total value matrix; The forward total output matrix corresponding to the subsequence is determined according to the intra-card output matrix and the inter-card output matrix.
2. The method according to claim 1, characterized in that The method of distributing the subsequence corresponding to the original sequence on the linear transformer in the distributed environment to the corresponding second processing device according to the pre-configured data distribution strategy includes: Determine the corresponding number of sequence parallel groups according to the pre-configured total number of distributed cards and sequence parallel scale; Determine the corresponding subsequence length according to the total sequence length of the original sequence and the sequence parallel scale; Determine a sequence parallel starting device index list according to a pre-acquired global device index list and the sequence parallel scale; Splitting the original sequence into corresponding multiple subsequences according to the subsequence lengths; transmitting the subsequence to a corresponding second processing device index in the parallel starting device index list; The subsequences are dispersedly sent from the parallel starting device index list to the second processing devices corresponding to the second processing device index in each sequence parallel communication group.
3. The method according to claim 1, characterized in that Also includes: The forward intermediate state corresponding to each of the subsequences is stored in a preset memory space of the second processing device.
4. The method according to claim 1, characterized in that: The parameter gradient includes: the gradient of the total query matrix, the gradient of the total keyword matrix and the gradient of the total value matrix; the method of using the pre-configured back propagation method and the forward total output matrix to determine the parameter gradient corresponding to the subsequence includes: Determine the gradient of the on-card query matrix corresponding to the subsequence according to the total output matrix gradient, the transposed matrix of the total value matrix, the pre-configured mask matrix and the total keyword matrix; Determine the gradient of the inter-card query matrix of the corresponding subsequence according to the pre-configured decay rate set, the total output matrix gradient and the forward intermediate state of the previous subsequence; Determine the gradient of the card keyword matrix corresponding to the subsequence according to the total output matrix gradient, the transposed matrix of the total value matrix, the pre-configured mask matrix and the total query matrix; Determine the gradient of the card-based value matrix of the corresponding subsequence according to the total query matrix, the transposed matrix of the total keyword matrix, the pre-configured mask matrix and the total output matrix gradient; Determine the gradient of the inter-card keyword matrix corresponding to the subsequence according to the received reverse intermediate state of the next subsequence, the inverse matrix of the pre-configured decay rate set, the decay rate corresponding to the subsequence, and the total value matrix; Determine the gradient of the inter-card value matrix of the corresponding subsequence according to the received reverse intermediate state of the next subsequence, the inverse matrix of the pre-configured decay rate set, the decay rate corresponding to the subsequence, and the total keyword matrix; Determine the gradient of the corresponding total query matrix according to the gradient of the intra-card query matrix and the gradient of the inter-card query matrix; Determining the gradient of the corresponding total keyword matrix according to the gradient of the intra-card keyword matrix and the gradient of the inter-card keyword matrix; The gradient of the corresponding total numerical matrix is determined according to the gradient of the intra-card numerical matrix and the gradient of the inter-card numerical matrix.
5. The method according to claim 4, characterized in that The method of using a pre-configured back propagation method and the forward total output matrix to determine the parameter gradient corresponding to the subsequence also includes: The reverse intermediate state of the subsequence is updated according to the reverse intermediate state of the next subsequence, the preconfigured decay rate and decay rate set, the total query matrix and the total output matrix gradient.
6. A sequential parallel device for linear attention, characterized in that: include: A distribution module, configured to distribute a plurality of subsequences corresponding to the original sequence on the linear transformer in the distributed environment to the corresponding second processing device according to a preconfigured data distribution strategy through the first processing device; A first determining module, configured to determine a forward total output matrix corresponding to the subsequence by using the second processing device in a preconfigured forward propagation manner; a second determination module, configured to determine, by the second processing device, a parameter gradient corresponding to the subsequence using a preconfigured back propagation method and the forward total output matrix, so as to update a parameter of the corresponding subsequence according to the parameter gradient; The first determining module includes: A first determining unit is used to determine a total query matrix, a total keyword matrix and a total value matrix corresponding to the subsequence according to the subsequence and the corresponding query weight coefficient, keyword weight coefficient and value weight coefficient; A second determining unit, configured to determine an on-card output matrix of a corresponding subsequence according to a product value between the total query matrix and a transposed matrix of the total keyword matrix, a preconfigured mask matrix, and the total value matrix; A third determination unit, configured to determine an inter-card output matrix of a corresponding subsequence according to the total query matrix, a pre-configured decay rate diagonal matrix, and a forward intermediate state of a previous subsequence; An updating unit, configured to update the forward intermediate state of the subsequence according to the product value of the forward intermediate state of the previous subsequence and the decay rate corresponding to the subsequence, and the product value of the inverse matrix of the pre-configured decay rate set, the transposed matrix of the product value between the decay rate corresponding to the subsequence and the total keyword matrix, and the total value matrix; The fourth determining unit is used to determine the total forward output matrix corresponding to the subsequence according to the intra-card output matrix and the inter-card output matrix.
7. An electronic device, characterized in that: The electronic device comprises: at least one processor; and a memory communicatively connected to the at least one processor; wherein, The memory stores a computer program executable by the at least one processor, and the computer program is executed by the at least one processor to enable the at least one processor to execute the sequential parallel method for linear attention described in any one of claims 1-5.
8. A computer-readable storage medium, characterized in that: The computer-readable storage medium stores computer instructions, which are used to enable a processor to implement the serial parallel method for linear attention described in any one of claims 1-5 when executed.
Citation Information
Patent Citations
Speech synthesis method and system based on linear self-attention
CN113707127A
Multi-dimensional parallel processing method, system and device based on artificial intelligence, and readable storage medium
CN114035936A