Federal learning survival prediction method based on hypergraph
By introducing hypergraph structure and knowledge distillation technology into federated learning, the high-order correlation of pathological images is effectively utilized, and the existing federated learning methods are solved, and efficient and accurate survival prediction is achieved.
Patent Information
- Application Number
- CN202411951989.4
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2024-12-27
- Publication Date
- 2025-05-27
- Estimated Expiration
- 2044-12-27
AI Technical Summary
The existing federated learning methods fail to effectively utilize the structural information of the data in survival prediction tasks, resulting in insufficient robustness and inference efficiency of multi-center data processing.
A federated learning method based on hypergraph is proposed to model higher-order correlations in gigapixel pathological images through hypergraph structures, and introduce hypergraph knowledge distillation technology to transfer higher-order structural information to shallow neural networks.
It significantly improves the robustness and consistency of the multi-center survival prediction task, reduces the model's inference time, and achieves efficient and accurate survival prediction while protecting privacy.
Smart Images

Figure CN120048507A_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the fields of distillation learning, distributed machine learning, and hypergraph representation learning, and particularly relates to a hypergraph-based federated learning survival prediction method. Background Art
[0002] Survival prediction is one of the key objectives of precision medicine, as accurate survival assessment can help doctors select appropriate treatment methods for individual patients. To achieve this goal, a large amount of data must be used to train the prediction model and prevent overfitting. However, due to differences in data sources among different institutions and concerns about privacy and ownership issues in data sharing, collecting patient data for disease prediction is challenging. To facilitate the integration of cancer data from different institutions without violating privacy laws, federated learning-based survival prediction tasks have gradually emerged, such as methods like AdFed. However, existing methods do not incorporate the application of the structural information of data in federated learning, resulting in deficiencies in enhancing the robustness and inference efficiency of multi-center data processing. Summary of the Invention
[0003] Aiming at the deficiencies existing in the prior art, the present invention provides a hypergraph-based federated learning survival prediction method. A novel hypergraph federated learning (HGFL) is proposed for survival prediction of multi-center pathology images. At the local center, HGFL first uses the hypergraph structure to model the high-order correlations within gigapixel pathology images. In addition, we introduce a hypergraph knowledge distillation technique to transfer the high-order structural information to a shallow neural network, retaining the ability to reason about high-order correlations while significantly improving the inference efficiency of the model. At the global center, the local center only needs to share the parameters of its shallow neural network with the central server. Through parameter aggregation, the shallow networks of each center can learn the high-order structural information of other centers. Experiments on multi-center datasets show that the proposed method not only effectively improves the robustness and consistency of the multi-center survival prediction task, but also significantly reduces the inference time of the model, thus achieving efficient and accurate survival prediction while protecting privacy.
[0004] The model proposed in the present invention uses the features and structural information in multimodal data for self-supervised clustering. First, for r centers, n pathological images are obtained for each center, each pathological image is evenly divided into p patch images, and the image embedding of each patch image is obtained using a pre-trained feature extraction model, and the coordinate information of each patch image is recorded, and finally the two are fused; then the nearest neighbor algorithm is used to calculate the nearest neighbors of the p patch images of each pathological image, and the patches that are neighbors of each other are connected to form a hypergraph to obtain an association matrix, in which each pathological image forms an independent hypergraph and association matrix; then the hypergraph convolutional neural network (HGSurvNet) is used to learn and obtain the features of the entire pathological image; furthermore, the hypergraph distillation is used to update the local network; then the global network is updated according to the local network of each center, and survival prediction is performed.
[0005] A hypergraph-based federated learning survival prediction method includes the following steps:
[0006] Step 1: For r centers, obtain n pathological images for each center, and evenly divide each pathological image into p patch images, where p is an optional hyperparameter. Use the pre-trained feature extraction model to obtain the image embedding of each patch image, and record the coordinate information of each patch image. Finally, fuse the two to obtain the feature information of the dataset.
[0007] Step 2: Use the nearest neighbor algorithm to calculate the nearest neighbors of the p patch images of each pathology image, connect the neighboring patches to form a hypergraph, and obtain the association matrix.
[0008] Step 3: Learn and obtain the features of the entire pathology image through the hypergraph convolutional neural network (HGSurvNet).
[0009] Step 4: Use hypergraph distillation to update the local network.
[0010] Step 5: Update the global network based on the local network of each center and make survival predictions.
[0011] Furthermore, the specific steps of step 1 are as follows:
[0012] For r centers, we independently obtain their own pathology image-related datasets. i The dataset D = {x i} i∈{1,…,N} , where N is the total number of samples. i The image is evenly divided into p patches and the image embedding of each pathology image is extracted using the pre-trained feature extraction model with frozen parameters as F. i ={f 1 ,f2 ,…f j ,…f p} ∈{1,…,N} , where f is the image embedding extracted by the pre-trained feature extraction model from the patch image.
[0013] Thus, the image embedding of the dataset D is F = {F i} i∈{1,…,N} .
[0014] After that, record the coordinates of the upper left corner of each patch image on the pathological image as c = (a, b) 0≤a≤h,0≤b≤w , where h is the height of the pathological image and w is the width of the pathological image. Further, the coordinate information of each pathological image is obtained as C i = {c 1 , c 2 ,…c j ,…c p} i∈{1,…,N} .
[0015] Thus, the coordinate information C of the dataset D is obtained as C = {C i} i∈{1,…,N} .
[0016] Finally, fuse the image embedding F and the coordinate information C of the dataset into the feature information of the dataset:
[0017] X = {X i} i∈{1,…,N}
[0018] where the feature information X of a single pathological image i is defined as follows:
[0019] X i = F i + C i = {f 1 + c 1 , f 2 + c 2 ,…f j + c j ,…f p + c p} i∈{1,…,N}
[0020] where f j + c j is the concatenation of the two vectors.
[0021] In one embodiment, the pre-trained feature extraction model uses ResNet-34.
[0022] Furthermore, the specific steps of step two are as follows:
[0023] For each pathological image xi The p patch images in are mapped to the vertex set of the hypergraph, where vertex v j is the mapping of x in the hypergraph, and define i as the vertex feature, where By calculating the Euclidean distance between vertex features to measure the distance between vertices. The Euclidean distance is defined as follows:
[0024]
[0025] where d(v a , v b ) represents the Euclidean distance between vertex v a and v b . and represent the features of vertex v a and v b respectively. C represents the dimension of the feature. and represent the values of the features of vertex v a and v b in the c-th dimension respectively. Each pathological image data X i in the dataset D is mapped to the vertices of its respective hypergraph using the above method.
[0026] Next, for each vertex, scale the Euclidean distance between this vertex and other vertices to [0, 1]. The specific formula is as follows:
[0027]
[0028] d norm (v a , v b ) = d(v a , v b ) / d(v a ) max
[0029] where d(v a ) max represents the maximum Euclidean distance from vertex v a to other vertices; d norm (v a , v b ) is the Euclidean distance from the scaled vertex v a to other vertices.
[0030] Then, the ∈-ball (∈-neighborhood) and K-NN (k-nearest neighbors) graph construction methods are used simultaneously to obtain the nearest neighbors of the modal data in the modal data set. Among them, K-NN is a graph structure constructed based on the relationship between each point and its k nearest neighbors; ∈-ball is a graph structure constructed based on the relationship between each vertex and its neighbors whose distance from it is less than a certain radius ∈. The specific steps are as follows:
[0031] First, use the ∈-ball graph construction method (∈-neighborhood) to calculate the neighbor subset of vertex v a as follows:
[0032] neigh th (v a ) = {v b | d norm (v a , v b ) < th}
[0033] where neigh th (v a ) is the neighbor subset of vertex v a obtained by the ∈-ball graph construction method, and th is a pre-set distance threshold.
[0034] At the same time, use the K-NN graph construction method (K-nearest neighbors) to calculate the neighbor subset of vertex v a as follows:
[0035] neigh ne (v a ) = {v b | (d(v a , v b ) < sortd(v a , v b ))[K]}
[0036] where K is the pre-set number of neighbor vertices, neigh ne (v a ) represents the neighbor subset of vertex v a obtained by the K-NN graph construction method, sort(·) represents the distance sorting algorithm, and | neigh ne (v a ) | = K.
[0037] In summary, through the two graph construction methods, for each vertex v a , neighbor vertex sets neigh th (v a ) and neigh ne (v a ) can be generated, and the final neighbor vertex set is as follows:
[0038] neigh(v a ) = {neigh th (v a ), neigh ne (v a )}
[0039] Define the hyperedge set according to the neighbor vertex set as follows:
[0040] ε = {neigh(v a )} a∈{1,…,N}
[0041] The incidence matrix can be obtained according to the neighbor vertex set , where the incidence matrix is defined as follows:
[0042]
[0043] where H(a, b) is the element value of the a-th row and b-th column of the incidence matrix H.
[0044] Furthermore, Step 3 is specifically as follows:
[0045] Learn and obtain the feature Z of the entire pathology image through the hypergraph convolutional neural network (HGSurvNet) i , first obtain the result of l - time convolution of the feature matrix of each pathology image through the hypergraph convolution formula The formula is as follows:
[0046]
[0047] where l is the number of convolutions of the hypergraph convolutional neural network (a hyperparameter that can be selected during training), H i is the incidence matrix corresponding to the i-th pathology image, is the result of the l - th convolution of the feature matrix of the i-th pathology image, Mask(·) is the mask operator for hyperedge features, Θ l is the parameter to be learned during training, and the pathology image feature set {X i} i∈{1,…,N} shares Θ l , where where c l , c l+1 are optional hyperparameters.
[0048] After l - layer convolution, the vertex features that fuse topological and semantic information are obtained. Further, the convolution results of the feature matrix of the i-th pathology image are added through the pooling operator and a pathology image feature vector is obtained through a linear layer where c o is an optional hyperparameter.
[0049] Furthermore, Step 4 is specifically as follows:
[0050] Define the global network of federated learning as f(w), and the local networks as {f j (w)}. j∈{1,…,r} Among them, the model architectures of the global network and the local networks are the same, and a suitable model architecture is selected according to the memory of the device and the requirements for the running speed.
[0051] Use hypergraph distillation to update the local networks {f j (w)}. To enhance federated learning, use HGSurvNet for soft label guidance and directly distill knowledge from HGSurvNet to f j∈{1,…,r} (w). Therefore, use HGSurvNet as the teacher network and f j (w) as the student network, and use the following objective function to extract knowledge: j where
[0052]
[0053] where is the number of pathological images and h θ (x i ) represents the survival risk probabilities output by HGSurvNet and f j (w) respectively.
[0054] Furthermore, Step 5 is specifically as follows:
[0055] The update of the global network is shown by the following objective:
[0056]
[0057] For machine learning problems, we usually take f j (w) = l(x j , y j ; w), that is, the loss of predicting the example (x j , y j ) using the model parameters w. Since there are r centers and the data is partitioned onto these centers, where is the index set of data points on center j, where Therefore, the above objective can be rewritten as:
[0058]
[0059] For the data of r centers, the pathological image data of each center is processed into X respectively by the method of step one 1 ,…X j ,…X r . The feature information X of the r data sets j is respectively input into its own HGSurvNet and local network f j (w). The local network is updated end-to-end according to the loss function of step four. Finally, the global network f(w) is updated through the above objectives
[0060] Finally, for any pathological image x i , it is processed into feature information X by the method of step 1 i , and then the feature vector Z is further obtained through the trained HGSurvNet i . We use Z i as the input of the global network f(w), and obtain a scalar representing the survival days after the patient is diagnosed as the output
[0061] Features and beneficial effects of the present invention
[0062] Aiming at the deficiencies of survival prediction in the field of distributed training and the limitations of traditional federated learning in terms of poor robustness and performance, the present invention proposes hypergraph federated learning, which utilizes the high-order correlation between WSI data and the excellent modeling ability of hypergraph neural networks to enhance federated learning. The present invention conducts experiments on the survival prediction task, where MLP is retained as the local and global models required for deployment. In addition, the present invention uses hypergraph distillation to narrow the gap between HGSurvNet and MLP, while achieving fast inference, retaining the ability to capture structural information. Experimental results show that through hypergraph convolutional neural networks, this method improves the problems of insufficient accuracy and robustness in the field of federated learning for survival prediction, and improves the accuracy of survival prediction and the robustness of predicting data from different centers Description of the drawings
[0063] Figure 1 is the flow chart of the method of the present invention
[0064] Figure 2 is the comparison effect between this method and other methods
[0065] Figure 3 is the result of this method trained in one center and tested in other centers
[0066] Figure 4 is the result of this method and other methods using the global model to test in other centers
[0067] Figure 5 is the ablation experiment of this method for the hypergraph distillation part Detailed implementation manners
[0068] The technical solution of the present invention will be further described below in conjunction with the accompanying drawings and embodiments.
[0069] As Figure 1 shown, a hypergraph-based federated learning survival prediction method includes the following steps:
[0070] Step 1: For r centers, obtain n pathological images for each center, evenly divide each pathological image into p patch images (p = 2000 in the embodiment), where p is an optional hyperparameter, use a pre-trained feature extraction model to obtain the image embedding of each patch image, and record the coordinate information of each patch image. Finally, fuse the two to obtain the feature information of the dataset. The specific steps are as follows:
[0071] For r centers, independently obtain their respective pathological image-related datasets. For a given dataset D = {x i} i composed of pathological images x i∈{1,…,N} , where N is the total number of samples. Evenly divide each pathological image x i into p patch images and use a pre-trained feature extraction model with frozen parameters to extract the image embedding of each pathological image as F i = {f 1 , f 2 , … f j , … f p} i∈{1,…,N} , where f is the image embedding extracted by the pre-trained feature extraction model from the patch image.
[0072] Thus, the image embedding of dataset D is obtained as F = {F i} i∈{1,…,N} .
[0073] After that, record the coordinates of the upper left corner of each patch image on the pathological image as c = (a, b) 0≤a≤h,0≤b≤w , where h is the height of the pathological image and w is the width of the pathological image. Further, the coordinate information of each pathological image is obtained as C i = {c 1 , c 2 , … c j , … c p} i∈{1,…,N} .
[0074] Thus, the coordinate information of dataset D is obtained as C = {C i} i∈{1,…,N} .
[0075] Finally, fuse the image embedding F and the coordinate information C of the dataset into the feature information of the dataset:
[0076] X = {X i} i∈{1,…,N}
[0077] Among them, the feature information X of a single pathological image i is defined as follows:
[0078] X i = F i + C i = {f 1 + c 1 , f 2 + c 2 , … f j + c j , … f p + c p} i∈{1,…,N}
[0079] Among them, f j + c j is the concatenation of two vectors.
[0080] The pre-trained feature extraction model used is ResNet-34.
[0081] Step 2: Use the nearest neighbor algorithm to calculate the nearest neighbors of the p patch images of each pathological image, connect the mutually neighboring patches to form a hypergraph, and obtain the adjacency matrix.
[0082] Map the p patch images in each pathological image x i to the vertex set of the hypergraph where the vertex v j is the mapping of x i in the hypergraph, and define as the vertex feature, where The distance between vertices is measured by calculating the Euclidean distance between vertex features , and the Euclidean distance is defined as follows:
[0083]
[0084] where d(v a , v b ) represents the Euclidean distance between vertices v a and v b . and represent the features of vertices v a and v b respectively, C represents the dimension of the feature, and respectively represent the value of the feature of vertex v a and v b in the c-th dimension. Each pathological image data X in the dataset D i is mapped to the vertices of its respective hypergraph using the above method.
[0085] Next, for each vertex, the Euclidean distance between this vertex and other vertices is scaled to [0, 1]. The specific formula is as follows:
[0086]
[0087] d norm (v a , v b ) = d(v a , v b ) / d(v a ) max
[0088] where d(v a ) max represents the maximum Euclidean distance from vertex v a to other vertices; d norm (v a , v b ) is the Euclidean distance from the scaled vertex v a to other vertices.
[0089] Then, both the ∈-ball (∈-neighborhood) and K-NN (k-nearest neighbors) graph construction methods are used to obtain the nearest neighbors of the modal data in the modal data set. Among them, K-NN is a graph structure constructed based on the relationship between each point and its k nearest neighbors; ∈-ball is a graph structure constructed based on the relationship between each vertex and its neighbors whose distance is less than a certain radius ∈. The specific steps are as follows:
[0090] First, use the ∈-ball graph construction method (∈-neighborhood) to calculate the neighbor subset of vertex v a . The calculation method is as follows:
[0091] neigh th (v a ) = {v b ∣d norm (v a , v b ) < th}
[0092] where neigh th (v a ) is the neighbor subset of vertex v a obtained by the ∈-ball graph construction method, and th is a pre-set distance threshold.
[0093] Meanwhile, calculate the neighbor subset of vertex v using the K-NN graph construction method (K-Nearest Neighbor), and the calculation method is as follows: a
[0094] neigh ne (v a ) = {v b | (d(v a , v b ) < sortd(v a , v b ))[K]}
[0095] Where K is the number of neighbor vertices set in advance, neigh ne (v a ) represents the neighbor subset of vertex v obtained by the K-NN graph construction method, sort(·) represents the distance sorting algorithm, and |neigh a (v ne )| = K. a
[0096] In summary, through the two graph construction methods, for each vertex v a , the neighbor vertex sets neigh th (v a ) and neigh ne (v a ) can be generated. Finally, the neighbor vertex set is generated as follows:
[0097] neigh(v a ) = {neigh th (v a ), neigh ne (v a )}
[0098] According to the neighbor vertex set, the hyperedge set is defined as follows:
[0099] ε = {neigh(v a )} a∈{1,…,N}
[0100] According to the neighbor vertex set, the incidence matrix can be obtained, where the incidence matrix is defined as follows:
[0101]
[0102] Where H(a, b) is the element value of the a-th row and b-th column of the incidence matrix H.
[0103] Step 3: Learn and obtain the features of the entire pathological image through the hypergraph convolutional neural network (HGSurvNet).
[0104] Learn and obtain the feature Z of the whole pathological image through the Hypergraph Convolutional Neural Network (HGSurvNet). i , first obtain the result of the l-th convolution of the feature matrix of each pathological image through the hypergraph convolution formula The formula is as follows:
[0105]
[0106] where l is the number of convolutions of the hypergraph convolutional neural network (a hyperparameter that can be selected during training), H i is the incidence matrix corresponding to the i-th pathological image, is the result of the l-th convolution of the feature matrix of the i-th pathological image, Mask(·) is the mask operator for hyperedge features, Θ l is the parameter to be learned during training, and the pathological image feature set {X i} i∈{1,…,N} shares Θ l , where where c l , c l+1 are optional hyperparameters.
[0107] After l layers of convolution, vertex features that fuse topological and semantic information are obtained. Further, the convolution results of the feature matrix of the i-th pathological image are added through the pooling operator and passed through a linear layer to obtain the pathological image feature vector where c o are optional hyperparameters.
[0108] Step 4: Use hypergraph distillation to update the local network.
[0109] Define the global network of federated learning as f(w), and the local networks as {f j (w)} j∈{1,…,r} , where the model architectures of the global network and the local networks are the same. An appropriate model architecture can be selected according to the memory of the device and the requirements for the running speed. In theory, any model architecture that does not exceed the maximum memory during training can be selected. For the consideration of less memory occupancy and faster speed, MLP is adopted in this example.
[0110] Use hypergraph distillation to update the local network {f j (w)} j∈{1,…,r} . To enhance federated learning, use HGSurvNet for soft label guidance and directly distill knowledge from HGSurvNet to f j (w). Therefore, use HGSurvNet as the teacher network, f j(w) is the student network, and the following objective function is used to extract knowledge:
[0111]
[0112] where, is the number of pathological images and h θ (x i ) represent the survival risk probabilities output by HGSurvNet and f j (w), respectively.
[0113] Step 5: Update the global network according to the local network of each center and perform survival prediction.
[0114] The update of the global network is shown by the following objective:
[0115]
[0116] For machine learning problems, we usually take f j (w) = l(x j , y j ; w), that is, the loss of predicting the example (x j , y j ) using the model parameter w. Since there are r centers and the data is partitioned onto these centers, where is the index set of data points on center j, where Therefore, the above objective can be rewritten as:
[0117]
[0118] For the data of r centers, the pathological image data of each center is processed into X 1 , … X j , … X r respectively by the method of Step 1. The feature information X j of the r data sets is respectively input into its own HGSurvNet and the local network f j (w) (MLP is adopted in this embodiment). Update the local network end-to-end according to the loss function in Step 4. Finally, update the global network f(w) through the above objective.
[0119] Finally, for any pathological image x i , we process the pathological image into feature information X i by the method of Step 1, and further obtain the feature vector Z i through the trained HGSurvNet. We will Z iAs the input of the global network f(w), a scalar representing the number of days a patient survives after being diagnosed with cancer is obtained as the output.
[0120] For the pathological image datasets collected from three medical centers, namely GY, ZY, and SY. First, we use ResNet to extract visual semantic information and obtain topological information based on the spatial location of the images. Then, we adopt HGSurvNet as the teacher network and MLP as the student network for distillation learning, where the batch size is set to 256 and the number of epochs is set to 30. We use the ADAM optimizer with an initial learning rate of 10 -3 , and the convolutional layer of the hypergraph is set to 2. Further, we update the three local MLP networks of GY, ZY, and SY to the global MLP network by averaging the model weights.
[0121] Figure 2 For the comparative experiment, the method of the present invention is about six percentage points higher than the traditional method. Figure 3 This is the result of testing the data of other centers with the local model of a single center by this method, which is on average 15 percentage points lower than that tested with the global model. Figure 4 When using this method to test with the global model on multiple centers, it is on average 10 percentage points higher than the traditional method. Figure 5 This is the ablation experiment of the hypergraph distillation part of this method. Using hypergraph distillation is on average 5 percentage points higher than not using hypergraph distillation, which proves the effectiveness of the hypergraph distillation network.
[0122] The above content is a further detailed description of the present invention in combination with specific / preferred embodiments, and it cannot be determined that the specific implementation of the present invention is only limited to these descriptions. For those of ordinary skill in the technical field to which the present invention pertains, without departing from the concept of the present invention, they can also make several substitutions or variations to these described embodiments, and these substitution or variation methods should all be regarded as belonging to the protection scope of the present invention.
[0123] The parts not detailed in the present invention belong to the well-known technologies in the art.
Claims
1. A hypergraph-based federated learning survival prediction method, characterized in that: The following steps are involved: Step 1: For r centers, obtain n pathological images for each center, evenly divide each pathological image into p patch images, where p is an optional hyperparameter, use the pre-trained feature extraction model to obtain the image embedding of each patch image, and record the coordinate information of each patch image, and finally fuse the two to obtain the feature information of the dataset; Step 2: Use the nearest neighbor algorithm to calculate the nearest neighbors of the p patch images of each pathology image, connect the neighboring patches to form a hypergraph, and obtain the association matrix; Step 3: Learn and obtain the features of the entire pathology image through the hypergraph convolutional neural network; Step 4: Use hypergraph distillation to update the local network; Step 5: Update the global network based on the local network of each center and make survival predictions.
2. The hypergraph-based federated learning survival prediction method according to claim 1, characterized in that: Step 1 The specific steps are as follows: For r centers, we independently obtain their own pathology image-related datasets. i The dataset D = {x i } i∈{1,...,N} , where N is the total number of samples; each pathological image x i The image is evenly divided into p patches and the image embedding of each pathology image is extracted using the pre-trained feature extraction model with frozen parameters as F. i ={f1, f2, ... f j , ...f p } i∈{1,...,N} , where f is the image embedding extracted from the patch graph by the pre-trained feature extraction model; The image embedding of the dataset D is obtained as F = {F i } i∈{1,...,N} ; Then record the coordinates of the upper left corner of each patch image on the pathology map as c = (a, b) 0≤a≤h,0≤b≤w , where h is the height of the pathological image, w is the width of the pathological image, and the coordinate information of each pathological image is further obtained as C i ={c1, c2, ...c j , ...c p } i∈{1,...,N} ; Thus, the coordinate information C of the data set D is obtained. i } i∈{1,...,N} ; Finally, the image embedding F and coordinate information C of the fusion data set are the feature information of the data set: X={X i } i∈{1,...,N} The characteristic information X of a single pathological image i The following definition: X i =F i +C i ={f1+c1,f2+c2,...f j +c j ,...f p +c p } i∈{1,…,N} where f j +c j Concatenates two vectors.
3. A hypergraph-based federated learning survival prediction method according to claim 1 or 2, characterized in that: The pre-trained feature extraction model uses ResNet-34.
4. The hypergraph-based federated learning survival prediction method according to claim 3, characterized in that: Step 2 The specific steps are as follows: Each pathological image x i The p patch graphs in the graph are mapped to the vertex set of the hypergraph The vertex v j For x i Mapping in the hypergraph, and define is the vertex feature, where By calculating the vertex features The distance between vertices is measured by the Euclidean distance between them, where the Euclidean distance is defined as follows: Where d(v a , v b ) represents the vertex v a and v b The Euclidean distance between and Represents the vertex v a and v b The feature of, C represents the dimension of the feature, and Represents the vertex v a and v b The value of the feature in the cth dimension, each pathological image data X in the dataset D i They are all mapped to the vertices of their respective hypergraphs using the above method; Next, for each vertex, the Euclidean distance between the vertex and other vertices is scaled to [0, 1]. The specific formula is as follows: Where d(v a ) max Represents the vertex v a The maximum Euclidean distance to other vertices; d norm (v a , v b ) is the scaled vertex v a Euclidean distance to other vertices; Then, the ∈-ball and K-NN graph construction methods are used simultaneously to obtain the nearest neighbors of the modal data in the modal data set, where K-NN is a graph structure constructed based on the relationship between each point and its k nearest neighbors; ∈-ball is a graph structure constructed based on the relationship between each vertex and its neighbors whose distance from it is less than a certain radius ∈; the specific steps are as follows: First, use the ∈-ball construction method to calculate the vertex v a The neighbor subset of is calculated as follows: neigh th (v a )={v b |d norm (v a ,v b )<th} Among them th (v a ) is the vertex v obtained by ∈-ball construction a The neighbor subset of , th is the pre-set distance threshold; At the same time, the K-NN graph construction method is used to calculate the vertex v a The neighbor subset of is calculated as follows: neigh ne (v a )={v b |(d(v a ,v b )<sort(d(v a ,v b ))[K]} Among them, K is the number of neighbor vertices set in advance, neigh ne (v a ) represents the vertex v obtained by K-NN construction method a Neighbor subset, sort(·) represents the distance sorting algorithm, |neigh ne (v a )|=K; In summary, through two construction methods, for each vertex v a , can generate neighbor vertex set neigh th (v a ) and neigh ne (v a ), and finally generate the neighbor vertex set as follows: neigh(v a )={neigh th (v a ),neigh ne (v a )} According to the neighbor vertex set, the hyperedge set is defined as follows: ε={neigh(v a )} a∈{1,...,N} According to the neighbor vertex set, the association matrix can be obtained The correlation matrix is defined as follows: Where H(a, b) is the element value in the ath row and bth column of the incidence matrix H.
5. The hypergraph-based federated learning survival prediction method according to claim 4, characterized in that: Step 3 is as follows: The feature Z of the entire pathology image is learned and obtained through the hypergraph convolutional neural network HGSurvNet i First, the hypergraph convolution formula is used to obtain the result of the l-time convolution of the feature matrix of each pathology image. The formula is as follows: Where l is the number of convolutions of the hypergraph convolutional neural network, H i is the correlation matrix corresponding to the i-th pathological image, is the result of the lth convolution of the feature matrix of the i-th pathological image, Mask(·) is the mask operator of the hyperedge feature, Θ l are the parameters that need to be learned during the training process, and the pathological image feature set {X i } i∈{1,...,N} Shared Θ l ,in where c l , c l+1 is an optional hyperparameter; After l layers of convolution, the vertex features that integrate topological and semantic information are obtained. The convolution result of the feature matrix of the i-th pathological image is further pooled through the pooling operator. Add and pass through a linear layer to get the pathology map feature vector where c o is an optional hyperparameter.
6. The hypergraph-based federated learning survival prediction method according to claim 5, characterized in that: Step 4 is as follows: Define the global network of federated learning as f(w) and the local network as {f j (w)} j∈{1,...,r} , where the model architecture of the global network and the local network is the same, and the appropriate model architecture is selected according to the device's memory and the requirements for the running speed; Using Hypergraph Distillation for Local Networks j (w)} j∈{1,...,r} To enhance federated learning, HGSurvNet is used for soft label guidance, and knowledge is directly distilled from HGSurvNet to f j (w); Therefore, HGSurvNet is used as the teacher network, f j (w) is the student network, and the following objective function is used to extract knowledge: in, is the number of pathological images and h θ (x i ) represents HGSurvNet and f j (w) Output survival risk probability.
7. The hypergraph-based federated learning survival prediction method according to claim 6, characterized in that: Step 5 is as follows: The update of the global network is represented by the following objectives: For machine learning problems, take f j (w) = l(x j ,y j ; w), that is, use the model parameter w to analyze the example (x j ,y j ) is used to predict the loss; since there are r centers, the data is divided into these centers, where is the index set of data points on center j, where Therefore, the above goal can be rewritten as: For the data of r centers, the pathological data of each center is processed into X by the method of step 1. 1 , ...X j , ...X r ; The feature information X of r data sets j Input to its own HGSurvNet and local network f respectively j (w); update the local network end-to-end according to the loss function of step 4; finally update the global network f(w) through the above objectives; Finally, for any pathological image x i , processed into feature information X by the method in step 1 i , and then further obtain the feature vector Z through the trained HGSurvNet i , Z i As the input of the global network f(w), a scalar representing the number of days the patient survives after diagnosis is obtained as the output.
Citation Information
Patent Citations
Federal learning scene link prediction method based on knowledge distillation
CN118133945A
Link prediction method based on hypergraph neural network
CN118214707A
Methods and apparatus for analyzing pathology patterns of whole-slide images based on graph deep learning
US20230334662A1