Personalized federal contrast learning image classification method based on super network

By using hypernetwork for feature processing and fusion in federated learning, and local training combined with comparative learning methods, the performance degradation of existing federated learning methods on non-independent and homogeneous data is solved, and the accuracy improvement and generalization ability of the personalized model are achieved.

CN120014337APending Publication Date: 2025-05-16CHINA UNIV OF PETROLEUM (EAST CHINA)
View PDF 0 Cites 1 Cited by

Patent Information

Application Number
CN202510080813.4
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-01-20
Publication Date
2025-05-16

AI Technical Summary

Technical Problem

When existing federated learning methods deal with non-independent and homogeneous data, it is difficult to effectively integrate data characteristics, making it difficult for personalized models to accurately capture the real characteristics of data and affect model performance.

Method used

A personalized federated contrast learning image classification method based on hypernetwork is adopted to extract client descriptors through embedded networks, and feature advanced processing and fusion are used for hypernetwork to generate classification models. This method combines the comparative learning method for local training to fully capture the similarity and difference information between various classes.

Benefits of technology

It improves the ability to adapt to the differences in data distribution across clients, improves the accuracy of personalized models, and improves the generalization ability of the model.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120014337A_ABST
    Figure CN120014337A_ABST
Patent Text Reader

Abstract

The invention discloses a supernetwork-based personalized federal contrast learning image classification method, which comprises the following steps that: a server initializes an embedded network, a supernetwork and a global category vector, and issues the embedded network and the global category vector to a client; the client extracts a client descriptor by using the embedded network model and sends the client descriptor to the server; the server generates a classification model for the client based on the client descriptor by using a super network model, and sends the classification model to each client; the client calculates a local category vector based on the local data and trains a classification model by using a comparative learning mode; the client sends the update quantity of the classification model and the local category vector to the server; the server updates the global category vector, the embedded network and the super network; the server issues the trained embedded network to the client; the client extracts a client descriptor based on the embedded network and sends the client descriptor to the server; the server generates a classification model by using a super network based on the client descriptor, and sends the classification model to each client; and the client performs image classification by using the classification model. The method has the beneficial effects that through joint modeling of the embedded network and the super network, local training is carried out in a comparative learning mode, and the image classification accuracy in a data heterogeneous scene is improved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the fields of federated learning, contrastive learning and image classification, and in particular to a hypernetwork-based personalized federated contrastive learning image classification method. Background Art

[0002] Image classification is a key technology in computer vision, which aims to assign predefined category labels to target images. With the increasing awareness of data privacy protection, it is challenging to centrally obtain data for training models.

[0003] Federated learning is an emerging distributed collaborative training method that adopts a decentralized model. It can achieve efficient model training and sharing among multiple clients while protecting data privacy. The training process is that the client and the server continuously distribute, train, and aggregate model parameters until the model performance reaches the standard or the client converges.

[0004] However, the data in the federated system is often not independent and identically distributed. Traditional federated learning algorithms rely on simple averaging methods for model aggregation, which is difficult to handle, resulting in a decrease in model performance.

[0005] To this end, personalized federated learning allows clients to learn personalized models based on their own data and task requirements, improving performance on non-IID data. Hypernetwork is a commonly used personalization method that generates personalized model parameters for each client based on the client descriptor and adapts to local data.

[0006] However, these technologies still have shortcomings. In personalized federated learning based on hypernetworks, client descriptors cannot effectively integrate the actual features of the data, making it difficult for personalized models to accurately capture the true characteristics of the data, affecting model performance. At the same time, the data in real federated systems is complex, and existing methods have low sensitivity to data features and are difficult to capture fine-grained features, limiting the accuracy and generalization ability of the model.

[0007] Currently, no effective solution has been proposed for the problems in the related technologies. Summary of the invention

[0008] In view of the problems in the related art, the present invention proposes a personalized federated contrastive learning image classification method based on a hypernetwork, which is characterized by overcoming the above technical problems existing in the existing related art. To this end, the specific technical solution adopted by the present invention is as follows:

[0009] A personalized federated contrastive learning image classification method based on a hypernetwork, characterized in that the method comprises the following steps:

[0010] S1. Define the global training round of the federated system as T, the local training round as E, the global learning rate as α, the local learning rate as η, the global category set as K, and the central server initializes the embedded network model η v , Hypernetwork Model η h and the global category vector P;

[0011] S2, the client sends the embedded network model and the global category vector to each client;

[0012] S3. Define the local category set of the i-th client as K i , the local data set is D i ={x i ,y i}, x i Represents a picture sample, y i Indicates the label corresponding to the image sample, D i Input η v , get the client descriptor v i ;

[0013] S4: After receiving the client descriptor sets sent by N clients, the central server inputs them into n h , obtain the client's local classification model θ set and send it to each client;

[0014] S5, the classification model θ received by the i-th client i After that, D i Input to θ i , using local category vector calculation and local training method based on contrastive learning to perform E rounds of local training, and obtain the local category vector p of the i-th client i and the model update Δθ i ;

[0015] S6, after receiving the local category vector sets sent by N clients, the central server updates the global category vector using the category vector average update method;

[0016] S7: After receiving the model update amount set sent by N clients, the central server updates the embedded network model η using a model update method based on the chain derivation rule. v and the hypernetwork model η h ;

[0017] S8, repeat S2 to S6 T times to obtain the trained embedding network model η v , Hypernetwork Model η h ;

[0018] S9. Deploy the model on the client that needs to perform image classification to perform image classification tasks.

[0019] Furthermore, the local category vector calculation and the local training method based on contrastive learning in S5 include the following steps:

[0020] S5-1, classification model θ of the i-th client i Contains feature extraction layers and the classification layer D i Input to Get the vector p, average the vectors of each category in p, and get the local category vector p i ={p k |k∈K i};

[0021] S5-2. Replace the vectors of each category in p with the corresponding category vector P in the global category vector P k Get the reference vector p r ;

[0022] S5-3, x i Input θ i , the output With y i Calculating cross entropy loss

[0023] S5-4. Calculate the contrast vector using the global category vector P

[0024] S5-5, using p, p r and P c Computing contrast loss Get the loss function l = l ce +l c , find the gradient Update local model Calculate the model update amount

[0025] Furthermore, the category vector average updating method in S6 includes the following steps:

[0026] Count the number of times category k appears in the local category vector set C k , sum all category k vectors in the local category vector set and divide by C k , get the global k-category vector P k , the global category vector P = {P k |k∈K}.

[0027] Furthermore, the model updating method based on the chain rule in S7 includes the following steps:

[0028] S7-1. Using the model update amount Δθ of the i-th client i , and obtain the updated gradient of the embedded network model HyperNetwork Model Update Gradient

[0029] S7-2. Use the embedded network model update gradient of each client and the super network model update gradient to obtain the updated embedded network model Hypernetwork Model

[0030] The beneficial effects of the present invention are:

[0031] (1) The present invention proposes a hypernetwork-based personalized federated contrastive learning image classification method, which uses an embedded network to extract client descriptors containing feature distribution information. The hypernetwork further performs advanced feature processing and fusion to generate a classification model. This method realizes the joint modeling of data features by the embedded network and the hypernetwork, improves the ability to adapt to cross-client data distribution differences, and improves the accuracy of the personalized model.

[0032] (2) Construct category vectors for all categories and use contrastive learning to perform local training. The loss function fully captures the similarities and differences between categories to improve the generalization ability of the model. BRIEF DESCRIPTION OF THE DRAWINGS

[0033] In order to more clearly illustrate the embodiments of the present invention or the technical solutions in the prior art, the drawings required for use in the embodiments will be briefly introduced below. Obviously, the drawings described below are only some embodiments of the present invention. For ordinary technicians in this field, other drawings can be obtained based on these drawings without paying creative work.

[0034] Figure 1 The figure is a flow chart of the technical solution proposed by the present invention. DETAILED DESCRIPTION

[0035] To further illustrate each embodiment, the present invention provides drawings, which are part of the disclosure of the present invention and are mainly used to illustrate the embodiments and can be used in conjunction with the relevant descriptions in the specification to explain the operating principles of the embodiments. With reference to these contents, ordinary technicians in the field should be able to understand other possible implementations and advantages of the present invention. The components in the figures are not drawn to scale, and similar component symbols are generally used to represent similar components.

[0036] According to an embodiment of the present invention, a personalized federated contrastive learning image classification method based on a hypernetwork is provided. Figure 1As shown, it is applied to a federated system consisting of a central server and N clients, and is performed in the following steps:

[0037] S1. Define the global training round of the federated system as T, the local training round as E, the global learning rate as α, the local learning rate as η, the global category set as K, and the central server initializes the embedded network model η v , Hypernetwork Model η h and the global category vector P;

[0038] S2. The client sends the embedded network model η v and the global category vector P to each client;

[0039] S3. Define the local category set of the i-th client as K i , the local data set is D i ={x i ,y i}, x i Represents a picture sample, y i Indicates the label corresponding to the image sample, D i Input η v , get the client descriptor v i ;

[0040] S4: After receiving the client descriptor sets sent by N clients, the central server inputs them into n h , obtain the personalized local classification model θ set of each client, and send it to each client respectively;

[0041] S5, the classification model θ received by the i-th client i After that, D i Input to θ i , using local category vector calculation and local training method based on contrastive learning to perform E rounds of local training, and obtain the local category vector p of the i-th client i and the model update Δθ i ;

[0042] S5-1, classification model θ of the i-th client i Contains feature extraction layers and the classification layer D i Input to Get the vector p, add up the vectors of any category k in p, and divide it by the number of times the category k vector appears to calculate the average category vector p k , get the local category vector p i ={p k |k∈K i};

[0043] S5-2. Replace the vector of any category k in p with the corresponding category vector P in the global category vector P k Get the reference vector p r ;

[0044] S5-3, x i Input θ i , the output With y i Calculating cross entropy loss

[0045] S5-4. Calculate the contrast vector using the global category vector P

[0046] S5-5, using p, p r and P c Computing contrast loss Get the loss function l = l ce +l c , find the gradient pass Update the local model and further calculate the model update amount

[0047] S6, after receiving the local category vectors sent by each of the N clients, the central server updates the global category vector using the category vector average update method;

[0048] S6-1. Count the number of times category k appears in the local category vector set C k , sum all category k vectors in the local category vector set and divide by C k , get the global k-category vector P k , the updated global category vector is P = {P k |k∈K};

[0049] S7: After receiving the model update amount set sent by N clients, the central server updates the embedded network model η using a model update method based on the chain derivation rule. v and the hypernetwork model η h ;

[0050] S7-1. Using the model update amount Δθ of the i-th client i , and obtain the updated gradient of the embedded network model HyperNetwork Model Update Gradient

[0051] S7-2. Use the embedded network model update gradient of each client and the super network model update gradient to obtain the updated embedded network model Hypernetwork Model

[0052] S8, repeat S2 to S6 T times to obtain the trained embedding network model η v , Hypernetwork Model η h ;

[0053] S9, deploy the model on the client that needs to perform image classification to perform image classification tasks;

[0054] S9-1, deploy the trained embedding network model and super network model on the client that needs to perform image classification, and input the local data set D into η v , get v, input v into η h , get the classification model θ;

[0055] S9-2. After θ is deployed on the client that needs to perform image classification, the image samples are input into the model to perform the image classification task and obtain the image classification result.

Claims

1. A personalized federated contrastive learning image classification method based on a hypernetwork, characterized in that: The method comprises the following steps: S1. Define the global training round of the federated system as T, the local training round as E, the global learning rate as α, the local learning rate as η, the global category set as K, and the central server initializes the embedded network model η v , Hypernetwork Model η h and the global category vector P; S2, the client sends the embedded network model and the global category vector to each client; S3. Define the local category set of the i-th client as K i , the local data set is D i ={x i ,y i }, x i Represents a picture sample, y i Indicates the label corresponding to the image sample, D i Input η v , get the client descriptor v i ; S4: After receiving the client descriptor sets sent by N clients, the central server inputs them into n h , obtain the client's local classification model θ set and send it to each client; S5, the classification model θ received by the i-th client i After that, D i Input to θ i , using local category vector calculation and local training method based on contrastive learning to perform E rounds of local training, and obtain the local category vector p of the i-th client i and model update Δθ i ; S6, after receiving the local category vector sets sent by N clients, the central server updates the global category vector using the category vector average update method; S7: After receiving the model update amount set sent by N clients, the central server updates the embedded network model η using a model update method based on the chain derivation rule. v and the hypernetwork model η h ; S8, repeat S2 to S6 T times to obtain the trained embedding network model η v , Hypernetwork Model η h ; S9. Deploy the model on the client that needs to perform image classification to perform image classification tasks.

2. According to claim 1, a personalized federated contrastive learning image classification method based on a hypernetwork is characterized in that: The local category vector calculation and the local training method based on contrastive learning in S5 include the following steps: S5-1, classification model θ of the i-th client i Contains feature extraction layers and the classification layer D i Input to Get the vector p, average the vectors of each category in p, and get the local category vector p i ={p k |k∈K i }; S5-2. Replace the vectors of each category in p with the corresponding category vector P in the global category vector P k Get the reference vector p r ; S5-3, x i Input θ i , the output With y i Calculating cross entropy loss S5-4. Calculate the contrast vector using the global category vector P S5-5, using p, p r and P c Computing contrast loss We get the loss function l = l ce +l c , find the gradient Update local model Calculate the model update amount 3. The method of image classification based on hypernetwork personalized federated contrastive learning according to claim 1, characterized in that: The method for updating the average category vector in S6 comprises the following steps: Count the number of times category k appears in the local category vector set C k , sum all category k vectors in the local category vector set and divide by C k , get the global k-category vector P k , the global category vector P = {P k |k∈K}.

4. The method of image classification based on personalized federated contrastive learning based on hypernetwork according to claim 1, characterized in that: The model updating method based on the chain derivation rule in S7 comprises the following steps: S7-1. Using the model update amount Δθ of the i-th client i , and obtain the updated gradient of the embedded network model HyperNetwork Model Update Gradient S7-2. Use the embedded network model update gradient of each client and the super network model update gradient to obtain the updated embedded network model Hypernetwork Model

Citation Information

Cited By

  • Federal personalized human activity recognition training method based on hypernetwork

    CN121031817A