Federal learning method based on comparative learning and knowledge distillation
By combining contrastive learning with knowledge distillation, the federated learning method solves the problems of weak model generalization ability and low stability caused by data heterogeneity and few-shot learning in federated learning, thus improving the accuracy and stability of the graded early warning of internet addiction behavior.
Patent Information
- Application Number
- CN202511480713.7
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-10-16
- Publication Date
- 2025-11-14
AI Technical Summary
Federated learning suffers from problems such as heterogeneous data distribution, limited data volume leading to weak model generalization ability, overfitting, and low model stability, which are particularly evident in the scenario of preventing addiction in online games.
We employ a federated learning approach based on contrastive learning and knowledge distillation. By combining local client data augmentation, contrastive learning training, and knowledge distillation with forgetting loss to calculate the aggregate weights of model parameters, we dynamically adjust the model update strategy to improve the model's stability and adaptability.
It effectively alleviates the overfitting problem caused by small sample learning, enhances the model's generalization ability and adaptability, and improves the accuracy and stability of the graded early warning of internet addiction behavior.
Smart Images

Figure CN120952111A_ABST
Abstract
Description
Technical Field
[0001] This invention relates to the field of federated learning technology, and in particular to a federated learning method based on contrastive learning and knowledge distillation. Background Technology
[0002] Federated learning is a distributed machine learning technology that trains models on local devices and uploads parameters to a central server for aggregation in an encrypted manner. It enables cross-domain collaborative modeling while protecting data privacy, effectively solving the data silos and privacy problems of traditional centralized learning. It has already demonstrated its application value in fields such as healthcare and finance.
[0003] However, traditional federated learning algorithms still have the following limitations: 1. Data Distribution Heterogeneity: In practical applications, the data from various clients in federated learning exhibits high heterogeneity, meaning the data distribution does not conform to the independent and identically distributed (non-IID) assumption. This heterogeneity poses a significant challenge to the model during global optimization, affecting its generalization ability and convergence speed. For example, the network usage behavior characteristics of different clients vary considerably, which may make it difficult for the model to balance the optimization objectives of each client during global updates, thereby reducing overall performance. 2. Challenges of Few-Shot Learning: In federated learning, the amount of data available to clients is typically limited, making the model prone to overfitting during training, resulting in weak generalization ability and unstable training. Furthermore, few-shot learning makes it difficult to accurately assess the data distribution, further increasing the difficulty of model optimization. 3. The intertwining of multiple challenges: The multiple challenges of federated heterogeneous optimization, few-shot learning, and knowledge forgetting are combined; for example, in order to protect privacy, data cannot be processed centrally, while distributed training needs to achieve efficient optimization under the conditions of data heterogeneity and limited sample size, which puts extremely high demands on algorithm design and system architecture.
[0004] These problems are amplified in the context of online game addiction prevention. Limited sample data from various participants and individual differences in user behavior exacerbate the non-independent and identically distributed nature of the data, making it difficult for the anti-addiction system to capture diverse addiction characteristics and limiting the model's adaptability. Furthermore, as the number of federated learning rounds increases, the model forgets early core risk strategies, resulting in low model stability. Summary of the Invention
[0005] This invention provides a federated learning method based on contrastive learning and knowledge distillation, which aims to improve the stability and adaptability of the model.
[0006] To achieve the above objectives, this invention provides a federated learning method based on contrastive learning and knowledge distillation, comprising: Step 1: Multiple local clients collect local data and data tags, and upload the data tags to the central server. The local data includes network behavior data, physiological signal data, behavioral posture data, and background data. Step 2: The central server filters the datasets that match the data labels in the public dataset based on the data labels, and then distributes the initialized model and the filtered datasets to each local client. Step 3: For each local client, the local client performs data augmentation on the local data and the received dataset to obtain training data. The initial model is then trained by comparison using the local data, the received dataset, and the training data to obtain the student model. The model parameters of the student model are then uploaded to the central server. Step 4: The central server calculates the aggregate weight of each set of model parameters uploaded using the forgetting loss, and aggregates all uploaded model parameters according to the aggregate weight of each set of model parameters to obtain the teacher model. Step 5: Distribute the teacher model to each local client, use local data, the received dataset, and training data to perform knowledge distillation on the teacher model and student model to obtain the local model, and upload the model parameters of the local model to the central server. Repeat the aggregation weight calculation and model parameter aggregation until the teacher model meets the preset training termination condition, and output the global model for the graded early warning of internet addiction behavior.
[0007] Furthermore, the local client performs data augmentation on local data and the received dataset to obtain training data, including: The local data and the image data in the received dataset are flipped to obtain the first augmented data; The local data and other data in the received dataset, excluding image data, are augmented to obtain the second enhanced data. The training data is composed of the first augmented data and the second augmented data.
[0008] Furthermore, the loss function expression for contrastive learning training of the initial model using local data, the received dataset, and the training data is as follows: ; in, This represents the contrastive learning loss value. This represents the local data features extracted through model initialization. This represents the features extracted from the training data through model initialization. This represents the features of the received dataset extracted through model initialization. Represents the similarity function. Indicates temperature parameter, This indicates the total number of data points.
[0009] Furthermore, the loss function expression for knowledge distillation of the teacher and student models using local data, the received dataset, and the training data is as follows: ; in, This represents the knowledge distillation loss value. Representing the teacher model, Representing the student model, This represents the features output by the teacher model. This represents the features output by the student model. Representing feature dimension, , This represents the feature probability distribution after normalization.
[0010] Furthermore, the expression for calculating the forgetting loss is: ; in, Indicates the loss value for forgetting. Represents the global model. Indicates the first A local model This represents the model parameters of the teacher model. Indicates the first Model parameters for a local model. This represents the Euclidean distance function.
[0011] Furthermore, the central server calculates the aggregate weights of each set of model parameters uploaded using the forgetting loss as follows: ; in, Indicates the first The aggregate weights of each set of model parameters uploaded by each local client. This indicates the total number of local clients. This represents an exponential function.
[0012] Furthermore, the expression for aggregating all uploaded model parameters based on the aggregated weights of each group of model parameters is as follows: ; in, This represents the model parameters of the updated teacher model.
[0013] Furthermore, the training termination conditions include: Training terminates when the number of iterations reaches a preset threshold. Training is terminated when the changes in the model parameters of the teacher model are less than a preset threshold in several consecutive iterations. Training is terminated when the difference between the model parameters of the teacher model and the model parameters of each local model is less than a preset threshold in several consecutive iterations. Training is terminated when the teacher model's performance metrics on the validation set reach or exceed a preset threshold.
[0014] The above-described solution of the present invention has the following beneficial effects: This invention collects local data and data labels through multiple local clients and uploads the data labels to a central server. The central server filters datasets matching the data labels from a public dataset and distributes the initial model and the selected datasets to each local client. Each local client performs data augmentation on its local data and the received dataset to obtain training data. Using the local data, the received dataset, and the training data, it performs comparative learning training on the initial model to obtain a student model. The student model's parameters are then uploaded to the central server, and the aggregation weights of each uploaded set of model parameters are calculated using forgetting loss to aggregate all uploaded model parameters, resulting in a teacher model. The teacher model is then distributed to each local client, and knowledge distillation is performed on the teacher and student models using the local data, the received dataset, and the training data to obtain a local model. The model parameters of the local model are uploaded to the central server, and the aggregation weight calculation and aggregation are repeated until the teacher model meets the preset training termination condition, and a global model is output for graded early warning of internet addiction behavior. Compared with the prior art, this invention filters public datasets by obtaining data labels from local clients, and trains the model together with local data after data augmentation on each local client. This effectively alleviates the overfitting problem caused by few-shot learning and enhances the generalization ability of the model. During the training process of the local model, contrastive learning is introduced to mine the intrinsic results of the data, enhancing the model's adaptability to data from different local clients. Knowledge distillation is introduced to constrain the model update process to prevent overfitting. In the model aggregation stage, the aggregation weights of the model parameters are calculated through forgetting loss, and the aggregation strategy is dynamically adjusted to reduce knowledge loss and improve the stability and adaptability of the model.
[0015] Other beneficial effects of the present invention will be described in detail in the following detailed description section. Attached Figure Description
[0016] Figure 1 This is a flowchart illustrating an embodiment of the present invention; Figure 2 This is a schematic diagram of the framework for federated learning in an embodiment of the present invention. Detailed Implementation
[0017] To make the technical problems, solutions, and advantages of this invention clearer, a detailed description will be provided below with reference to the accompanying drawings and specific embodiments. Obviously, the described embodiments are only some, not all, of the embodiments of this invention. All other embodiments obtained by those skilled in the art based on the embodiments of this invention without creative effort are within the scope of protection of this invention.
[0018] In the description of this invention, it should be noted that the terms "center," "upper," "lower," "left," "right," "vertical," "horizontal," "inner," and "outer," etc., indicate the orientation or positional relationship based on the orientation or positional relationship shown in the accompanying drawings. They are used only for the convenience of describing the invention and for simplifying the description, and do not indicate or imply that the device or element referred to must have a specific orientation, or be constructed and operated in a specific orientation. Therefore, they should not be construed as limitations on the invention. Furthermore, the terms "first," "second," and "third" are used for descriptive purposes only and should not be construed as indicating or implying relative importance.
[0019] In the description of this invention, it should be noted that, unless otherwise explicitly specified and limited, the terms "installation," "connection," and "linking" should be interpreted broadly. For example, they can refer to a locking connection, a detachable connection, or an integral connection; they can refer to a mechanical connection or an electrical connection; they can refer to a direct connection or an indirect connection through an intermediate medium; and they can refer to the internal connection of two components. Those skilled in the art can understand the specific meaning of the above terms in this invention based on the specific circumstances.
[0020] Furthermore, the technical features involved in the different embodiments of the present invention described below can be combined with each other as long as they do not conflict with each other.
[0021] This invention addresses existing problems by providing a federated learning method based on contrastive learning and knowledge distillation.
[0022] like Figure 1 , Figure 2 As shown, embodiments of the present invention provide a federated learning method based on contrastive learning and knowledge distillation, including: Step 1: Collect local data and data tags from multiple local clients and upload the data tags to the central server. The local data includes network behavior data, physiological signal data, behavioral posture data, and background data. Step 2: The central server filters the datasets that match the data labels in the public dataset based on the data labels, and then distributes the initialized model and the filtered datasets to each local client. Step 3: For each local client, the local client performs data augmentation on the local data and the received dataset to obtain training data. The initial model is then trained by comparison using the local data, the received dataset, and the training data to obtain the student model. The model parameters of the student model are then uploaded to the central server. Step 4: The central server calculates the aggregate weight of each set of model parameters uploaded using the forgetting loss, and aggregates all uploaded model parameters according to the aggregate weight of each set of model parameters to obtain the teacher model. Step 5: Distribute the teacher model to each local client, use local data, the received dataset, and training data to perform knowledge distillation on the teacher model and student model to obtain the local model, and upload the model parameters of the local model to the central server. Repeat the aggregation weight calculation and model parameter aggregation until the teacher model meets the preset training termination condition, and output the global model for the graded early warning of internet addiction behavior.
[0023] In this embodiment of the invention, the framework of federated learning includes: a central server and multiple local clients; The federated learning process involves using local network devices and mobile smart devices as local clients. In this embodiment, the network devices are used to collect data on minors' online behavior and categorized labels for internet addiction, such as browsing history, frequency of social media use, online gaming time, and the type and duration of videos watched. The categorized labels for internet addiction include levels of internet addiction behavior based on the minors' online behavior data. The mobile smart devices collect physiological signal data, behavioral posture data, and behavioral category labels from minors. Physiological signal data includes indicators such as heart rate, brain waves, and eye tracking, which can reflect the minors' physical reactions and psychological state when using the internet. Behavioral posture data is used to capture the minors' online behavior. The posture, distance, angle, and emotional state of adults when using network devices can reflect the behavioral patterns and habits of minors. Behavioral category labels refer to classifying minors' online behavior into different categories, such as normal behavior, excessive entertainment behavior, and learning behavior. Network behavior data, physiological signal data, and behavioral posture data are used as local data. The hierarchical labels and behavioral category labels of internet addiction behavior are used as local data labels. The initial model issued by the central server is compared and trained using local data, received datasets, and enhanced data to obtain a student model. Knowledge distillation is performed on the student model and teacher model using local data, received datasets, and enhanced data to obtain a local model. The model parameters of the local model are then uploaded to the central server. In this embodiment of the invention, the central server in federated learning can be a computing device such as a desktop computer, laptop, handheld computer, server, server cluster, or cloud server. It is used to distribute the initialized model and the dataset that matches the data label in the public dataset to each local client, calculate the aggregate weight of each set of uploaded model parameters, and aggregate all uploaded model parameters according to the aggregate weight of each set of model parameters to obtain a global model. Finally, the global model is used to classify and warn about internet addiction behavior of target individuals based on their network behavior data, physiological signal data, behavioral posture data, and background data. The resulting warning can be classified as the level of internet addiction behavior, which is divided into severe internet addiction, moderate internet addiction, and no internet addiction. The warning result is determined according to the level. For example, if the level is severe internet addiction, an emergency intervention warning is issued; if the level is moderate internet addiction, an addiction prevention warning is issued; and if the level is no internet addiction, no warning is required.
[0024] The initialization model can employ a convolutional neural network (CNN) for multimodal feature extraction, used to process network behavior data, physiological signal data, and behavioral pose data. This network structure can effectively fuse data features from different modalities, thereby improving the model's generalization ability and adaptability.
[0025] The input layer receives network behavior data, physiological signal data, behavioral posture data, and background data. Subsequently, the feature extraction layer extracts features from the data across different modalities. These features are then fused in the fusion layer using an attention mechanism to generate a comprehensive feature representation. The output layer uses the Softmax activation function to output a tiered warning result for network addiction behavior. Furthermore, the model incorporates contrastive learning loss, knowledge distillation loss, and classification loss to optimize performance, ensuring effective handling of multimodal data and improving the model's generalization ability and adaptability.
[0026] In this embodiment of the invention, federated learning needs to consider the following constraints: 1. Data privacy protection constraints: Throughout the federated learning process, the transmission and sharing of data are limited to model parameters and data labels. Direct sharing of raw data is strictly prohibited. Local data of each local client must be kept on the local client and cannot be directly uploaded to the central server or other local clients to ensure data privacy and security. 2. Data tag consistency constraint: The data tags uploaded by each local client to the central server must accurately reflect the category and characteristics of the local data, and the tag format and semantics must be consistent so that the central server can accurately filter and allocate public datasets. The selection and allocation of public datasets should be based on the matching degree of data tags to ensure their compatibility and relevance with the local data of the local clients. 3. Model update stability constraints: During the model update process on the local client, the introduction of contrastive learning and knowledge distillation must ensure the stability of the model update and avoid overfitting or training instability caused by data heterogeneity or small sample size. The model update magnitude on the local client in each iteration must be constrained to prevent the performance of the global model from deteriorating due to local optimization. 4. Constraints on the rationality of aggregation strategy: During the model aggregation phase, the dynamic weight allocation mechanism based on forgetting loss must ensure the rationality and dynamism of weight allocation. The calculation of weights should accurately reflect the degree of forgetting of global knowledge by the model in the local client, and avoid deviation of the global model due to unreasonable weight allocation. The aggregation process of the global model should take into account both the preservation of global knowledge and the satisfaction of the client's personalized needs, and ensure that the aggregated global model has good generalization ability and adaptability.
[0027] In this embodiment of the invention, given the heterogeneity and limited sample size of the data held by each participant in federated learning, in order to optimize the learning effect, each local client needs to upload the data labels of its own collected local data to the central server. After receiving the data labels from each local client, the central server will filter out the datasets that match the data labels from the public dataset based on the characteristics of the data labels. Subsequently, the central server will send the filtered datasets and the initialized model to each local client.
[0028] Under the federated learning framework, in order to effectively carry out contrastive learning tasks and meet the requirements of contrastive learning for data diversity, it is necessary to preprocess the local data held by each local client and the received dataset. Therefore, in this embodiment of the invention, the local client performs data augmentation on the local data and the received dataset to obtain training data, including: The local data and the image data in the received dataset are flipped to obtain the first augmented data; The local data and other data in the received dataset, excluding image data, are augmented to obtain the second enhanced data. The training data is composed of the first augmented data and the second augmented data.
[0029] This data augmentation process aims to expand the dataset through ensemble transformation, enhance the model's ability to learn features from different perspectives, and thus improve the effectiveness of contrastive learning and the model's generalization ability.
[0030] This invention embodiment flips the local data and the image data in the received dataset to obtain the first enhanced data, including: The image data in both the local data set and the received data set is horizontally flipped, that is, the left and right halves of the image data are symmetrically flipped, as shown in the following formula: ; ; ; in, Represents the coordinates of the original image. This represents the coordinates of the horizontally flipped image. Indicates the width of the image. This represents the image data after horizontal flipping; The image data in both the local data set and the received data set is vertically flipped, that is, the upper and lower halves of the image data are symmetrically flipped, as shown in the following formula: ; ; ; in, This represents the coordinates of the image after vertical flipping. Indicates the height of the image. This represents the image data after vertical flipping.
[0031] For text data, a data augmentation method using synonym replacement is used. Certain words in the text are randomly selected and replaced with their synonyms, maintaining semantic integrity while increasing vocabulary diversity.
[0032] For two-dimensional tabular data, oversampling is used for data augmentation. Oversampling is a data augmentation technique used to address class imbalance, especially when the minority class sample size is small. Oversampling increases the proportion of minority class samples in the dataset by generating new minority class samples, thereby improving the model's ability to identify the minority class. Suppose we have a minority class sample set... Each sample It is A feature vector of dimension. First, a sample is randomly selected from the minority class sample set. Subsequently, in the minority class sample set, the same as The k nearest neighbors Finally from Randomly select a sample And generate a new sample :
[0033] In federated learning, client data distribution is typically highly heterogeneous (non-IID). Contrastive learning, by bringing similar samples closer together and dissimilar samples further apart, can better adapt to the differences in data distribution among different clients by learning the inherent structure of the data, resulting in more robust feature representations. Therefore, this embodiment of the invention employs contrastive learning to train the initialization model, as detailed below: 1. Perform data augmentation on local data and received datasets, including flipping image data, time-series transformation of network behavior data, and noise addition to physiological signal data, to increase data diversity.
[0034] 2. Use the initialization model to extract features from the augmented data to obtain feature representations. Next, construct positive and negative sample pairs, where the positive sample pairs are local data and its augmented data, and the negative sample pairs are local data and data from other clients.
[0035] 3. Calculate the similarity between feature vectors using cosine similarity, and optimize the model using contrastive learning loss function to minimize the feature distance of positive sample pairs and maximize the feature distance of negative sample pairs.
[0036] 4. By updating the model parameters through backpropagation and optimization algorithms, and through multiple iterations, the feature representation capability of the model is gradually optimized, thereby improving the model's generalization ability and stability.
[0037] Specifically, the loss function expression for comparative learning training of the initial model using local data, the received dataset, and the training data is as follows: ; in, This represents the contrastive learning loss value. This represents the local data features extracted through model initialization. This represents the features extracted from the training data through model initialization. This represents the features of the received dataset extracted through model initialization. The similarity function is represented by cosine similarity, which ranges from -1 to 1. The more similar the two sets of features are, the closer the value is to 1; otherwise, the value is closer to -1. This represents a temperature parameter used to control the degree of similarity concentration; smaller values indicate lower similarity. This will make the similarity distribution sharper and larger. This will make the similarity distribution smoother. This indicates the total number of data points.
[0038] By minimizing this loss function, the model is trained to bring similar samples closer together and push dissimilar samples further apart in the feature space, enabling the model to learn the inherent structure of the data.
[0039] In the federated learning framework, knowledge distillation is introduced during the model update phase of each local client. Its core objective is to improve the model's generalization performance, training stability, and privacy protection capabilities, while effectively alleviating the challenges caused by heterogeneous data distribution (non-IID). Knowledge distillation suppresses overfitting by constraining the update magnitude of model parameters and enhances the model's adaptability to different client data distributions. Thus, even with limited or unevenly distributed data, it can still learn more robust feature representations.
[0040] In this embodiment of the invention, during the (t+1)th iteration of federated learning, the model generated by the central server in the tth iteration is defined as the "Teacher Model," while the model updated by each client based on local data in the (t+1)th iteration is called the "Student Model." During this process, the Student Model absorbs knowledge from the Teacher Model through a knowledge distillation loss function, achieving knowledge distillation. Because the Teacher Model integrates global information, its knowledge is more comprehensive and consistent. With the guidance of the Teacher Model during the update process, the Student Model can more effectively learn the general characteristics of the global data distribution, thereby significantly enhancing its generalization ability and stability.
[0041] The loss function expression for knowledge distillation of the teacher and student models using local data, the received dataset, and training data is as follows: ; ; ; in, This represents the knowledge distillation loss value, used to measure the difference in features output by the student model and the teacher model. Representing the teacher model, Representing the student model, This represents the features output by the teacher model. This represents the features output by the student model. Representing feature dimension, , This represents the feature probability distribution after normalization.
[0042] In this embodiment of the invention, the total loss of each local client training model consists of contrastive learning loss and knowledge distillation loss, expressed as: ; in, This represents the total loss of the model trained on the local client. This represents a hyperparameter used to control the degree of knowledge distillation loss.
[0043] Specifically, in the model aggregation process of federated learning, this embodiment of the invention employs a dynamic weight allocation mechanism based on forgetting loss to optimize the update of the global model. The difference between the global model and the local models is quantified using Euclidean distance to obtain the forgetting loss value, reflecting the degree to which the local model forgets global knowledge during training. Subsequently, through exponential decay and normalization operations, weights are assigned to the model parameters uploaded by each local client to ensure the rationality and dynamism of the weight allocation. Finally, the parameter updates of the teacher model are achieved through a weighted average, enabling the teacher model to retain global knowledge while better adapting to the personalized needs of each local client, thereby achieving more efficient knowledge transfer and model optimization.
[0044] Specifically, step 4 includes: The central server receives model parameters from various local clients; The aggregate weights of each uploaded set of model parameters are calculated using the forgetting loss, which measures the difference between the teacher model and the local model. This forgetting loss is specifically implemented using Euclidean distance, and the calculation expression is as follows: ; in, Indicates the loss value for forgetting. Represents the global model. Indicates the first A local model This represents the model parameters of the teacher model. Indicates the first Model parameters for a local model. This represents the Euclidean distance function, used to quantify the difference between two parameter vectors. The larger the difference, the more global knowledge the local model forgets during personalized training. Based on the forgetting algorithm, aggregate weights are assigned to the model parameters uploaded by each local client to reflect their importance in global knowledge preservation. The expression for calculating the aggregate weights is as follows: ;
[0045] in, Indicates the first The aggregate weights of each set of model parameters uploaded by each local client. This indicates the total number of local clients. This represents an exponential function used to decay the forgetting loss value, ensuring that model parameters uploaded from the local client with larger forgetting loss values receive higher weights. This represents a normalization operation, used to ensure that the sum of the aggregated weights of the model parameters of all local clients is 1, thereby achieving a reasonable allocation of aggregated weights; After the aggregated weights are assigned, the expression for aggregating all uploaded model parameters based on the aggregated weights of each group of model parameters is as follows: ; in, This represents the model parameters of the updated teacher model.
[0046] In this way, the teacher model not only retains global knowledge during the update process, but also absorbs personalized information from each local client, thereby improving the model's generalization ability and adaptability.
[0047] Specifically, the training termination conditions include: In practical applications, after multiple iterations, the performance improvement of the model will gradually stabilize. In order to prevent the model training process from going on indefinitely and to provide a time constraint for the model training, the training is terminated when the number of iterations reaches a preset threshold. The preset threshold is the maximum number of iterations set, which is used to ensure the efficiency of model training and the rationality of resource utilization to a certain extent. Training is terminated when the change in the model parameters of the teacher model is less than a preset threshold for several consecutive iterations. For example, if the change in the model parameters of the teacher model is less than 0.001 for 5 consecutive iterations, it can be considered that the teacher model has reached a stable state. Continuing training may not bring significant performance improvement and may even lead to overfitting. Stopping training at this time can ensure that the model maintains good generalization ability while avoiding unnecessary waste of computing resources. Training is terminated when the difference between the model parameters of the teacher model and the model parameters of each local model is less than a preset threshold in several consecutive iterations. For example, if the forgetting loss between the local model and the teacher model is less than 0.01 in three consecutive iterations, it means that the local model has absorbed the knowledge of the teacher model well and the teacher model can also adapt well to the personalized needs of each local client. At this time, training can be stopped to ensure the performance consistency of the teacher model on different local clients. Training is terminated when the teacher model's performance metrics on the validation set reach or exceed a preset threshold. For example, if the teacher model's accuracy on the validation set reaches 90% or more, and its recall and F1 score both meet the system's requirements, it indicates that the teacher model has high performance and can be stopped and put into use. The performance metrics can be accuracy, recall, or F1 score, etc.
[0048] To verify the effectiveness of the embodiments of the present invention, a comparative experiment was designed to compare the performance of the provided method (contrastive learning + knowledge distillation) with existing federated learning algorithms (FedAvg, FedProx, FedMoon) in handling federated learning tasks. A total of 40 clients were set up in the experiment. This indicates the proportion of clients selected in each training round, such as... When the learning rate is 0.3, 30% of the 40 clients are randomly selected for training in each round. The Dirichlet method is used for data allocation, ensuring high heterogeneity among the clients to simulate data allocation in real-world applications. The client-side local learning rate is set to 0.01, and the learning rate decay is set to 0.99. The number of client-side local update rounds is set to 5 rounds, for a total of 200 rounds.
[0049] surface Comparison of experimental results table ; As can be seen from Table 1 above, compared with the other three federated learning algorithms, the method provided by the present invention improves the performance in scenarios where 30%, 50%, and 80% of clients are randomly selected, proving that the combination of contrastive learning and knowledge distillation can improve the performance of federated learning.
[0050] This invention employs multiple local clients to collect local data and data labels, and uploads the data labels to a central server. The central server then filters a public dataset based on the data labels, selecting datasets that match them, and distributes the initial model and the selected datasets to each local client. Each local client performs data augmentation on its local data and the received dataset to obtain training data. Using the local data, the received dataset, and the training data, it performs comparative learning training on the initial model to obtain a student model. The student model's parameters are then uploaded to the central server, and the aggregation weights of each uploaded set of model parameters are calculated using forgetting loss. All uploaded model parameters are then aggregated to obtain a teacher model. The teacher model is then distributed to each local client, and knowledge distillation is performed on both the teacher and student models using the local data, the received dataset, and the training data to obtain a local model. The model parameters of the local model are uploaded to the central server, and the aggregation weight calculation and aggregation are repeated until the teacher model meets the preset training termination condition. The global model is then output for graded early warning of internet addiction behavior. Compared with the prior art, this embodiment of the invention filters public datasets by obtaining data labels from local clients, and trains the model together with local data after data augmentation on each local client. This effectively alleviates the overfitting problem caused by few-shot learning and enhances the generalization ability of the model. During the training process of the local model, contrastive learning is introduced to mine the intrinsic results of the data, enhancing the model's adaptability to data from different local clients. Knowledge distillation is introduced to constrain the model update process to prevent overfitting. In the model aggregation stage, the aggregation weights of the model parameters are calculated through forgetting loss, and the aggregation strategy is dynamically adjusted to reduce knowledge loss and improve the stability and adaptability of the model.
[0051] The above description represents the preferred embodiments of the present invention. It should be noted that those skilled in the art can make various improvements and modifications without departing from the principles of the present invention, and these improvements and modifications should also be considered within the scope of protection of the present invention.
Claims
1. A federated learning method based on contrastive learning and knowledge distillation, characterized in that, include: Step 1: Multiple local clients collect local data and data tags, and upload the data tags to the central server. The local data includes network behavior data, physiological signal data, behavioral posture data, and background data. Step 2: The central server filters the datasets that match the data labels in the public dataset based on the data labels, and sends the initialized model and the filtered datasets to each local client. Step 3: For each local client, the local client performs data augmentation on the local data and the received dataset to obtain training data. The initial model is then trained using the local data, the received dataset, and the training data to obtain a student model. The model parameters of the student model are then uploaded to the central server. Step 4: The central server calculates the aggregate weight of each set of uploaded model parameters using the forgetting loss, and aggregates all uploaded model parameters according to the aggregate weight of each set of model parameters to obtain the teacher model. Step 5: Distribute the teacher model to each local client, use the local data, the received dataset, and the training data to perform knowledge distillation on the teacher model and the student model to obtain a local model, and upload the model parameters of the local model to the central server. Repeat the aggregation weight calculation and model parameter aggregation until the teacher model meets the preset training termination condition, and output the global model for the graded early warning of internet addiction behavior.
2. The federated learning method based on contrastive learning and knowledge distillation according to claim 1, characterized in that, The local client performs data augmentation on the local data and the received dataset to obtain training data, including: The local data and the image data in the received dataset are flipped to obtain the first enhanced data; The local data and other data in the received dataset, excluding image data, are augmented to obtain second enhanced data; The training data is composed of the first augmented data and the second augmented data.
3. The federated learning method based on contrastive learning and knowledge distillation according to claim 2, characterized in that, The loss function expression for performing contrastive learning training on the initialized model using the local data, the received dataset, and the training data is as follows: ; in, This represents the contrastive learning loss value. This represents the local data features extracted through model initialization. This represents the features extracted from the training data through model initialization. This represents the features of the received dataset extracted through model initialization. Represents the similarity function. Indicates temperature parameter, This indicates the total number of data points.
4. The federated learning method based on contrastive learning and knowledge distillation according to claim 3, characterized in that, The loss function expression for knowledge distillation of the teacher model and the student model using the local data, the received dataset, and the training data is as follows: ; in, This represents the knowledge distillation loss value. Representing the teacher model, Representing the student model, This represents the features output by the teacher model. This represents the features output by the student model. Representing feature dimension, , This represents the feature probability distribution after normalization.
5. The federated learning method based on contrastive learning and knowledge distillation according to claim 4, characterized in that, The formula for calculating the forgetting loss is: ; in, Indicates the loss value for forgetting. Represents the global model. Indicates the first A local model This represents the model parameters of the teacher model. Indicates the first Model parameters for a local model. This represents the Euclidean distance function.
6. The federated learning method based on contrastive learning and knowledge distillation according to claim 5, characterized in that, The central server calculates the aggregate weights of each set of model parameters uploaded using the forgetting loss method using the following expression: ; in, Indicates the first The aggregate weights of each set of model parameters uploaded by each local client. This indicates the total number of local clients. This represents an exponential function.
7. The federated learning method based on contrastive learning and knowledge distillation according to claim 6, characterized in that, The expression for aggregating all uploaded model parameters based on the aggregated weights of each group of model parameters is: ; in, This represents the model parameters of the updated teacher model.
8. The federated learning method based on contrastive learning and knowledge distillation according to claim 7, characterized in that, The training termination conditions include: Training terminates when the number of iterations reaches a preset threshold. Training is terminated when the change in the model parameters of the teacher model is less than a preset threshold in several consecutive iterations. Training is terminated when the difference between the model parameters of the teacher model and the model parameters of each local model is less than a preset threshold in several consecutive iterations. Training is terminated when the performance metrics of the teacher model on the validation set reach or exceed a preset threshold.
Citation Information
Patent Citations
Federal learning algorithm based on diffusion model and weight adaptive knowledge distillation
CN116665000A
Self-adaptive knowledge screening method and system in federal learning
CN117422147A
Federal learning knowledge distillation method based on data-free driving
CN119578504A
Federal learning method based on partial label mask weighted distillation
CN120163260A