A Method for Predicting the Remaining Useful Life of Rolling Bearings Based on Personalized Federated Learning
By adopting a personalized federated learning method in the field of rolling bearings, combining dynamic weighted federated aggregation strategy and a parallel cross-attention feature extraction module of multi-scale convolution and gated loop units, the problems of data distribution differences and data island effect are solved, and higher prediction accuracy and generalization are achieved.
Patent Information
- Application Number
- CN202510285952.0
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2025-03-12
- Publication Date
- 2025-06-24
- Estimated Expiration
- 2045-03-12
AI Technical Summary
The rolling bearing residual life prediction method based on classic deep learning is difficult to effectively capture data characteristics in different scenarios. It is affected by data distribution differences and data island effects, resulting in limited accuracy and generalization of the prediction model.
Using a method based on personalized federated learning, a personalized prediction model is built by adopting a dynamic weighted federated aggregation strategy on the server side and designing a parallel cross-attention feature extraction module for multi-scale convolution and gated loop units on the client side, combining adversarial training and parameter adaptive fusion strategies to build a personalized prediction model to improve the generalization and prediction accuracy of the model.
Through personalized federated learning methods, the accuracy and generalization of the remaining life prediction of rolling bearings are improved, and data information in different scenarios can be used more effectively to adapt to data distribution under different operating conditions.
Smart Images

Figure CN119783565B_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the technical field of rotating machinery fault diagnosis, and particularly relates to a rolling bearing remaining life prediction method based on personalized federated learning. Background Art
[0002] As a core supporting component of rotating machinery, the health state of rolling bearings directly affects the operation reliability and operation efficiency of the whole machine. Due to harsh operating environments and complex and changeable operating conditions, etc., rolling bearings will inevitably have problems such as wear and performance degradation. Therefore, carrying out research on the health state monitoring and remaining life prediction of rolling bearings is of great significance for ensuring the safe and stable operation of mechanical equipment and the smooth progress of industrial production. Accurate prediction of the remaining life of rolling bearings is like the "wise eyes" installed on mechanical equipment, which can help equipment maintenance personnel timely master the performance degradation state of rolling bearings so as to take preventive measures, thereby reducing the risk of unexpected shutdowns and improving production efficiency.
[0003] With the continuous development of artificial intelligence and the improvement of hardware computing power, deep learning has been widely applied in fields such as image processing, speech recognition, and natural language processing. Various deep learning methods have also emerged in the field of rolling bearing remaining life prediction. However, the remaining life prediction methods based on classical deep learning often need to rely on a large amount of identically distributed data for centralized training, while in the actual industrial environment, the vibration data collected from rolling bearings often has significant distribution differences. In addition to the problem of data distribution differences, the data island effect is also a major challenge in the health management of industrial equipment. Due to reasons such as data security and privacy protection, enterprises are often reluctant or unable to centrally share data, resulting in the dispersion and independence of data. This heterogeneity of data distribution and data island problems make it difficult for the remaining life prediction model based on centralized data processing to effectively capture data characteristics in different scenarios, thus limiting the accuracy and generalization of the prediction model. Therefore, how to effectively utilize data information in different scenarios while protecting data privacy, realize the remaining life prediction of rolling bearings under the federated learning framework, and improve the accuracy and generalization of the prediction model is still an urgent problem to be solved. Summary of the Invention
[0004] The technical problem to be solved by the present invention is to provide a rolling bearing remaining life prediction method based on personalized federated learning, improve the federated learning strategy, use the data of different clients to construct personalized prediction models for the remaining life of rolling bearings under different operating conditions, strengthen the generalization of the prediction model, and improve the accuracy of the remaining life prediction.
[0005] To solve the above technical problems, the technical solution adopted by the present invention is:
[0006] A method for predicting the remaining life of rolling bearings based on personalized federated learning, comprising the following steps:
[0007] S1. Collect the full-life cycle vibration data of the rolling bearing under different operating conditions;
[0008] S2. Store the full-life cycle vibration data of the rolling bearing under different operating conditions in different clients. At each client, convert the collected full-life cycle vibration data into two-dimensional images by using continuous wavelet transform to construct a training data set, a validation data set, and a test data set;
[0009] S3. Use the training data set and the validation data set to initialize the training of the client local model, and upload the parameters of the trained client local model to the server side;
[0010] S4. At the server side, adopt a dynamic weighted federated strategy to aggregate the uploaded model parameters of multiple clients to obtain a global model, and send the aggregated global model parameters to each client;
[0011] S5. At the client side, adopt an adversarial training and parameter adaptive fusion strategy to fuse the aggregated global model parameters with the client local model parameters, and use the training data set to fine-tune the local model parameters;
[0012] S6. Use the trained client local model to analyze the local test data set, so as to obtain the remaining life prediction result of the test bearing.
[0013] A further improvement of the technical solution of the present invention lies in: in step S2, the conversion of the full-life cycle vibration data into two-dimensional images by using continuous wavelet transform includes:
[0014] Analyze each data sample in the full-life cycle vibration data by using continuous wavelet transform to obtain wavelet coefficients, and convert the normalized wavelet coefficients into two-dimensional images.
[0015] A further improvement of the technical solution of the present invention lies in: the client local model in step S3 includes:
[0016] A multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module, which consists of two parallel branches and a cross-attention module. The first branch in the parallel branches consists of a multi-scale convolution module and a channel attention module connected in series, and the second branch in the parallel branches consists of three layers of gated recurrent units connected in series in turn. The outputs of the two parallel branches are subjected to feature fusion through the cross-attention module for spatio-temporal feature extraction of the input data samples;
[0017] The prediction output module, which consists of two fully connected layers, is used to output the prediction results of the remaining life of the rolling bearing.
[0018] A further improvement of the technical solution of the present invention lies in that: in step S4, a dynamic weighted federated strategy is adopted on the server side to aggregate the uploaded model parameters of multiple clients to obtain a global model, including: the uploaded model parameters include the parameters of the first branch in the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in each client's local model and the validation loss of each client's local model wherein k (k = 1, 2, …, K ) represents the k k-th client, K and K represents the total number of clients.
[0019] The dynamic weighted federated strategy includes the following steps:
[0020] First, the server side uses the validation losses of all K clients to calculate the weight of each client wherein, ωmin is a preset minimum weight value, ωmax is the maximum value among the validation losses of all clients;
[0021] Second, normalize to obtain the aggregation weight ωk of the k-th client;
[0022] Finally, the server side uses the parameters of the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module of each client and the aggregation weight ωk to generate the global model parameters θglobal, and distribute the aggregated global model parameters
[0023] A further improvement of the technical solution of the present invention lies in that: in step S5, an adversarial training and parameter adaptive fusion strategy is adopted on the client side to fuse the aggregated global model parameters with the client's local model parameters, and use the training data set to fine-tune the local model parameters, including:
[0024] The adversarial training is to use the global model sent by the server as the first branch in the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client local model, and form a global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module with the second branch in the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client local model and the cross-attention module, which is used to learn the useful features in the client local training dataset; then, the data features learned by the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module and the data features learned by the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client local model are input into a binary classifier composed of two fully connected layers to determine which feature extraction module the input comes from; at the same time, the data features learned by the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module are input into the prediction output module in the client local model for prediction analysis. The optimization objective of the adversarial training is Among them, N k represents the number of samples in the k th client training dataset, x i is the i th training sample data, D is the binary classifier, y i is the true label of the i th training sample, is the predicted output of the prediction output module for the i th training sample, λ is the weight coefficient. During the adversarial training process, first fix the parameters of the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module and the parameters of the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client local model, and update the parameters of the binary classifier: then, fix the parameters of the binary classifier and update the parameters of the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module and the parameters of the prediction output module. The parameter adaptive fusion strategy is to perform weighted fusion on the updated parameters of the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module obtained by the adversarial training and the parameters of the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client local to obtain a new parameter of the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module; The calculation formula is as follows Among them, is the adaptive fusion coefficient, is the k validation loss of the local model of the k-th client.
[0025] The fine-tuning of the local model parameters using the training dataset is to form a remaining useful life prediction model by parallelly cross-attention feature extraction modules and prediction modules of new multi-scale convolution and gated recurrent units, and optimize the parameters of the remaining useful life prediction model using the local training dataset and the validation dataset. The optimization objective is where N k is the number of samples in the training dataset of the k-th client, y i is the i true label of the k-th training sample, is the predicted output of the prediction output model for the i k-th training sample.
[0026] Due to the above technical solutions, the technical progress achieved by the present invention is:
[0027] The method for predicting the remaining useful life of rolling bearings based on personalized federated learning proposed by the present invention adopts a dynamic weighted federated aggregation strategy on the server side, adaptively adjusts the weights according to the validation loss of each client, aggregates multiple client models into a global model, and improves the perception ability of the global model for multi-condition data; a multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module is designed on the client side to capture spatio-temporal feature information in the data; an adversarial training and adaptive fusion strategy is proposed to effectively fuse the global model aggregated on the server side with the client local model, realize the construction of the client personalized prediction model under different operating conditions, and further improve the adaptability of the client model to the local data distribution. The present invention improves the accuracy and generalization of the remaining useful life measurement of rolling bearings under different conditions. Description of the Drawings
[0028] Figure 1 is the flow chart of the method for predicting the remaining useful life of rolling bearings based on personalized federated learning proposed by the present invention;
[0029] Figure 2 is the structure diagram of the client local model proposed by the present invention;
[0030] Figure 3 is the remaining useful life prediction result of the rolling bearing obtained by the method proposed by the present invention; Figure 3 (a) is the prediction result of Bearing1_5, Figure 3(b) is the prediction result of Bearing2_5, Figure 3 (c) is the prediction result of Bearing3_4;
[0031] Figure 4 They are the remaining useful life prediction results of rolling bearings obtained by three clients under a non-federated learning framework; Figure 4 (a) is the prediction result of Bearing1_5, Figure 4 (b) is the prediction result of Bearing2_5, Figure 4 (c) is the prediction result of Bearing3_4. Detailed implementation manner
[0032] To more clearly illustrate the technical solutions in the embodiments of the present invention or the prior art, the following will briefly introduce the drawings required for description in the embodiments or the prior art. Obviously, the drawings in the following description are some embodiments of the present invention. For those of ordinary skill in the art, without creative efforts, other drawings can be obtained based on these drawings.
[0033] As Figure 1 shown, a method for predicting the remaining useful life of a rolling bearing based on personalized federated learning includes the following steps:
[0034] S1. Collect the full-life cycle vibration data of the rolling bearing under different operating conditions;
[0035] S2. Store the full-life cycle vibration data of the rolling bearing under different operating conditions in different clients. At each client, convert the collected full-life cycle vibration data into two-dimensional images using continuous wavelet transform to construct a training data set, a validation data set, and a test data set;
[0036] S3. Initialize and train the client local model using the training data set and the validation data set, and upload the parameters of the trained client local model to the server side;
[0037] S4. Aggregate the uploaded model parameters of multiple clients at the server side using a dynamic weighted federated strategy to obtain a global model, and send the aggregated global model parameters to each client;
[0038] S5. At the client side, adopt an adversarial training and parameter adaptive fusion strategy to fuse the aggregated global model parameters with the client local model parameters, and fine-tune the local model parameters using the training data set;
[0039] S6. Analyze the test data set using the trained client local model to obtain the remaining useful life prediction result of the test bearing.
[0040] In one embodiment, in step S1, first, a vibration acceleration sensor is used to collect the full-life cycle vibration data of a rolling bearing under different operating conditions. The sampling frequency is set to 25.6 kHz, the sampling time interval is 1 minute, and the sampling duration of each data sample is 1.28 seconds. The device under test is a rolling bearing fatigue life acceleration test bench, and its operating conditions are as follows: Condition 1 - rotational speed of 2100 r / min and load of 12 kN, Condition 2 - rotational speed of 2250 r / min and load of 11 kN, and Condition 3 - rotational speed of 2400 r / min and load of 10 kN. Under each condition, the full-life cycle vibration data of 5 rolling bearings are collected. The names of the rolling bearings tested under the three conditions are bearing1_1 - bearing1_5, bearing2_1 - bearing2_5, and bearing3_1 - bearing3_5 respectively.
[0041] In step S2, the full-life cycle vibration data of the rolling bearings collected under Condition 1, Condition 2, and Condition 3 are sequentially stored in 3 clients respectively. On each client, continuous wavelet transform is used to convert each data sample in the full-life cycle vibration data of each tested bearing into a two-dimensional image with a size of 128×128: Among them, WT represents the wavelet coefficient obtained by continuous wavelet transform, and respectively represent the minimum value and the maximum value in the wavelet coefficient WT ; COEF represents the two-dimensional image, represents the pixel point in the two-dimensional image.
[0042] On the first client, bearing1_1, bearing1_2, and bearing1_4 in Condition 1 are used as the training data set, bearing1_3 is used as the validation set, and bearing1_5 is used as the test set;
[0043] On the second client, bearing2_1, bearing2_2, and bearing2_4 in Condition 2 are used as the training data set, bearing2_3 is used as the validation set, and bearing2_5 is used as the test set;
[0044] On the third client, bearing3_1, bearing3_2, and bearing3_5 in Condition 3 are used as the training data set, bearing3_3 is used as the validation set, and bearing3_4 is used as the test set;
[0045] The normalized root mean square value is used to generate the label of the remaining life of the rolling bearing. The calculation formula of the root mean square value is Among them, x ( j ) is the x th sampling point in the sample data j , N is the length of the sample data, X rms is the root mean square value. Then, all the root mean square values calculated from the vibration data of the full life cycle of each rolling bearing are normalized so that each root mean square value is in the range of 0 to 1. The normalization formula is as follows Among them, and respectively represent the maximum and minimum values of the root mean square value, is the normalized root mean square value. Then, the exponential fitting method is further used to fit the normalized root mean square value to more accurately describe the change trend of the performance degradation of the rolling bearing. The fitting formula is as follows Among them, a and b are fitting parameters, which are learned from the data through nonlinear optimization; y t is the normalized remaining life corresponding to the t time calculated using the fitted root mean square value. The y t is discretized according to the number of samples of the full life cycle vibration data collected for each test bearing to obtain the discrete value y i ( i = 1, 2, …, M , M represents the number of sample data), which is used as the true label of each sample data.
[0046] The local model of each client is as Figure 2 shown, and it includes a multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module and a prediction output module;
[0047] The multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module consists of two parallel branches and a cross-attention module. The first branch in the parallel branches is composed of a multi-scale convolution module and a channel attention module in series, and the second branch in the parallel branches is composed of three layers of gated recurrent units in series. The outputs of the two parallel branches are subjected to feature fusion through the cross-attention module to extract spatio-temporal features of the input data samples;
[0048] The prediction output module is composed of two fully connected layers and is used to output the prediction result of the remaining life of the rolling bearing.
[0049] In step S3, first, the local model is trained using the training dataset of each client, and the performance of the local model is verified using the validation dataset; the purpose of training is to adapt each client model to the specific data distribution of the client. The loss function for client model training is as follows: Where N k is the number of samples in the training dataset of the k-th client, y i is the true label of the i -th training sample, is the predicted output of the prediction output model for the i -th training sample.
[0050] The batch size for training is 64, the learning rate is 0.001, the number of local training rounds is 30, and the learning rate decay parameter is set to 0.5. After training, the parameters in the first branch of the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in each client's local model and the validation loss of each client's local model are uploaded to the server side, where k (k = 1, 2, …, K) represents the k -th client, and K = 3 represents the total number of clients.
[0051] In step S4, first, the server side calculates the weight of each client using the validation losses of all 3 clients, where is the preset minimum weight value, and in the present invention = 0.2, is the maximum value among the validation losses of all clients;
[0052] Secondly, is normalized to obtain the aggregated weight of the k-th client;
[0053] Finally, the server side uses the parameters of the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module of each client and the aggregated weight to generate the global model parameters , and the aggregated global model parameters are sent to each client.
[0054] In step S5, each client uses the global model sent by the server as the first branch in the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client's local model, and combines it with the second branch in the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client's local model and the cross-attention module to form a global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module for learning useful features in the client's local training dataset;
[0055] Then, the data features learned by the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module and the data features learned by the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client's local model are input into a binary classifier composed of two fully connected layers to determine which feature extraction module the input comes from; at the same time, the data features learned by the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module are input into the prediction output module in the client's local model for prediction analysis. The optimization objective of the adversarial training is where, N k represents the number of samples in the k th client training dataset, x i is the i th training sample data, D is the binary classifier, y i is the true label of the i th training sample, is the predicted output of the prediction output module for the i th training sample, λ is the weight coefficient.
[0056] During the adversarial training process, first fix the parameters of the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module and the parameters of the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client's local model, and update the parameters of the binary classifier; then, fix the parameters of the binary classifier and update the parameters of the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module and the parameters of the prediction output module. In the adversarial training, the batch size is 64, the learning rate is 0.001, the number of rounds of adversarial training is 5, and the weight coefficient λ is set to 0.3.
[0057] After the adversarial training is completed, the updated parameters of the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module obtained from the adversarial training Parameters of the parallel cross-attention feature extraction module for client-local multi-scale convolution and gated recurrent unit Perform weighted fusion to obtain a new set of parameters for the parallel cross-attention feature extraction module of multi-scale convolution and gated recurrent unit ; The calculation formula is as follows Wherein, is the adaptive fusion coefficient, is the validation loss of the k th client-local model
[0058] Finally, on each client, a remaining useful life prediction model is formed by combining the new parallel cross-attention feature extraction module of multi-scale convolution and gated recurrent unit and the prediction module, and the parameters of the remaining useful life prediction model are optimized and fine-tuned using the local training dataset and validation dataset. The fine-tuning optimization objective is , where N k is the number of samples in the training dataset of the kth client, y i is the true label of the i th training sample, is the predicted output of the prediction output model for the i th training sample. After the fine-tuning on the local client is completed, the client-local model is obtained as the initialization parameter for the next round of federated learning
[0059] In step S6, the trained client-local model is used to analyze the test dataset, thereby obtaining the remaining useful life prediction results of the test bearings
[0060] The root mean square error is used to measure the deviation between the predicted value and the true value. The remaining useful life prediction results of the three test datasets obtained by the method proposed in the present invention are shown as Figure 3 . At the same time, for comparison, the prediction results of three clients for three test datasets under the non-federated learning framework are shown as Figure 4 . It can be seen that the remaining useful life prediction errors obtained by the method proposed in the present invention in the three test datasets are all smaller than those of the comparative method, which indicates that the personalized federated learning strategy proposed in the present invention can effectively improve the generalization performance of the client-local model and the accuracy of the remaining useful life prediction
[0061] By adopting the method for predicting the remaining useful life of rolling bearings based on personalized federated learning, the present invention can, while realizing the protection of privacy data, improve the prediction performance of the client-local model, successfully extract the effective feature information contained in the client vibration data, make the research target more targeted, and is of great significance for accurately predicting the remaining useful life of rolling bearings
[0062] The embodiments described above are only descriptions of the preferred embodiments of the present invention, and do not limit the scope of the present invention. Without departing from the design spirit of the present invention, various deformations and improvements made by those of ordinary skill in the art to the technical solutions of the present invention shall fall within the protection scope determined by the claims of the present invention.
Claims
1. A method for predicting the remaining life of a rolling bearing based on personalized federated learning, characterized in that: The steps include: S1. Collect vibration data of rolling bearings in their entire life cycle under different operating conditions; S2. The full life cycle vibration data of rolling bearings under different operating conditions are stored in different clients. The collected full life cycle vibration data are converted into two-dimensional images by continuous wavelet transform on each client to construct training data sets, verification data sets and test data sets. S3, using the training data set and the verification data set to initialize the client local model training, and upload the trained client local model parameters to the server; S4. A dynamic weighted federation strategy is used on the server to aggregate the model parameters uploaded by multiple clients to obtain a global model, and the aggregated global model parameters are sent to each client. S5. Adopt adversarial training and parameter adaptive fusion strategy on the client side to fuse the aggregated global model parameters with the client local model parameters, and use the training data set to fine-tune the local model parameters; The parameter adaptive fusion strategy is to parallel cross-attention feature extraction module parameters of the updated global multi-scale convolution and gated recurrent unit obtained from adversarial training Cross-attention feature extraction module parameters in parallel with client-side local multi-scale convolution and gated recurrent units Perform weighted fusion to obtain a new multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module parameters The calculation formula is as follows in, is the adaptive fusion coefficient, is the validation loss of the k-th client local model; Among them, the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module is composed of two parallel branches and a cross-attention module, the first branch of the parallel branches is composed of a multi-scale convolution module and a channel attention module connected in series, and the second branch of the parallel branches is composed of three layers of gated recurrent units connected in series in sequence. The outputs of the two parallel branches are subjected to feature fusion through the cross-attention module for extracting spatiotemporal features of the input data samples; S6. Use the trained client local model to analyze the local test data set to obtain the remaining life prediction result of the test bearing.
2. The method for predicting the remaining life of a rolling bearing based on personalized federated learning according to claim 1 is characterized in that: In step S2, the life cycle vibration data is converted into a two-dimensional image using a continuous wavelet transform, specifically, each data sample in the life cycle vibration data is analyzed using a continuous wavelet transform to obtain a wavelet coefficient, and the normalized wavelet coefficient is converted into a two-dimensional image; Among them, WT represents the wavelet coefficient obtained by continuous wavelet transform, WT min and WT max They represent the minimum and maximum values in the wavelet coefficient WT respectively; COEF represents a two-dimensional image, and COEF(i) represents a pixel point in the two-dimensional image.
3. The method for predicting the remaining life of a rolling bearing based on personalized federated learning according to claim 2 is characterized in that: The normalized RMS value is used to generate the label of the remaining life of the rolling bearing. The RMS value is calculated as follows: Where x(j) is the jth sampling point in the sample data x, N is the length of the sample data, and X rms is the root mean square value; then, all the root mean square values calculated from the full life cycle vibration data of each rolling bearing are normalized so that each root mean square value is in the range of 0 to 1. The normalization formula is as follows in, and represent the maximum and minimum values of the RMS value, respectively. is the normalized RMS value; The exponential fitting method is further used to fit the normalized RMS value to more accurately describe the changing trend of rolling bearing performance degradation. The fitting formula is as follows: Among them, a and b are fitting parameters, which are learned from the data through nonlinear optimization; y t is the normalized remaining life corresponding to time t calculated using the fitted RMS value, and y t Discretize the sample number of the full life cycle vibration data collected for each test bearing to obtain the discrete value y i (i=1,2,…,M, M represents the number of sample data), as the true label of each sample data.
4. The method for predicting the remaining life of a rolling bearing based on personalized federated learning according to claim 1 is characterized in that: The client local model in step S3 includes: The prediction output module consists of two fully connected layers and is used to output the prediction results of the remaining life of the rolling bearing; The loss function for client model training is: Among them, N k is the number of samples in the training data set of the kth client, y i is the true label of the i-th training sample, is the predicted output of the prediction output model for the i-th training sample.
5. The method for predicting the remaining life of a rolling bearing based on personalized federated learning according to claim 1, characterized in that: The uploaded model parameters in step S4 include the parameters θ of the first branch in the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in each client local model k and the validation loss of each client's local model Where k (=1, 2, ..., K) represents the kth client, and K represents the total number of clients; The dynamic weighted federation strategy described in step S4 includes the following steps: First, the server uses the verification loss of all K clients to calculate the weight of each client Among them, ω min is the preset minimum weight value, The maximum value among all client validation losses; Secondly, for ω k Perform normalization to obtain the aggregate weight of the kth client Finally, the server side uses the multi-scale convolution and gated recurrent units of each client to cross-attentionally adjust the parameters θ of the first branch in the feature extraction module in parallel. k and aggregation weight Generate global model parameters And the aggregated global model parameters θ g Distributed to each client.
6. The method for predicting the remaining life of a rolling bearing based on personalized federated learning according to claim 1, characterized in that: The adversarial training described in step S5 is to use the global model sent by the server as the first branch in the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client local model, and the second branch and the cross-attention module in the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client local model to form a global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module for learning useful features in the client local training data set; then, the data features learned by the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module and the data features learned by the multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module in the client local model are input into a binary classifier composed of two fully connected layers to determine which feature extraction module the input comes from; at the same time, the data features learned by the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module are input into the prediction output module in the client local model for prediction analysis.
7. The method for predicting the remaining life of a rolling bearing based on personalized federated learning according to claim 6, characterized in that: The optimization goal of adversarial training is Among them, N k represents the number of samples in the k-th client training data set, x i is the i-th training sample data, D is a binary classifier, y i is the true label of the i-th training sample, is the predicted output of the prediction output module for the i-th training sample, λ is the weight coefficient; In the adversarial training process, we first fix the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module parameters θ gk Parallel cross-attention feature extraction module parameters of multi-scale convolution and gated recurrent unit in client local model Update the parameters θ of the binary classifier D :Then, fix the parameters of the binary classifier and update the parameters θ of the global multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module gk and the predicted output module parameters θ y .
8. The method for predicting remaining life of rolling bearings based on personalized federated learning according to claim 1, characterized in that: In step S5, the local model parameters are fine-tuned using the training data set, which is to form a remaining life prediction model with a new multi-scale convolution and gated recurrent unit parallel cross-attention feature extraction module and prediction module, and the parameters of the remaining life prediction model are optimized using the local training data set and the verification data set. The optimization target is Where N k is the number of samples in the training data set of the kth client, y i is the true label of the i-th training sample, is the predicted output of the prediction output model for the i-th training sample.
9. The method for predicting the remaining life of a rolling bearing based on personalized federated learning according to claim 8, characterized in that: After fine-tuning on the local client is completed, the client local model is obtained as the initialization parameter for the next round of federated learning.
Citation Information
Patent Citations
Image classification method based on federal knowledge distillation and ensemble learning
CN117523291A
AIGC federal learning method for designer style fusion and privacy protection
CN117993480A