Hybrid expert network training method, question and answer method, related equipment and program product
By fusing the second expert layer FC2 and the Unpermute operation layer in the MoE model into one layer, the problem of excessive video memory usage during the MoE model training process is solved, and efficient utilization of video memory and improved computing efficiency is achieved.
Patent Information
- Application Number
- CN202510737431.4
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-06-04
- Publication Date
- 2025-07-04
- Estimated Expiration
- 2045-06-04
AI Technical Summary
The problem of increased memory usage during the training of Hybrid Expert Network (MoE) model, especially during the forward propagation process, requires the storage of a large number of intermediate activation values for gradient calculation of the backpropagation process, resulting in a significant increase in memory usage.
The second expert layer FC2 and the inverse rearrangement Unpermute operation layer in the traditional MoE model are fused into one layer (the first fusion layer). Only the input and probability Probs of the first fusion layer are saved during the forward propagation process, and the input of the first fusion layer is used to calculate the gradient to reduce the memory usage.
It significantly reduces the amount of activated video storage that needs to be saved during training, improves computing efficiency, and reduces storage requirements.
Smart Images

Figure CN120258046A_ABST
Abstract
Description
Technical Field
[0001] This application relates to the technical field of large model training, and more specifically, to a method for training a mixture-of-experts network, a question-and-answer method, related devices, and program products. Background Art
[0002] In the past few years, large language models (LLMs) have experienced significant development, and many LLMs have started to adopt a mixture-of-experts (MoE) network structure. The MoE structure dynamically activates a part of the experts through expert routing to replace the dense FFN layer. This mechanism enables computing resources to be concentrated on the most relevant parts, reducing unnecessary computational overhead. Compared with traditional dense Dense models, MoE models have a more efficient training speed.
[0003] Compared with traditional Dense models, the training process of MoE models is more complex, which is likely to cause some problems. One of the more common problems is the increase in video memory occupancy. That is, during the training process of the MOE model, the activation values calculated by each layer during the forward propagation process need to be saved in the video memory for gradient calculation during the backpropagation process. Although the MoE model reduces the amount of computation through sparse activation, each input token needs to be processed by multiple experts (such as top-k), resulting in a multiple increase in the storage volume of intermediate activation values. For example, if the activation dimension after processing by a single expert is dff, then each token needs to store k × dff activation values in the MoE layer (assuming k experts are selected), while the dense model only needs to store dff activation values. It can be seen that the activation video memory occupancy generated during the training process of the MOE model linearly increases with the batch size and the input sequence length, resulting in an increase in video memory occupancy. Summary of the Invention
[0004] In view of the above problems, this application is proposed to provide a method for training a mixture-of-experts network, a question-and-answer method, related devices, and program products to reduce the video memory occupancy during the training process of the MoE model. The specific solutions are as follows:
[0005] In a first aspect, a method for training a mixture-of-experts network is provided. The mixture-of-experts network includes a permutation Permute operation layer, a first expert layer FC1, an activation network, and a first fusion layer. The first fusion layer is fused with a second expert layer FC2 and an un-permutation Unpermute operation layer. The method includes:
[0006] Obtain the subsequence input by the current device and the subsequence input by other devices to form a complete input sequence;
[0007] Forward propagate the complete input sequence through the Permute operation layer, the FC1, and the activation network in sequence to obtain the output of the activation network;
[0008] Use the output of the activation network as the input to the first fusion layer. During the forward propagation of the first fusion layer, sequentially perform the processing of the FC2 and the Unpermute operation layer, and save the input of the first fusion layer and the probability Probs of the expert corresponding to each token in the complete input sequence, where the Probs of the expert corresponding to each token in the complete input sequence are calculated by the routing module in the mixture-of-experts network;
[0009] During the backward propagation of the first fusion layer, calculate the gradients of the parameters of the first fusion layer and the gradients of the Probs respectively based on the saved input of the first fusion layer and the Probs;
[0010] Update the parameters of the mixture-of-experts network according to the gradients of the parameters of each layer during the backward propagation process.
[0011] In a possible design, in another implementation manner of the first aspect of the embodiments of the present application, the Permute operation layer and the FC1 form a second fusion layer;
[0012] During the forward propagation of the second fusion layer, save the subsequence input by the current device, obtain the subsequences input by other devices, form the complete input sequence with the subsequence input by the current device, sequentially perform the processing of the Permute operation layer and the FC1 on the complete input sequence, and use the processing result as the input to the activation network;
[0013] During the backward propagation of the second fusion layer, obtain the subsequences input by other devices, and calculate the gradients of the parameters of the second fusion layer based on the saved subsequence input by the current device and the subsequences input by other devices.
[0014] In a possible design, in another implementation manner of the first aspect of the embodiments of the present application, the process of calculating the gradients of the parameters of the second fusion layer based on the saved subsequence input by the current device and the subsequences input by other devices includes:
[0015] Calculate the gradients of the subsequence input by the current device;
[0016] Calculate the gradients of the weight parameters based on the subsequences input by other devices and the subsequence input by the current device;
[0017] Among them, the process of calculating the gradient of the subsequence input by the current device is executed in parallel with the process of obtaining the subsequences input by other devices.
[0018] In a possible design, in another implementation of the first aspect of the embodiments of the present application, the process of calculating the gradient of the parameters of the first fusion layer and the gradient of the Probs respectively based on the saved input of the first fusion layer and the Probs includes:
[0019] Obtain the gradient of the output calculated by other devices, fuse it with the gradient of the output calculated by the current device, and obtain the gradient grad_output of the complete output;
[0020] Perform a Permute operation on the Probs to obtain the rearranged Probs, and perform a dot product of the rearranged Probs and the input of the first fusion layer to obtain the scaled input;
[0021] Based on the grad_output, calculate the gradient of the scaled input, the gradient of the weight parameters, and the gradient of the rearranged Probs;
[0022] Calculate the gradient of the input of the first fusion layer according to the gradient of the scaled input;
[0023] Calculate the gradient of the Probs before rearrangement according to the gradient of the rearranged Probs.
[0024] In a possible design, in another implementation of the first aspect of the embodiments of the present application, the process of obtaining the gradient of the output calculated by other devices is executed in parallel with the process of obtaining the scaled input.
[0025] In a possible design, in another implementation of the first aspect of the embodiments of the present application, in the backpropagation process of the second fusion layer, the process of obtaining the subsequences input by other devices includes:
[0026] In the backpropagation process of the second fusion layer, use the all gather communication method to obtain the subsequences input by other devices.
[0027] In a possible design, in another implementation of the first aspect of the embodiments of the present application, the process of obtaining the gradient of the output calculated by other devices includes:
[0028] Use the all gather communication method to obtain the gradient of the output calculated by other devices.
[0029] Second aspect, a question and answer method is provided, including:
[0030] Obtain question information;
[0031] Send the question information into the configured mixture-of-experts large model to obtain the answer information output by the mixture-of-experts large model;
[0032] Wherein, the mixture-of-experts large model is a model trained by using the mixture-of-experts network training method described in any one of the foregoing first aspects.
[0033] In a third aspect, a mixture-of-experts network training device is provided. The mixture-of-experts network includes a rearrangement Permute operation layer, a first expert layer FC1, an activation network, and a first fusion layer. The first fusion layer fuses a second expert layer FC2 and an inverse rearrangement Unpermute operation layer. The device includes:
[0034] A calculation unit, configured to obtain the subsequence input by the current device and the subsequences input by other devices to form a complete input sequence; perform forward propagation on the complete input sequence through the Permute operation layer, the FC1, and the activation network in sequence to obtain the output of the activation network; use the output of the activation network as the input of the first fusion layer. In the forward propagation process of the first fusion layer, perform the processing of the FC2 and the Unpermute operation layer in sequence, and save the input of the first fusion layer and the probability Probs of the expert corresponding to each token in the complete input sequence, where the Probs of the expert corresponding to each token in the complete input sequence is calculated by a routing module in the mixture-of-experts network; in the backward propagation process of the first fusion layer, calculate the gradient of the parameters of the first fusion layer and the gradient of the Probs respectively based on the saved input of the first fusion layer and the Probs;
[0035] A parameter update unit, configured to update the parameters of the mixture-of-experts network according to the gradients of the parameters of each layer in the backward propagation process.
[0036] In a fourth aspect, an electronic device is provided, including: a memory and a processor;
[0037] The memory is used for storing programs;
[0038] The processor is configured to execute the programs to implement each step of the mixture-of-experts network training method described in any one of the foregoing first aspects of the present application, or to implement each step of the question-and-answer method described in the foregoing second aspect.
[0039] Fifth aspect, a readable storage medium is provided, on which a computer program is stored. When the computer program is executed by a processor, each step of the hybrid expert network training method described in any one of the foregoing first aspects of the present application is implemented, or each step of the question and answer method described in the foregoing second aspect is implemented.
[0040] Sixth aspect, a computer program product is provided, including a computer program. When the computer program is executed by a processor, each step of the hybrid expert network training method described in any one of the foregoing first aspects of the present application is implemented, or each step of the question and answer method described in the foregoing second aspect is implemented.
[0041] The traditional hybrid expert network consists of a Permute operation layer, a first expert layer FC1, an activation network, a second expert layer FC2, and an Unpermute operation layer. To accelerate the calculation, in the forward propagation process, the Permute operation layer rearranges the input to arrange the hidden layer states of the tokens to be processed by the same expert together, and then performs the processing of FC1 and the activation network. The Unpermute operation layer performs a weighted sum of the results processed by different experts corresponding to the same token according to the corresponding probabilities Probs, and then performs an inverse rearrangement Unpermute operation to obtain the output of the current device. To calculate the gradients of the FC2 parameters and Probs during the backpropagation process, it is necessary to save the input of FC2 (i.e., the output of the activation network) during the forward propagation process for the gradient calculation of FC2 during the backpropagation process, and it is also necessary to save Probs and the output of FC2 during the forward propagation process for the gradient calculation of Probs during the backpropagation process. Based on the above introduction of the technical solution of the present application, the present application integrates the second expert layer FC2 and the inverse rearrangement Unpermute operation layer in the traditional hybrid expert network into one layer (the first fusion layer). In this way, only the input of the first fusion layer and Probs need to be saved during the forward propagation process. During the backpropagation process, the input of the first fusion layer can be used to calculate both the gradient of the first fusion layer parameters and the gradient of Probs, without the need to save the output of FC2 additionally during the forward propagation process, greatly reducing the video memory occupancy of the activations that need to be saved. Description of the Drawings
[0042] By reading the detailed description of the preferred embodiments below, various other advantages and benefits will become clear to those of ordinary skill in the art. The drawings are only for the purpose of showing the preferred embodiments and are not considered to be a limitation of the present application. Moreover, throughout the drawings, the same reference numerals are used to represent the same components. In the drawings:
[0043] Figure 1Illustrates the overall processing framework schematic diagram of the MLP module in the existing MoE model;
[0044] Figure 2 Illustrates Figure 1 The schematic diagram of the calculation process corresponding to the framework;
[0045] Figure 3 Schematic diagram of an implementation system architecture for the hybrid expert network training method provided by an embodiment of the present application;
[0046] Figure 4 Schematic diagram of the process of a hybrid expert network training method provided by an embodiment of the present application;
[0047] Figure 5 Illustrates a schematic diagram of the calculation process using the hybrid expert network training method of the present application;
[0048] Figure 6 Schematic diagram of the structure of a hybrid expert network training device provided by an embodiment of the present application;
[0049] Figure 7 Schematic diagram of the structure of an electronic device provided by an embodiment of the present application. Detailed implementation manners
[0050] Next, the technical solutions in the embodiments of the present application will be clearly and completely described in conjunction with the accompanying drawings in the embodiments of the present application. Obviously, the described embodiments are only a part of the embodiments of the present application, rather than all the embodiments. Based on the embodiments in the present application, all other embodiments obtained by those of ordinary skill in the art without creative efforts shall fall within the protection scope of the present application.
[0051] The training process of the MoE model is more complex than that of traditional dense models, and problems such as load imbalance, increased communication complexity, and increased video memory occupancy are likely to occur. The present application focuses on solving the problem of increased video memory occupancy during the training process of the MoE model.
[0052] First, analyze the problem of video memory occupancy during the training process of the existing MoE model.
[0053] In some implementations, the training of the MoE model generally adopts an expert parallel scheme, that is, different experts are distributed on different devices (computing nodes), and each device only stores and calculates its own expert Expert. Taking the training strategy of strict application load balance as an example, that is, each expert group calculates all tokens. In the existing scheme, the overall process of the MLP module in the MoE model is as Figure 1 shown.
[0054] In the tensor and sequence parallel modes, each segmented input subsequence ( Figure 1The input subsequence 1 (Input1) and input subsequence 2 (Input 2) in the example pass through the gating Gate network to obtain the corresponding probability distribution logits output. The logits of each subsequence are fused ( Figure 1 fused through the All Gather communication operation in the example) together to obtain the complete probability distribution Total Logits. Further passing through the routing network Routing, the expert id corresponding to each token (corresponding to Figure 1 Indices in the example), the probability of each token corresponding to the expert (corresponding to Figure 1 Probs in the example), and the number of Tokens that each expert needs to process (corresponding to Figure 1 Tokens Per Expert in the example) are obtained. At the same time, it is necessary to AllGather the hidden states of all tokens of all input subsequences Input to obtain the complete input sequence Total Input.
[0055] Generally, in order to accelerate the operation process, the hidden states can be rearranged according to the Indices through the Permute operation layer to achieve arranging the hidden states of the tokens that the same expert needs to process together to obtain the Permuted result. Then, the calculation processes of the first expert layer, activation network, and second expert layer are carried out. The first expert layer and the second expert layer can adopt a fully connected layer network, and examples are Figure 2 FC1 Gmm and FC2 Gmm in the example. Among them, FC represents a fully connected layer, and Gmm represents Grouped Matmul, that is, grouped matrix multiplication. The first and second expert layers can efficiently process multi-expert parallel calculations and improve the calculation efficiency by adopting grouped matrix multiplication in the fully connected layer. The activation network can adopt the Geglu activation function.
[0056] The output of the second expert layer is used as the input of the Unpermute operation layer. In the Unpermute operation layer, the results processed by different experts corresponding to the same token are weighted (the weight is the probability Probs) and summed. For the weighted sum result, the Unpermute operation is further carried out with reference to the position mapping relationship row_id_map, and the processed result is distributed ReduceScatter to the corresponding processing devices. Among them, row_id_map records the position mapping relationship of the hidden states of each token before and after permute, and based on this, the data dimension before permute can be restored through the Unpermute operation.
[0057] To more intuitively show the changes in the corresponding data in the above calculation process, this application Figure 2 shows Figure 1 the calculation process shown below and marks the data shape information. Among them, S represents the length of the complete input sequence, B represents the batch size (Batch_Size), H represents the dimension size of the hidden state vector of the token, H_e represents the output dimension of the experts within the first expert layer, T represents the tensor parallelism scale (TP Size), and TopK represents the number of selected experts.
[0058] Combined with Figure 2 shown below, the subsequence input to the current device, i.e., the local input Local Input, is obtained, and its data shape is [S / T, B, H]. Further, the complete input sequence Total Input is obtained through the All Gather communication operation, and its data shape is [S, B, H]. After being processed by the Permute operation layer, the Permuted result is obtained, and its data shape is . After being processed by the first expert layer FC1 Gmm, the FC1 output result is obtained, and its data shape is . Further, after being processed by the activation network, taking the Geglu activation function used in the activation network as an example, this activation function can reduce the dimension of the hidden state vector by half, and the data shape of the obtained Geglu output result is . The Geglu output result is further sent to the second expert layer FC2 Gmm for processing. The second expert layer selected here can restore the dimension of the hidden state vector to the size of H, and the data shape of the obtained FC2 output result is . The FC2 output result is further sent to the Unpermute operation layer for processing to obtain the Unpermuted result, and its data shape is [S, B, H]. Further, it is distributed and reduced by Reduce Scatter to other devices to obtain the output of the current device, i.e., the local output Local Output, and its data shape is [S / T, B, H].
[0059] Figure 2 illustrates the data shapes of the activations generated by each layer in the forward propagation process of the MOE model. To facilitate gradient calculation during the backpropagation process, the main activation video memory information to be saved is shown in Table 1 below:
[0060] Table 1
[0061]
[0062] It should be noted that the above Table 1 only exemplifies some of the main activations that need to be saved, not all activations. In addition to the main activations shown in Table 1, Figure 2 indices, row_id_map, Probs, etc. in Figure 2 also need to be saved. However, the sizes of these parameters are much smaller than the main activations shown in Table 1, which can also be clearly seen from the shapes of the parameters.
[0063] Taking the sequence length S = 8192, TopK = 8, H = 6144, and H_e = 2048 as an example, the storage amount of the activations that need to be saved shown in Table 1 is 1920MB (assuming storage in BF16 type). It can be seen that the video memory of this part of the activations is very large.
[0064] In order to reduce the occupancy of video memory during the training phase of the MOE model, the present application provides an improved training scheme for the mixture-of-experts network.
[0065] The present application provides a method for training a mixture-of-experts network, which can be applied to the system architecture as shown in Figure 3 The system may include a terminal 100 and a server 200. The server 200 may include one or more servers ( Figure 3 illustrated by taking one server as an example).
[0066] Either the terminal 100 or the server 200 can be used alone to execute the method for training the mixture-of-experts network provided in the embodiments of the present application. In addition, the terminal 100 and the server 200 can also be used in cooperation to execute the method for training the mixture-of-experts network provided in the embodiments of the present application.
[0067] The method for training a mixture-of-experts network provided in this embodiment improves the structure of the mixture-of-experts network by fusing the second expert layer FC2 and the un-permute operation layer in the traditional MOE model into one layer, which is defined as the first fusion layer. In this way, the gradient calculation of Probs can be transferred from the un-permute operation to the input part before the first fusion layer, that is, the input of the first fusion layer can be used to calculate the gradient of Probs, thus eliminating the need to save the output result of FC2 and greatly reducing the occupancy of video memory by activations.
[0068] In a possible implementation, the improved mixture-of-experts network includes a permute operation layer, a first expert layer FC1, an activation network, and a first fusion layer, and the first fusion layer fuses the second expert layer FC2 and the un-permute operation layer. On this basis, combined with Figure 4 the following introduces a method for training a mixture-of-experts network provided in the embodiments of the present application, which specifically may include the following steps:
[0069] Step S100: Obtain the subsequences input by the current device and the subsequences input by other devices, and form a complete input sequence.
[0070] Step S110: Propagate the complete input sequence forward through the Permute operation layer, FC1, and the activation network in sequence to obtain the output of the activation network.
[0071] Specifically, the complete input sequence is rearranged through the Permute operation layer, and the hidden states of the tokens that need to be processed by the same expert are arranged together to facilitate improving the subsequent calculation speed. Then, after being processed by the first expert layer FC1 and the activation network, the output of the activation network is obtained.
[0072] Among them, various types of activation functions can be adopted for the activation network. For example, the Geglu activation function, relu activation function, etc. are adopted. In the subsequent embodiments of this application, the Geglu activation function is taken as an example for illustration.
[0073] Step S120: Use the output of the activation network as the input of the first fusion layer. During the forward propagation process of the first fusion layer, the processing of FC2 and the Unpermute operation layer is performed in sequence, and the input of the first fusion layer and the probability Probs of the expert corresponding to each token in the complete input sequence are saved.
[0074] Among them, the Probs of the expert corresponding to each token in the complete input sequence are calculated by the routing module in the mixture-of-experts network. This process can refer to Figure 1 the relevant introduction of the corresponding embodiment, which will not be elaborated here.
[0075] In this embodiment, the second expert layer FC2 and the Unpermute operation layer are fused as the first fusion layer, and the output of the activation network is used as the input of the first fusion layer. The first fusion layer performs the processing of FC2 and the Unpermute operation on the input in sequence during the forward propagation process to obtain the output of the current device. In addition, for facilitating the calculation of the gradient during the backpropagation process, the first fusion layer can also save the input of the first fusion layer (i.e., the output of the activation network) and the above-mentioned Probs during the forward propagation process.
[0076] Step S130: During the backpropagation process of the first fusion layer, based on the saved input of the first fusion layer and Probs, calculate the gradients of the parameters of the first fusion layer and the gradient of Probs respectively.
[0077] Specifically, since the present application integrates the FC2 and Unpermute operations into one layer, based on this, the gradient of Probs can be calculated based on the input of the first fusion layer, without the need to save the output of FC2 as in the prior art to calculate the gradient of Probs, thus saving the video memory occupancy.
[0078] Step S140: Update the parameters of the mixture-of-experts network according to the gradients of the parameters of each layer in the backpropagation process.
[0079] In the backpropagation process, the gradients of the parameters of each layer are calculated layer by layer from the output forward to update the parameters of the mixture-of-experts network. Since the traditional FC2 and Unpermute operation layers are integrated into one layer in this embodiment, the calculation strategy of the backpropagation process of the first fusion layer is introduced in the above step S130. The gradients of the parameters of other layers can be calculated according to the existing backpropagation process calculation strategy, or other improved algorithms can also be used to calculate the parameter gradients, which are not limited in this embodiment.
[0080] Compared with the training process of the traditional MOE model, the traditional MOE model training process needs to save the input of FC2 (i.e., the output of the activation network) in the forward propagation process for calculating the gradient of FC2 in the backpropagation process, and also needs to save the outputs of Probs and FC2 in the forward propagation process for calculating the gradient of Probs in the backpropagation process. However, the method provided in this embodiment integrates the second expert layer FC2 and the Unpermute operation layer in the mixture-of-experts network into one layer (the first fusion layer). In this way, only the input of the first fusion layer and Probs need to be saved in the forward propagation process. In the backpropagation process, the input of the first fusion layer can be used to calculate both the gradient of the parameters of the first fusion layer and the gradient of Probs, without the need to additionally save the output of FC2 in the forward propagation process, greatly reducing the video memory occupancy of the saved activations.
[0081] In some possible implementations, in order to further reduce the video memory occupancy during the training process, the structure of the mixture-of-experts network can be further improved in this embodiment. Specifically:
[0082] The Permute operation layer and FC1 can be integrated into one layer, defined as the second fusion layer.
[0083] In the forward propagation process of the second fusion layer, save the subsequence input by the current device, obtain the subsequences input by other devices, form a complete input sequence with the subsequence input by the current device, and sequentially perform the processing of the Permute operation layer and FC1 on the complete input sequence, and use the processing result as the input of the activation network.
[0084] During the backpropagation process of the second fusion layer, obtain the subsequences input by other devices. Based on the saved subsequences input by the current device and the obtained subsequences input by other devices, calculate the gradients of the parameters of the second fusion layer.
[0085] Using the method of this embodiment, only the subsequences input by the current device need to be saved during the forward propagation process in the second fusion layer, and its data shape is [S / T, B, H]. When calculating the gradients during backpropagation, first communicate with other devices to obtain the subsequences input by other devices, and form a complete input sequence with the local subsequences, and then the gradients of the parameters of the second fusion layer can be calculated. Compared with the traditional MOE model that uses a structure with a Permute operation layer and FC1 separated, the data shape of the activated data that needs to be saved for the FC1 layer during the training process of the traditional MOE model is , it can be seen by comparison that after fusing the Permute operation layer and FC1 into one layer (the second fusion layer), the memory occupancy of the activated data required to be saved during the forward propagation process is further reduced.
[0086] In a possible implementation, during the backpropagation process of the second fusion layer, the communication process of obtaining the subsequences input by other devices can be executed in parallel with the gradient calculation process, that is, the communication process and the gradient calculation process can be overlapped, thereby improving the running efficiency.
[0087] Specifically, the process of calculating the gradients of the parameters of the second fusion layer based on the saved subsequences input by the current device and the obtained subsequences input by other devices may include:
[0088] First, calculate the gradients of the subsequences input by the current device.
[0089] After obtaining the subsequences input by other devices, based on the subsequences input by other devices and the subsequences input by the current device, obtain a complete input sequence, and calculate the gradients of the weight parameters based on the complete input sequence.
[0090] Among them, the process of calculating the gradients of the subsequences input by the current device can be executed in parallel with the communication process of obtaining the subsequences input by other devices, improving the running efficiency.
[0091] In a possible implementation, during the backpropagation process of the second fusion layer, the communication process of obtaining the subsequences input by other devices can use the All Gather communication method to obtain the subsequences input by other devices. In addition, other communication methods can also be used to obtain the subsequences input by other devices. Among them, All-gather is a collective communication operation that aims to collect the data of each process (or node) to all processes. This means that at the end of the All-gather operation, each process will have the data of all other processes.
[0092] In some embodiments of the present application, further for the aforementioned step S130, the process of calculating the gradients of the parameters of the first fusion layer and the gradients of Probs respectively based on the saved input and Probs of the first fusion layer during the backpropagation process of the first fusion layer will be described. The specific implementation process may include the following sub-steps:
[0093] S1. During the backpropagation process of the first fusion layer, obtain the gradients of the outputs calculated by other devices, and fuse them with the gradients of the outputs calculated by the current device to obtain the gradients grad_output of the complete output.
[0094] This process can use the all gather communication method to obtain the gradients of the outputs calculated by other devices. Of course, other communication methods can also be used to obtain the gradients of the outputs calculated by other devices.
[0095] S2. Perform a Permute operation on Probs to obtain the rearranged Probs, and multiply the rearranged Probs with the input of the first fusion layer to obtain the scaled input.
[0096] Specifically, since the input of the first fusion layer is processed by the previous permute operation layer and is in the rearranged permuted state. Therefore, in this step, it is also necessary to perform a Permute operation on Probs to obtain the rearranged Probs.
[0097] Define the input of the first fusion layer as x and the rearranged Probs as Permuted_Probs:
[0098] Permuted_Probs = Permute(Probs, row_id_map).
[0099] Wherein, the definition of row_id_map refers to the previous introduction and will not be elaborated here.
[0100] Furthermore, multiply the rearranged Probs with the input x of the first fusion layer to obtain the scaled input scaled_x:
[0101] .
[0102] S3. Based on the gradients grad_output of the complete output obtained in step S1, calculate the gradients of the scaled input scaled_x, the weight parameter w, and the rearranged Probs.
[0103] Since scaled_x is in the Permuted state, grad_output should also be in the same Permuted state. Therefore, perform a Permute operation on grad_output to obtain the rearranged grad_output:
[0104] Permuted_grad = Permute(grad_output, indices).
[0105] Among them, the definition of indices refers to the previous introduction and will not be elaborated here.
[0106] Furthermore, calculate the gradient of scaled_x. Since there are multiple experts, the grouped matrix multiplication Gmm can be used. The gradient of scaled_x is denoted as scaled_dx, and the calculation formula is:
[0107] scaled_dx = Gmm(Permuted_grad, weight.T).
[0108] Among them, weight represents the parameters of the experts, and weight.T represents the transpose operation on weight.
[0109] Calculate the gradient dw of the weight w:
[0110] dw = Gmm(scaled_x.T, Permuted_grad).
[0111] Among them, scaled_x.T represents the transpose operation on scaled_x.
[0112] Calculate the gradient of Permuted_Probs:
[0113] .
[0114] S4. According to the gradient scaled_dx of the scaled input, calculate the gradient dx of the input x of the first fusion layer. According to the gradient Permuted_dProbs of the rearranged Probs, calculate the gradient dProbs of the Probs before rearrangement.
[0115] Specifically:
[0116] ;
[0117] dProbs = Permute_bwd(Permuted_dProbs, row_id_map).
[0118] Among them, the Permute_bwd operation represents the reverse operation of the Permute operation. Since Permuted_Probs is obtained by performing the Permute operation on Probs, in this step, calculating the gradient by taking the derivative of Probs can be converted to performing the Permute_bwd operation on Permuted_dProbs.
[0119] In a possible implementation, the process of obtaining the gradient of the output calculated by other devices in step S1 above and the process of obtaining the scaled input in step S2 can be executed in parallel. That is, the process of communicating with other devices to obtain the output gradient by the current device and the process of calculating the scaled input scaled_x by the current device can be overlapped, thereby improving the operation efficiency.
[0120] Refer to Figure 5 , which exemplifies the calculation process after adopting the improved training method of the mixture-of-experts network of the present application and annotates the data shape information therein.
[0121] Obtain the subsequence of the input of the current device, that is, the local input Local Input, as the input of the second fusion layer.
[0122] In the second fusion layer (Permute&FC1):
[0123] Forward propagation Forward process: Step 1. Save the local input for gradient calculation in the backward propagation process, that is, Save local input for backward. The data shape of the saved local input subsequence is [S / T, B, H]. Step 2. Use the All Gather communication method to obtain the subsequences of the inputs of other devices and form a complete input sequence with the subsequence of the input of the current device. Steps 3-4. Perform the Permute operation and the FC1 operation on the complete input sequence in sequence.
[0124] Backward propagation Backward process: Step 1. Use the All Gather communication method to obtain the subsequences of the inputs of other devices and form a complete input sequence with the saved subsequence of the input of the current device, that is, All Gather Input. Step 2. Calculate the gradient of the subsequence x of the local input, that is, Compute grad of x. Step 3. Distribute the gradient of the subsequence x of the local input to other devices, that is, Reduce Scatter grad of x. Step 4. Calculate the gradient of the weight parameter w, that is, Compute grad of w.
[0125] The output of the second fusion layer is obtained in the forward propagation process, and its data shape is The output of the second fusion layer serves as the input to the activation network, and various types of activation functions can be selected for the activation network. Figure 5 Taking the Geglu activation function as an example in
[0126] After being processed by the Geglu activation network, the output is obtained, and its data shape is The output of the Geglu activation network serves as the input to the first fusion layer.
[0127] In the first fusion layer (FC2&Unpermute):
[0128] Forward propagation process: Steps 1 - 2. Sequentially go through the FC2 and Unpermute operations to obtain the local output. Step 3. Save the input and Probs of the first fusion layer for gradient calculation in the backpropagation process. Figure 5 The data shape of the input of the first fusion layer saved is further exemplified in Step 4. Distribute the output of the current device to other devices, that is, Reduce Scater Output.
[0129] Backward propagation process: Step 1. Use the All gather communication method to obtain the gradient of the output calculated by other devices, fuse it with the gradient of the output calculated by the current device to obtain the gradient of the complete output grad_output, that is, All gather grad_output. Step 2. Perform the Permute operation on Probs. Step 3. Multiply the rearranged Probs with the input x of the first fusion layer to obtain the scaled input (corresponding to Figure 5 in ). Step 4. Based on grad_output, calculate the gradient of the scaled input scaled_x, the gradient of the weight parameter w, and the gradient of the rearranged Probs, that is, Compute grad of scaled_x, w, Permuted_Probs. Step 5. According to the gradient of the scaled input, calculate the gradient of the input x of the first fusion layer, and according to the gradient of the rearranged Probs, calculate the gradient of the Probs before rearrangement, that is, Compute grad of x, Probs.
[0130] Based on Figure 5 the improved training method of the mixture of experts model shown, the main activation video memory information to be saved is shown in Table 2 below:
[0131] Table 2
[0132]
[0133] It should be noted that the above Table 2 only exemplifies some of the main activations that need to be saved, not all activations. In addition to the main activations shown in Table 2, Figure 5 indices, row_id_map, Probs, etc. in also need to be saved. However, the sizes of these parameters are much smaller than the main activations shown in Table 2, which can be clearly seen from the shapes of the parameters.
[0134] In order to compare with the amount of data of the main activations that need to be saved in the traditional training method of the mixture-of-experts network (that is, to compare Table 1 and Table 2), in this embodiment, taking the sequence length S = 8192, TopK = 8, H = 6144, and H_e = 2048 as an example, the amount of activation storage that needs to be saved shown in Table 2 is 396 MB (assuming storage in BF16 type). Compared with the 1920 MB of the amount of activation storage that needs to be saved shown in Table 1, the amount of activation storage is reduced by nearly 5 times using the improved method of this application. Moreover, when H, T, and TopK are larger, the benefits of using the improved method of this application are higher.
[0135] Some embodiments of this application further provide a question-answering method, which can be implemented based on a mixture-of-experts large model. The mixture-of-experts large model can be trained using the mixture-of-experts network training method of the foregoing embodiments. The question-answering method includes the following steps:
[0136] First, obtain the question information.
[0137] Further, send the question information into the mixture-of-experts large model to obtain the answer information output by the mixture-of-experts large model.
[0138] Next, a description is given of the mixture-of-experts network training device provided in the embodiments of this application. The mixture-of-experts network training device described below can be correspondingly referred to the mixture-of-experts network training method described above.
[0139] See Figure 6 , Figure 6 which is a schematic structural diagram of a mixture-of-experts network training device disclosed in the embodiments of this application. In this embodiment, the mixture-of-experts network includes a reordering Permute operation layer, a first expert layer FC1, an activation network, and a first fusion layer. The first fusion layer fuses a second expert layer FC2 and an un-reordering Unpermute operation layer. As Figure 6 shown, this mixture-of-experts network training device may include:
[0140] The computing unit 11 is configured to obtain the subsequences input by the current device and the subsequences input by other devices to form a complete input sequence; perform forward propagation on the complete input sequence through the Permute operation layer, the FC1, and the activation network in sequence to obtain the output of the activation network; use the output of the activation network as the input of the first fusion layer. During the forward propagation of the first fusion layer, perform the processing of the FC2 and the Unpermute operation layer in sequence, and save the input of the first fusion layer and the probability Probs of the expert corresponding to each token in the complete input sequence, where the Probs of the expert corresponding to each token in the complete input sequence are calculated by the routing module in the mixture-of-experts network; during the backward propagation of the first fusion layer, calculate the gradients of the parameters of the first fusion layer and the gradients of the Probs respectively based on the saved input of the first fusion layer and the Probs.
[0141] The parameter update unit 12 is configured to update the parameters of the mixture-of-experts network according to the gradients of the parameters of each layer during the backward propagation process.
[0142] In a possible implementation, the Permute operation layer and the FC1 form a second fusion layer; during the forward propagation of the second fusion layer, the computing unit saves the subsequence input by the current device, obtains the subsequences input by other devices, combines them with the subsequence input by the current device to form the complete input sequence, performs the processing of the Permute operation layer and the FC1 on the complete input sequence in sequence, and uses the processing result as the input of the activation network; during the backward propagation of the second fusion layer, obtains the subsequences input by other devices, and calculates the gradients of the parameters of the second fusion layer based on the saved subsequence input by the current device and the subsequences input by other devices.
[0143] In a possible implementation, the process of the computing unit calculating the gradients of the parameters of the second fusion layer based on the saved subsequence input by the current device and the subsequences input by other devices includes:
[0144] Calculating the gradient of the subsequence input by the current device;
[0145] Calculating the gradient of the weight parameter based on the subsequences input by other devices and the subsequence input by the current device;
[0146] Among them, the process of calculating the gradient of the subsequence input by the current device is executed in parallel with the process of obtaining the subsequences input by other devices.
[0147] In a possible implementation, the process in which the computing unit calculates the gradients of the parameters of the first fusion layer and the gradient of the Probs respectively based on the saved input of the first fusion layer and the Probs includes:
[0148] Obtain the gradient of the output calculated by other devices, fuse it with the gradient of the output calculated by the current device, and obtain the gradient grad_output of the complete output;
[0149] Perform a Permute operation on the Probs to obtain the rearranged Probs, multiply the rearranged Probs with the input of the first fusion layer, and obtain the scaled input;
[0150] Based on the grad_output, calculate the gradient of the scaled input, the gradient of the weight parameter, and the gradient of the rearranged Probs;
[0151] According to the gradient of the scaled input, calculate the gradient of the input of the first fusion layer;
[0152] According to the gradient of the rearranged Probs, calculate the gradient of the Probs before rearrangement.
[0153] In a possible implementation, the process in which the computing unit obtains the gradient of the output calculated by other devices is executed in parallel with the process of obtaining the scaled input.
[0154] In a possible implementation, the process in which the computing unit obtains the subsequence of the input of other devices during the backpropagation process of the second fusion layer includes:
[0155] During the backpropagation process of the second fusion layer, use the all gather communication method to obtain the subsequence of the input of other devices.
[0156] In a possible implementation, the process in which the computing unit obtains the gradient of the output calculated by other devices includes:
[0157] Use the all gather communication method to obtain the gradient of the output calculated by other devices.
[0158] This application embodiment also provides an electronic device. Refer to Figure 7 As shown, it shows a schematic structural diagram of an electronic device suitable for implementing the electronic device in this application embodiment. The electronic device in this application embodiment may include, but is not limited to, terminals such as mobile phones, tablet computers, computers, and so on. Figure 7 The electronic device shown is only an example and should not impose any limitation on the functions and usage scope of this application embodiment.
[0159] AsFigure 7 As shown, the electronic device may include a processing device (such as a central processing unit, a graphics processing unit, etc.) 601, which may perform various appropriate actions and processes according to the program stored in the read-only memory (ROM) 602 or the program loaded from the storage device 608 into the random access memory (RAM) 603, so as to implement the hybrid expert network training method of the foregoing embodiments of the present application, or implement the question-and-answer method of the foregoing embodiments of the present application. When the electronic device is powered on, various programs and data required for the operation of the electronic device are also stored in the RAM 603. The processing device 601, the ROM 602, and the RAM 603 are connected to each other through a bus 604. The input / output (I / O) interface 605 is also connected to the bus 604.
[0160] Generally, the following devices may be connected to the I / O interface 605: an input device 606 including, for example, a touch screen, a touchpad, a keyboard, a mouse, a camera, a microphone, an accelerometer, a gyroscope, etc.; an output device 607 including, for example, a liquid crystal display (LCD), a speaker, a vibrator, etc.; a storage device 608 including, for example, a memory card, a hard disk, etc.; and a communication device 609. The communication device 609 may allow the electronic device to communicate with other devices wirelessly or wiredly to exchange data. Although Figure 7 an electronic device with various devices is shown, it should be understood that it is not required to implement or have all the shown devices. Instead, more or fewer devices may be implemented or had.
[0161] In an embodiment of the present application, there is also provided a computer program product including computer-readable instructions, which, when running on an electronic device, enable the electronic device to implement any one of the hybrid expert network training methods or question-and-answer methods provided in the embodiments of the present application.
[0162] In an embodiment of the present application, there is also provided a computer-readable storage medium carrying one or more computer programs, which, when executed by an electronic device, can enable the electronic device to implement any one of the hybrid expert network training methods or question-and-answer methods provided in the embodiments of the present application.
[0163] In addition, it should be noted that the device embodiments described above are merely illustrative. The units described as separate components may or may not be physically separated, and the components shown as units may or may not be physical units, that is, they may be located in one place or distributed to multiple network units. Some or all of the modules may be selected according to actual needs to achieve the purpose of the solution of this embodiment. In addition, in the drawings of the device embodiments provided in the present application, the connection relationship between the modules indicates that they have a communication connection, which can be specifically implemented as one or more communication buses or signal lines.
[0164] Through the description of the above embodiments, those skilled in the art can clearly understand that the present application can be implemented by means of software plus necessary general hardware. Of course, it can also be implemented by dedicated hardware including application-specific integrated circuits, dedicated CPUs, dedicated memories, dedicated components, etc. Generally, functions accomplished by computer programs can easily be implemented by corresponding hardware, and the specific hardware structures for implementing the same function can also be diverse, such as analog circuits, digital circuits, or dedicated circuits. However, for the present application, in more cases, implementation by software programs is a better embodiment. Based on such understanding, the technical solution of the present application, in essence, or the part that makes a contribution to the prior art, can be embodied in the form of a software product. This computer software product is stored in a readable storage medium, such as a floppy disk, USB flash drive, mobile hard disk, ROM, RAM, magnetic disk, or optical disc of a computer, and includes several instructions for causing a computer device (which can be a personal computer, training device, or network device, etc.) to execute the methods described in various embodiments of the present application.
[0165] In the above embodiments, it can be implemented in whole or in part by software, hardware, firmware, or any combination thereof. When implemented using software, it can be implemented in whole or in part in the form of a computer program product.
[0166] The computer program product includes one or more computer instructions. When the computer program instructions are loaded and executed on a computer, the processes or functions described in the embodiments of the present application are generated in whole or in part. The computer can be a general computer, a dedicated computer, a computer network, or other programmable devices. The computer instructions can be stored in a computer-readable storage medium, or transmitted from one computer-readable storage medium to another computer-readable storage medium. For example, the computer instructions can be transmitted from a website, computer, training device, or data center to another website, computer, training device, or data center in a wired manner (such as coaxial cable, optical fiber, digital subscriber line (DSL)) or a wireless manner (such as infrared, wireless, microwave, etc.). The computer-readable storage medium can be any available medium that a computer can store, or a data storage device such as a training device or data center that includes one or more integrated available media. The available medium can be a magnetic medium (such as a floppy disk, hard disk, magnetic tape), an optical medium (such as a DVD), or a semiconductor medium (such as a solid state disk (SSD)).
[0167] The various embodiments in this specification are described in a progressive manner. Each embodiment focuses on the differences from other embodiments. The various embodiments can be combined as needed, and the same or similar parts can be referred to each other.
Claims
1. A method for training a mixture of experts network, characterized in that, The mixture of experts network includes a Permute operation layer, a first experts layer FC1, an activation network, and a first fusion layer. The first fusion layer fuses a second experts layer FC2 and an Unpermute operation layer. The method includes: Obtain the subsequence input by the current device and the subsequences input by other devices to form a complete input sequence; Perform forward propagation on the complete input sequence through the Permute operation layer, the FC1, and the activation network in sequence to obtain the output of the activation network; Use the output of the activation network as the input of the first fusion layer. During the forward propagation of the first fusion layer, sequentially perform the processing of the FC2 and the Unpermute operation layer, and save the input of the first fusion layer and the probability Probs of the expert corresponding to each token in the complete input sequence, where the Probs of the expert corresponding to each token in the complete input sequence are calculated by the routing module in the mixture of experts network; During the backward propagation of the first fusion layer, calculate the gradients of the parameters of the first fusion layer and the gradients of the Probs respectively based on the saved input of the first fusion layer and the Probs; Update the parameters of the mixture of experts network according to the gradients of the parameters of each layer during the backward propagation process.
2. The method according to claim 1, wherein The Permute operation layer and the FC1 form a second fusion layer; During the forward propagation of the second fusion layer, save the subsequence input by the current device, obtain the subsequences input by other devices, form the complete input sequence with the subsequence input by the current device, sequentially perform the processing of the Permute operation layer and the FC1 on the complete input sequence, and use the processing result as the input of the activation network; During the backward propagation of the second fusion layer, obtain the subsequences input by other devices, and calculate the gradients of the parameters of the second fusion layer based on the saved subsequence input by the current device and the subsequences input by other devices.
3. The method according to claim 2, wherein The process of calculating the gradients of the parameters of the second fusion layer based on the saved subsequence input by the current device and the subsequences input by other devices includes: Calculate the gradients of the subsequence input by the current device; Calculate the gradients of the weight parameters based on the subsequences input by other devices and the subsequence input by the current device; Among them, the process of calculating the gradients of the subsequence input by the current device is executed in parallel with the process of obtaining the subsequences input by other devices.
4. The method according to claim 1, characterized in that, The process of calculating the gradients of the parameters of the first fusion layer and the gradients of the Probs respectively based on the saved input of the first fusion layer and the Probs includes: Obtain the gradients of the output calculated by other devices, fuse them with the gradients of the output calculated by the current device to obtain the gradients grad_output of the complete output; Perform a Permute operation on the Probs to obtain the permuted Probs, multiply the permuted Probs with the input of the first fusion layer to obtain the scaled input; Based on the grad_output, calculate the gradients of the scaled input, the weight parameters, and the re-arranged Probs. According to the gradient of the scaled input, calculate the gradient of the input of the first fusion layer. According to the gradient of the re-arranged Probs, calculate the gradient of the Probs before re-arrangement.
5. The method according to claim 4, wherein The process of obtaining the gradients of the outputs calculated by other devices is executed in parallel with the process of obtaining the scaled input.
6. The method according to claim 2, characterized in that, During the backpropagation process of the second fusion layer, the process of obtaining the subsequences input by other devices includes: During the backpropagation process of the second fusion layer, use the all-gather communication method to obtain the subsequences input by other devices.
7. The method according to claim 4, characterized in that The process of obtaining the gradients of the outputs calculated by other devices includes: Use the all-gather communication method to obtain the gradients of the outputs calculated by other devices.
8. A question-and-answer method, characterized in that, Includes: Obtain the question information. Send the question information into the configured mixture-of-experts large model to obtain the answer information output by the mixture-of-experts large model. Wherein, the mixture-of-experts large model is a model trained by using the mixture-of-experts network training method described in any one of claims 1 to 7.
9. A hybrid expert network training device, characterized in that, The mixture-of-experts network includes a re-arrangement Permute operation layer, a first expert layer FC1, an activation network, and a first fusion layer. The first fusion layer is fused with a second expert layer FC2 and an inverse re-arrangement Unpermute operation layer. The device includes: A calculation unit, configured to obtain the subsequence input by the current device and the subsequences input by other devices to form a complete input sequence; perform forward propagation on the complete input sequence through the Permute operation layer, the FC1, and the activation network in sequence to obtain the output of the activation network; use the output of the activation network as the input of the first fusion layer. During the forward propagation process of the first fusion layer, sequentially perform the processing of the FC2 and the Unpermute operation layer, and save the input of the first fusion layer and the probabilities Probs of the experts corresponding to each token in the complete input sequence, where the Probs of the experts corresponding to each token in the complete input sequence are calculated by the routing module in the mixture-of-experts network; during the backpropagation process of the first fusion layer, based on the saved input of the first fusion layer and the Probs, calculate the gradients of the parameters of the first fusion layer and the gradients of the Probs respectively. A parameter update unit, configured to update the parameters of the mixture-of-experts network according to the gradients of the parameters of each layer during the backpropagation process.
10. An electronic device, characterized in that, Includes: A memory and a processor; The memory is used to store programs. The processor is configured to execute the program to implement each step of the mixture-of-experts network training method described in any one of claims 1 to 7, or to implement each step of the question-and-answer method described in claim 8.
11. A readable storage medium having a computer program stored thereon, characterized in that, When the computer program is executed by the processor, it implements each step of the mixture-of-experts network training method described in any one of claims 1 to 7, or implements each step of the question-and-answer method described in claim 8.
12. A computer program product, comprising a computer program, characterized in that, When the computer program is executed by a processor, it implements each step of the method for training a mixture of experts network according to any one of claims 1 to 7, or implements each step of the question-answering method according to claim 8.
Citation Information
Patent Citations
Checkpoint selection method and device based on DNN model and storage medium
CN114692829A
Neural network model training method and device, electronic equipment and storage medium
CN115688917A
Back propagation optimization method and device of Attention operator
CN118114737A
Video memory optimization online method and system based on FastMoE model
CN118505491A
Lightweight hybrid expert model architecture system and implementation method thereof
CN119026693A
Cited By
Expert model training method and device, storage medium and electronic equipment
CN120806040A