Brain age prediction method based on twinborn pruning attention neural network

By combining the pruned Transformer structure with the twin neural network, the redundant information processing unit is dynamically screened, and the twin neural network model is constructed, which solves the problem of high complexity in high-dimensional rs-fMRI image data calculation data, improves the accuracy and generalization ability of brain age prediction, and is suitable for individual brain health assessment and neurodegenerative disease screening.

CN120388010APending Publication Date: 2025-07-29NANTONG UNIV
View PDF 0 Cites 1 Cited by

Patent Information

Application Number
CN202510582084.2
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-05-07
Publication Date
2025-07-29

AI Technical Summary

Technical Problem

The existing Transformer model has high computational complexity when processing high-dimensional rs-fMRI image data, which is not suitable for direct application in brain age prediction. The traditional method lacks generalization ability on brain image data with significant differences in small samples and individuals.

Method used

Combining the pruning mechanism and twin neural network, through the pruning Transformer structure, redundant information processing units are dynamically screened to build a twin neural network model, and a joint loss function is used to consider structural similarity and label similarity to realize brain age prediction.

Benefits of technology

It significantly reduces calculation costs, improves the accuracy and generalization ability of brain age prediction, enhances the interpretability of the model in clinical application, and is suitable for individual brain health assessment and neurodegenerative disease screening.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120388010A_ABST
    Figure CN120388010A_ABST
Patent Text Reader

Abstract

The invention provides a brain age prediction method based on a twinborn pruning attention neural network, and belongs to the technical field of medical image intelligent diagnosis. While the accuracy of brain age prediction is ensured, the calculation overhead of the model for high-dimensional rs-fMRI image data is reduced, and the generalization ability of the model in a small sample and individual difference significant scene is improved. According to the technical scheme, the method comprises the following steps that S1, resting state functional magnetic resonance imaging of a subject is collected; s2, constructing a pruning module; s3, constructing a twin neural network model; s4, designing a joint loss function to comprehensively consider structural similarity and label similarity; and S5, after model training is completed, inputting test set samples into the trained twin network structure one by one for prediction analysis. The method has the beneficial effect that the accuracy and generalization ability of brain age prediction are improved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technical field of medical image intelligent diagnosis, and in particular to a brain age prediction method based on a twin pruned attention neural network. Background Art

[0002] As humans age, cognitive abilities gradually decline, and the brain undergoes a series of structural and functional changes. These changes manifest as reduced brain tissue volume, enlarged ventricles, decreased gray matter density, and decreased functional connectivity, particularly in middle-aged and elderly individuals. Furthermore, in certain populations, such as those with neurological diseases or those who have experienced trauma, the aging process of the brain exhibits greater individual variability. Therefore, developing metrics that can objectively assess the actual state of an individual's brain has become particularly important. In recent years, the concept of brain age has provided a new quantitative approach to studying brain health. Brain age refers to physiological brain age estimated through objective metrics such as imaging data. Discrepancies between brain age and chronological age are believed to be closely associated with the risk of early-stage disease, neurodegeneration, and cognitive impairment. Brain age prediction technology has potential clinical value in assisting the diagnosis of Alzheimer's disease, Parkinson's disease, stroke sequelae, and post-traumatic stress disorder. Therefore, developing efficient, robust, and interpretable brain age prediction models has become a key research direction in neuroimaging analysis.

[0003] The Transformer has achieved outstanding performance in natural language processing and image processing tasks due to its powerful global modeling capabilities. However, its standard architecture has high computational complexity and is not suitable for direct application to brain imaging data at the voxel level or with extremely high temporal dimensions. Summary of the Invention

[0004] The objective of the present invention is to provide a brain age prediction method based on a twin pruning attention neural network. By introducing a processing unit pruning mechanism, the present invention dynamically screens redundant information processing units and only retains the key processing units that contribute to the task objective, thereby significantly reducing the computational overhead while ensuring the model performance. On the other hand, a twin neural network is a neural network with a dual-branch structure, where the two branches share parameters and are commonly used to learn the similarity between samples. In the brain age prediction task, the twin neural network can effectively characterize the structural or functional similarity between the input rs-fMRI images and guide the model to learn the age expression of deep features by comparing the labels of similar samples. Compared with the model structure of direct regression, the method based on similarity matching is more suitable for brain imaging data with small samples and significant individual differences. Therefore, combining the pruned Transformer structure with the twin neural network can not only efficiently extract the discriminative key features in rs-fMRI but also achieve robust brain age prediction through similarity measurement, with stronger generalization ability and clinical application potential. Based on this, the present invention proposes a brain age prediction method that combines a processing unit pruning mechanism and a twin neural network structure to improve the processing efficiency and prediction accuracy of rs-fMRI images and provide a better solution for brain health status assessment.

[0005] The inventive concept of the present invention is as follows: First, collect the brain resting-state functional magnetic resonance image data of the subjects, i.e., rs-fMRI images, and perform standardized preprocessing operations to generate three-dimensional structured image data, which is paired with the corresponding actual age information to form a complete sample set, and further divide it into a training set and a test set. Second, construct a pruning module to determine the retention or elimination of image patches according to the contribution degree of the image patches, construct a binary mask to represent the contribution of the image patches, and calculate the contribution degree of the image patches, using the L1 norm to measure the weight size of each patch. Then, construct a twin neural network model, using a Transformer encoder with a pruning module as a feature extractor to obtain two sets of feature vectors with the same scale. Third, design a joint loss function that comprehensively considers structural similarity and label similarity. The joint loss function includes a contrast loss function module Li to measure the Euclidean distance of samples in the feature space and an expected value loss function L(w,t).

[0006] To achieve the above inventive objective, the technical solution adopted by the present invention is specifically as follows: A brain age prediction method based on a twin pruning attention neural network, comprising the following steps:

[0007] S1: Collect the resting-state functional magnetic resonance imaging of the subjects, that is, rs-fMRI image data, to form an original data set, including image data and corresponding brain age prediction labels, and perform preprocessing operations such as slice timing correction, head motion correction, and spatial registration on the rs-fMRI image data to generate structured three-dimensional image data, and pair it with the corresponding actual age information to form a sample set, and further divide it into a training sample set and a test sample set at a ratio of 7:3;

[0008] S2: Construct a pruning module to select which image patches should be retained and which can be pruned according to the importance of the patches. First, construct a binary mask to represent the contribution of the image patches, then calculate the contribution degree of the image patches, and use the L1 norm to measure the weight size of each patch;

[0009] S3: Construct a siamese neural network model, using a Transformer encoder with a pruning module as a feature extractor. Its inputs are known and unknown rs-fMRI image data respectively. After being extracted by the same feature extractor, two groups of feature vectors with the same scale are obtained, and then the similarity of the two groups of vectors is measured by calculating the loss function;

[0010] S4: Design a joint loss function to comprehensively consider the structural similarity and label similarity. The joint loss function includes a contrastive loss function module L i , to measure the Euclidean distance of the samples in the feature space, and an expected value loss function L(w,t);

[0011] S5: After the model training is completed, input the test set samples into the trained siamese network structure one by one for prediction analysis. For each test sample to be predicted, the system will calculate its similarity with all samples in the training set in the feature space, and take the simplest average of the three most similar samples to obtain the predicted brain age of each test data sample.

[0012] Further, the specific steps of step S2 are as follows:

[0013] Step S2.1: First, define the importance of each image patch. For each image patch, calculate its contribution degree, that is, through the formula:

[0014]

[0015] where C1 represents the pruning function, M is a diagonal matrix composed of 0 and 1, 1 represents that the image patch is retained, 0 represents that the image patch is pruned, D is the data set, B() is the result of the pruned function passed to the encoder, Z is the input feature vector, W is the weight matrix, is the Hadamard product;

[0016] Step S2.2: Then further reduce redundancy and use regularization to constrain the structural complexity of pruning, that is, through the formula:

[0017]

[0018] where C2 represents the regularization function, l represents the pruning layer index, M l represents the binary mask used for pruning in the l-th layer, L is the total number of layers of the encoder, and ||·|| represents the L1 norm calculation;

[0019] Step S2.3: Finally, combine the pruning function and the regularization function to obtain the final pruning model formula:

[0020]

[0021] where λ is the hyperparameter of regularization, used to control the trade-off between computational cost and loss performance, r l represents the final pruning ratio;

[0022] Furthermore, the specific steps of step S3 are as follows:

[0023] Step S3.1: Assume that x is the sequence of processing units of the input rs-fMRI image. First, obtain the standard Transformer input mapping through the formula:

[0024] q, k, v = W q x, W k x, W v x(4)

[0025] where x ∈ R N×D , N represents the length of the sequence, that is, the number of processing units, D represents the dimension of each processing unit, q, k, and v are the query vector, key vector, and value vector respectively; W q , W k , W v are the corresponding learnable weight matrices;

[0026] Step S3.2: Then calculate the attention weights, that is, through the formula:

[0027]

[0028] where qk T is the dot product similarity of the query and the key, divided by is for scaling to prevent gradient explosion, s is a proportional attention vector, often used to improve the attention mechanism, Ψ is the Softmax operation, which normalizes the attention scores for all processing units, and finally the standard attention weight matrix A ∈ R N×N is obtained through the formula:

[0029] x = Av (6)

[0030] Apply the weight attention A to the value vector to obtain a weighted output, and the representations of each processing unit are fused after being simply weighted by all processing units;

[0031] Step S3.3: Then perform the pruning operation. First, calculate the number of processing units to be retained, that is, through the formula:

[0032] K = N - (r l ×N) (7)

[0033] where N is the total number of current input processing units, r l represents the current pruning ratio, l represents the pruning stage number, K represents the number of processing units to be retained, that is, subtract the number of processing units to be pruned from the original processing units;

[0034] Step S3.4: Calculate the importance of each processing unit to the processing unit, that is, through the formula:

[0035]

[0036] where represents the attention score of the processing unit to all processing units in the h-th attention head, that is, take the average of all attention heads. The processing unit represents the classification head, h represents the index of the attention head, c represents the index related to the calculation of the classification head, and A c,: ∈R N is the final importance score of each processing unit; to avoid the processing unit from being pruned, set it to infinity to make its attention score to itself the largest, ensuring that this processing unit will definitely be retained during pruning;

[0037] Step S3.5: Then sort the importance scores of the processing units obtained in the previous step from high to low and return the sorting index, that is, through the formula:

[0038] index = argsort(A c,: ) (9)

[0039] where argsort is the sorting function, index is the finally returned index value, and then select the top k indices with the highest processing units, that is, the most important k processing units, that is, through the formula:

[0040] Source index = index[...,:k] (10)

[0041] where Source indexDenote the top k processing units with the highest contribution degrees selected;

[0042] Step S3.6: Finally, extract the top k most important processing units corresponding in the original processing unit sequence according to the selected indexes, that is, through the formula:

[0043] x prune = gather(x,source index ) (11)

[0044] where x prune is the processed unit sequence after pruning, and gather is an index extraction function that can extract specified data according to indexes;

[0045] Furthermore, the specific steps of the said Step S4 are as follows:

[0046] Step S4.1: Let G W (x prune ) and G W (y prune ) be the feature vectors extracted by the feature extractor with a pruning module from the known brain age and the unknown brain age respectively, and calculate the Euclidean distance between the two, that is, through the formula:

[0047] D w = || G W (x prune ) - G W (y prune )|| (12)

[0048] where D w represents the Euclidean distance between samples and is used to initially measure the similarity of the feature vectors of two samples. x prune and y prune are the rs-fMRI images of the unknown and known age labels respectively, and ||·|| represents calculating the Euclidean distance;

[0049] Step S4.2: Take D w as the input and define a contrast loss function as shown in the following formula:

[0050] L i = (1 - ζ) D 2 w + ζ(max(0, m - D 2 w )) (13)

[0051] where L i is the contrast loss function, m is the margin interval used to control the separation degree of negative sample pairs, and ζ is the defined similarity degree coefficient, that is, when x prune belongs to yprune When it is the case, ζ is 0; otherwise, it is 1.

[0052] Step S4.3: Obtain the expected loss through the Euclidean distance and the contrast loss function, that is, through the formula:

[0053] L(w,τ) = E xi,τ [||G w (τ(x prune )) - μx prune ||] (14)

[0054] where L(w,τ) is the expected loss, which represents the training loss under the current network parameters w and the augmentation process τ. The goal is to minimize this loss function. E xi,τ is the expectation symbol, G w (τ(x prune )) is the feature vector after being extracted and augmented by the feature extractor, μ is the mapping of the target category, which is set according to the actual label;

[0055] Step S4.4: First, initialize w0 and μ0, where w0 and μ0 are the initial values of the weight and the learning rate respectively;

[0056] Step S4.5: Then update the parameters by minimizing the expected loss and the contrast loss function, and update them through the gradient descent direction propagation method. Among them, w and μ are updated through the formulas:

[0057]

[0058] are updated, where θ is the parameter to be optimized, that is, w or μ, η is the learning rate, is the gradient of the loss function with respect to θ. When updating and iterating w and μ, μ and w are kept as invariant respectively, and finally the optimal network model is obtained.

[0059] Compared with the prior art, the beneficial effects of the present invention are as follows: By integrating the pruning mechanism and the siamese neural network structure, while significantly reducing the computational cost, the accuracy, generalization ability and clinical interpretability of brain age prediction are improved. The specific reasons are as follows:

[0060] 1. The present invention combines the advantages of Transformer in sequence modeling with the ability of the siamese neural network in similarity learning, and introduces a processing unit pruning mechanism to reduce redundant information and computational cost, and significantly improves the computational efficiency while maintaining the model accuracy. This method is applicable to the brain age estimation task of large-scale neuroimaging data, provides technical support for individual brain health assessment, neurodegenerative disease screening, etc., and has broad clinical application prospects.

[0061] 2. The present invention improves the accuracy of feature extraction: The present invention uses Transformer as the backbone encoder, which has good global modeling ability and can capture long-range dependence features in rs-fMRI images. It has stronger extraction ability than traditional convolutional networks and is suitable for high-dimensional and unstructured brain imaging data.

[0062] 3. The present invention significantly reduces the computational cost: By introducing a processing unit pruning module into the Transformer structure, information redundancy or processing units with little contribution to brain age prediction are dynamically screened, effectively reducing the computational amount in each layer of the Transformer, reducing the dependence on hardware resources, and improving the deployment efficiency and inference speed of the model.

[0063] 4. The present invention enhances the generalization ability of the model by introducing a siamese network structure: The present invention embeds the pruned Transformer into the siamese neural network framework to achieve feature contrast learning with shared parameters in the dual path. Compared with traditional direct regression prediction models, this structure can learn the relative differences between samples more fully and has stronger individual adaptability and generalization performance.

[0064] 5. The present invention constructs a joint loss function based on similarity and label difference: By designing a joint loss function that fuses contrast loss and expected value loss, the model can accurately align age semantics in the feature space, improving the stability and accuracy of brain age prediction.

[0065] 6. The present invention enhances the interpretability of the model in clinical application scenarios: The present invention estimates the brain age based on the label average of similar samples during the prediction stage, and the output result is more interpretable, which is convenient for doctors to understand the results and make clinical auxiliary judgments, improving the trust and application value of the system. BRIEF DESCRIPTION OF THE DRAWINGS

[0066] The drawings are used to provide a further understanding of the present invention, and constitute a part of the specification. They are used to explain the present invention together with the embodiments of the present invention, and do not constitute a limitation to the present invention.

[0067] Figure 1 It is the overall framework diagram of the brain age prediction method based on the siamese pruning attention neural network of the present invention.

[0068] Figure 2 It is the structural diagram of the similarity siamese convolutional neural network model of the brain age prediction method based on the siamese pruning attention neural network of the present invention.

[0069] Figure 3 It is the structural diagram of the pruning module of the brain age prediction method based on the siamese pruning attention neural network of the present invention. DETAILED DESCRIPTION OF THE EMBODIMENTS

[0070] To make the objectives, technical solutions and advantages of the present invention more clear and understandable, the present invention will be further described in detail below with reference to the accompanying drawings and embodiments. Of course, the specific embodiments described herein are only used to explain the present invention and are not used to limit the present invention.

[0071] Embodiment 1

[0072] See Figures 1 to 3 , this embodiment provides its technical solution as a brain age prediction method based on a twin pruning attention neural network. Taking a pair of pictures x1, y1 selected from the dataset as an example, from the original data image to the final predicted result, the following steps are included:

[0073] S1: Collect the resting-state functional magnetic resonance imaging (rs-fMRI) image data of the subject to form an original dataset, including image data and corresponding brain age prediction labels, and perform preprocessing operations such as slice timing correction, head motion correction, and spatial registration on the rs-fMRI image data to generate structured three-dimensional image data, and pair it with the corresponding actual age information to form a sample set, and further divide it into a training sample set and a test sample set in a ratio of 7:3;

[0074] S2: Construct a pruning module to select which image blocks should be retained and which can be pruned according to the importance of the blocks. First, construct a binary mask to represent the contribution of the image blocks, then calculate the contribution degree of the image blocks, and use the L1 norm to measure the weight size of each block;

[0075] S3: Construct a twin neural network model, using a Transformer encoder with a pruning module as a feature extractor. Its inputs are known and unknown rs-fMRI image data respectively. After being extracted by the same feature extractor, two sets of feature vectors with the same scale are obtained, and then the similarity of the two sets of vectors is measured by calculating the loss function;

[0076] S4: Design a combined loss function to comprehensively consider the structural similarity and label similarity. The combined loss function includes a contrast loss function module L i , to measure the Euclidean distance of the samples in the feature space, and the expected value loss function L(w,t);

[0077] S5: After the model training is completed, input the test set samples one by one into the trained twin network structure for prediction analysis. For each test sample to be predicted, the system will calculate its similarity with all samples in the training set in the feature space, and take the simplest average of the three most similar samples to obtain the predicted brain age of each test data sample.

[0078] As the brain age prediction method based on the twin pruning attention neural network provided in this embodiment, the specific steps of step S2 are as follows:

[0079] Step S2.1: First, define the importance of each image patch. For each image patch, calculate its contribution degree, that is, through the formula:

[0080]

[0081] where C1 represents the pruning function, M is a diagonal matrix composed of 0 and 1, 1 indicates that the image patch is retained, 0 indicates that the image patch is pruned, and here M is:

[0082]

[0083] D is the dataset, B() is the result passed to the encoder by the pruned function, Z is the input feature vector, and here x1 and y1 are taken

[0084]

[0085] W is the weight matrix, and here it is

[0086]

[0087] is the Hadamard product, and the finally obtained C1 is 0.30;

[0088] Step S2.2: Then further reduce redundancy and use regularization to constrain the structural complexity of pruning, that is, through the formula:

[0089]

[0090] where C2 represents the regularization function, l represents the pruning layer index, M l represents the binary mask used for pruning in the l-th layer, and here it is

[0091]

[0092] L is the total number of layers of the encoder, which is 3 here, ||·|| represents the L1 norm calculation, and the finally calculated C2 is 4;

[0093] Step S2.3: Finally, combine the pruning function and the regularization function to obtain the final pruning model formula:

[0094]

[0095] where λ is the hyperparameter of regularization, used to control the trade-off between computational cost and loss performance, and here λ is 0.05, r l represents the final pruning ratio, which is 0.5 here;

[0096] As the brain age prediction method based on the twin pruning attention neural network provided in this embodiment, the specific steps of step S3 are as follows:

[0097] Step S3.1: The input image data x1 are all matrices of shape 128×64. First, a standard Transformer input mapping is obtained through the formula:

[0098] q, k, v = W q x1, W k x1, W v x1 (9)

[0099] where x1 ∈ R 128×64 , 128 represents the length of the sequence, that is, the number of processing units, 64 represents the dimension of each processing unit, and W q , W k , W v are the corresponding learnable weight matrices, which are respectively:

[0100]

[0101] q, k, and v are the query vector, key vector, and value vector respectively, and here they are respectively:

[0102]

[0103] Step S3.2: Then calculate the attention weights, that is, through the formula:

[0104]

[0105] where qk T is the dot product similarity between the query and the key, divided by is to do scaling to prevent gradient explosion, s is a proportional attention vector, which is often used to improve the attention mechanism, and Ψ is the Softmax operation, which normalizes the attention scores for all processing units. Finally, the obtained is the standard attention weight matrix A ∈ R N×N , through the formula:

[0106] x1 = Av (17)

[0107] Apply the weighted attention A to the value vector to obtain the weighted output. The representations of each processing unit are simply weighted by all processing units and then fused together. The finally obtained x1 is a 128×64 vector. At the same time, the same operation is also performed on y1. The finally obtained x1 and y1 are respectively:

[0108]

[0109] Step S3.3: Then perform the pruning operation. First, calculate the number of processing units to be retained, that is, through the formula:

[0110] K = N - (r l × N) (20)

[0111] where N is the total number of current input processing units, here N = 128, r l represents the current pruning ratio, here taken as 0.5, l represents the pruning stage number, K represents the number of processing units to be retained, that is, 64, that is, subtract the number of processing units to be pruned from the original processing units;

[0112] Step S3.4: Calculate the importance of each processing unit to the processing unit, that is, through the formula:

[0113]

[0114] where represents the attention score of the processing unit to all processing units in the h-th attention head, that is, take the average of all attention heads. The processing unit represents the classification head, h represents the index of the attention head, c represents the index related to the calculation of the classification head, A c,: ∈ R N is the final importance score of each processing unit; to avoid the processing unit from being pruned, set it to infinity to make its attention score to itself the largest, ensuring that this processing unit will definitely be retained during pruning;

[0115] Step S3.5: Then sort the importance scores of the processing units obtained in the previous step from high to low and return the sorting index, that is, through the formula:

[0116] index = argsort(A c,: ) (22)

[0117] where argsort is the sorting function, index is the finally returned index value, and then select the top 64 indices with the highest values of the processing units, that is, the most important 64 processing units, that is, through the formula:

[0118] Source index = index[...,:k] (23)

[0119] where Source index represents the 64 processing units with the highest contribution selected;

[0120] Step S3.6: Finally, extract the top 64 most important processing units corresponding to the original processing unit sequence according to the selected indices, that is, through the formula:

[0121] x prune = gather(x, source index )(24)

[0122] where x prune is the sequence of processed units after pruning. The gather function is an index extraction function that can extract specified data according to the index. The finally extracted vector here is 64×64;

[0123] As the brain age prediction method based on the twin pruning attention neural network provided in this embodiment, the specific steps of step S4 are as follows:

[0124] Step S4.1: Let G W (x prune ) and G W (y prune ) be the feature vectors extracted by the feature extractor with a pruning module from the known brain age and the unknown brain age respectively. Here, G W (x prune ) and G W (y prune ) are respectively:

[0125]

[0126] Calculate the Euclidean distance between the two, that is, through the formula:

[0127] D w = ||G W (x prune ) - G W (y prune )||(27)

[0128] where D w represents the Euclidean distance between samples and is used to initially measure the similarity of the feature vectors of two samples. Here, D w is 0.803, x prune and y prune are the rs-fMRI images of the unknown and known age labels respectively, and ||·|| represents calculating the Euclidean distance;

[0129] Step S4.2: Using D w as the input, define the contrast loss function as the following formula:

[0130] L i = (1 - ζ)D 2 w + ζ(max(0, m - D 2 w ))(28)

[0131] Among them, L i is the contrast loss function, m is the margin, which is used to control the separation degree of negative sample pairs, ζ is the defined similarity coefficient, that is, when x prune belongs to y prune , ζ is 0, otherwise it is 1. Here, x prune belongs to y prune , so ζ is 0, and thus we get:

[0132] L i = D 2 w (29)

[0133] Here, the calculation result of L i is 0.645;

[0134] Step S4.3: Obtain the expected loss through the Euclidean distance and the contrast loss function, that is, through the formula:

[0135] L(w,τ) = E xi,τ [||G w (τ(x prune )) - μx prune ||] (30)

[0136] Among them, L(w,τ) is the expected loss, which represents the training loss under the current network parameters w and the enhancement process τ. The goal is to minimize this loss function. E xi,τ is the expectation symbol, G w (τ(x prune )) is the feature vector after being extracted and enhanced by the feature extractor, μ is the mapping of the target category, which is set according to the actual label;

[0137] Step S4.4: First, initialize w0 and μ0, where w0 and μ0 are the initial values of the weight and the learning rate respectively. Here, let their initial values be 0.1 and 0.01 respectively;

[0138] Step S4.5: Then update the parameters by minimizing the expected loss and the contrast loss function, and update through the gradient descent direction propagation method. Among them, w and μ are updated through the formulas:

[0139]

[0140] are updated, where θ is the parameter to be optimized, that is, w or μ, η is the learning rate, is the gradient of the loss function with respect to θ. When updating and iterating w and μ, μ and w are kept as invariants respectively, and finally the best network model is obtained.

[0141] Example 2

[0142] Referring to Embodiment 1, in this embodiment, we adopted the parameters and results obtained in Embodiment 1 and further conducted a comparative analysis with existing traditional methods. By systematically comparing the performance of multiple models under different evaluation metrics, the experimental results clearly show that this embodiment is superior to the traditional method in terms of performance and has more significant advantages.

[0143] 1. Traditional model

[0144] Referring to relevant domestic and foreign research, the following comparative models were selected in this embodiment:

[0145] 1) ViT model: The ViT (Vision Transformer) model is a method that transforms the image recognition task into a sequence modeling problem. It first divides the entire image into image patches of a fixed size (usually 16×16 pixels), then flattens and converts these image patches into vectors as inputs similar to text sequences. Subsequently, these vectors are fed into a standard Transformer network for processing. Through the multi-head self-attention mechanism, ViT can effectively capture long-range dependencies and global features in the image.

[0146] 2) Swin Transformer model: The Swin Transformer (Shifted Window Transformer) model is a method that improves the traditional Transformer structure through a local window self-attention mechanism. It first divides the image into multiple small windows, independently calculates the self-attention within each window, and then establishes cross-region connections between windows through a sliding window strategy. Such a design not only maintains the excellent long-range modeling ability of the Transformer but also effectively reduces the computational complexity. Swin Transformer can adapt to image features of different scales and performs excellently in tasks such as image classification, detection, and segmentation.

[0147] 3) DenseNet model: The DenseNet (Densely Connected Convolutional Networks) model is a convolutional neural network that strengthens feature transmission through a dense connection mechanism. It directly connects the output of each layer to the inputs of all subsequent layers, achieving efficient reuse of features and unobstructed transmission of gradients, thereby improving the stability of model training and the utilization efficiency of parameters. DenseNet performs well in medical image analysis, especially in small-sample learning scenarios.

[0148] 4) Siamese CNN Model: Siamese CNN (Siamese Convolutional Neural Network) is a method for feature extraction and similarity measurement through a dual-branch convolutional structure with parallel shared weights. It encodes the features of a pair of input images separately and learns the similarity or difference between samples by calculating the distance between the features. Siamese CNN is widely used in tasks such as face recognition, fingerprint matching, and medical image pairing, and can effectively handle small-sample problems.

[0149] 2. Comparison Metrics

[0150] In this embodiment, corresponding evaluation metrics are designed for classification tasks and prediction tasks respectively. In the classification task, we use Matthews Correlation Coefficient (MCC), Receiver Operating Characteristic curve (ROC), True Negative Rate (TNR), True Positive Rate (TPR), and Accuracy (ACC) to comprehensively evaluate the classification performance of the model at different levels. In the prediction task, the Concordance Index (c-Index) is selected as the main evaluation criterion to measure the accuracy of the model in predicting continuous or ranking-related outputs. Through the comprehensive analysis of multi-dimensional metrics, the actual performance of the model in different tasks can be more comprehensively reflected.

[0151] MCC: It is a metric for comprehensively evaluating the performance of a binary classification model, considering the balance of true positives, false positives, true negatives, and false negatives. The value of MCC ranges from -1 to 1, where 1 indicates completely correct, 0 indicates random prediction, and -1 indicates completely wrong. MCC is particularly suitable for cases of class imbalance and can provide a single evaluation criterion for each model, avoiding focusing only on the errors of a certain class. However, MCC does not specifically distinguish the specific impacts of the positive and negative classes. It is an overall performance evaluation metric and may therefore ignore the importance of certain specific classes (such as the negative class).

[0152] ROC: It is used to evaluate the performance of a classification model. By plotting the relationship between the True Positive Rate (TPR) and the False Positive Rate (FPR), it helps analyze the performance of the model at different thresholds. The Area Under the Curve (AUC) is used as the evaluation criterion. The closer the AUC value is to 1, the stronger the classification ability of the model. The ROC curve is commonly used to evaluate the overall performance of a classification model, especially in cases of class imbalance, as it can comprehensively display the advantages and disadvantages of the model. However, when the data is severely imbalanced, the ROC curve may overestimate the model's ability to identify the positive class because it considers the performance of both the negative and positive classes, and the negative class samples may occupy a large proportion.

[0153] TNR: It refers to the proportion of samples that are actually negative and are successfully predicted as negative by the model. This metric is used to evaluate the model's ability to identify negative samples, especially in tasks that have high requirements for negative class discrimination, such as anomaly detection or exclusion diagnosis. However, it only reflects the model's performance in dealing with negative classes and cannot reflect the effect of positive class recognition. TPR: It refers to the proportion of samples that are actually positive and are successfully identified as positive by the model. This metric is mainly used to evaluate the model's ability to identify positive classes, especially in application scenarios with high requirements for positive class detection, such as disease screening or security detection. However, it only focuses on the recognition effect of positive classes and cannot reflect the model's ability to distinguish negative samples. Therefore, it is usually necessary to combine with other metrics when comprehensively evaluating the model's performance.

[0154] ACC: It refers to the proportion between the number of samples correctly classified by the model and the total number of samples, and is used to measure the overall classification accuracy of the model. As an intuitive and easy-to-understand evaluation metric, accuracy has good reference value in scenarios where the class distribution is relatively balanced. However, when faced with a dataset with severely imbalanced classes, accuracy may not be able to truly reflect the model's performance on minority class samples.

[0155] c-Index: It is an evaluation metric commonly used in the field of survival analysis, and is used to measure the accuracy of the model in predicting the survival time or risk order of samples. It reflects the model's ability to distinguish high-risk and low-risk individuals by comparing the consistency between the predicted rankings and the actual rankings of all sample pairs. This metric is particularly suitable for evaluating the performance in survival time prediction and risk scoring tasks, and can effectively reflect the reliability of the model in the ranking prediction scenario.

[0156] Each metric has its unique role and limitations in different application scenarios. In this embodiment, a comprehensive evaluation method is adopted to objectively compare this embodiment with other models, aiming to highlight the superiority of this embodiment.

[0157] 3. Comparison Results

[0158] As can be seen from the charts and data, this embodiment performs excellently in all evaluation metrics and is significantly better than other comparison models. First, in terms of MCC, this embodiment achieved a high score of 86.90%. In contrast, the ViT model and the Swin Transformer model were 72.13% and 75.60% respectively, the DenseNet model was 60.05%, and the Siamese CNN model was 67.80%. This reflects that this embodiment has stronger reliability in the accuracy of positive class prediction.

[0159] This embodiment also demonstrates excellent performance in terms of the ROC score, reaching 88.00%. As an indicator that combines precision and recall, the ROC score can comprehensively measure the classification performance of the model. In contrast, the ROC scores of the ViT model, the Swin Transformer model, the DenseNet model, and the Siamese CNN model are 69.65%, 74.72%, 60.09%, and 64.97% respectively. This embodiment is significantly superior in terms of the stability and accuracy of positive class recognition.

[0160] In terms of TNR, this embodiment also performs very well, reaching 89.33%, which is significantly better than the 72.78% of the Swin Transformer model and the 68.91% of the ViT model, while the DenseNet model and the Siamese CNN model are 61.87% and 66.42% respectively. This indicates that this embodiment has stronger accuracy in identifying negative class samples and can effectively avoid misjudgments. In terms of TPR, this embodiment also performs outstandingly, reaching 86.75%. In contrast, the Swin Transformer model and the ViT model are 74.85% and 67.42% respectively, while the DenseNet model and the Siamese CNN model are 60.14% and 62.31% respectively. This shows that this embodiment has higher sensitivity in identifying positive class samples and is better at detecting key positive individuals.

[0161] In terms of the accuracy ACC, this embodiment reaches 88.21%, which is significantly higher than the 70.85% of the ViT model, the 74.12% of the Swin Transformer model, the 63.30% of the DenseNet model, and the 65.27% of the Siamese CNN model. This fully demonstrates that this embodiment has an obvious advantage in overall classification accuracy and can more reliably distinguish positive class and negative class samples.

[0162] Although in terms of the c-Index, the 79.45% of the DenseNet model is slightly higher than some models, this embodiment ranks first with a score of 86.77%, showing stronger ranking prediction ability. This indicates that in tasks involving survival time prediction or risk ranking, this embodiment also has excellent adaptability.

[0163] Generally speaking, this embodiment far exceeds other comparison models in multiple core evaluation indicators, fully demonstrating its superiority and stability in classification tasks, and having stronger practical value and promotion potential.

[0164] Table 1 Comparison table of this embodiment and other models

[0165]

[0166] The above are only the preferred embodiments of the present invention and are not intended to limit the present invention. Any modifications, equivalent replacements, improvements, etc. made within the spirit and principle of the present invention shall be included within the protection scope of the present invention.

Claims

1. A brain age prediction method based on a twin pruning attention neural network, characterized in that, It includes the following steps: S1: Collect the resting-state functional magnetic resonance imaging (rs-fMRI) of the subjects, that is, the rs-fMRI image data, to form an original data set, including the image data and the corresponding brain age prediction labels. Perform slice timing correction, head motion correction, and spatial registration preprocessing operations on the rs-fMRI image data to generate structured three-dimensional image data, and pair it with the corresponding actual age information to form a sample set, which is divided into a training sample set and a test sample set at a ratio of 7:3; S2: Construct a pruning module to select which image patches are retained and which are pruned according to the importance of the patches. First, construct a binary mask to represent the contribution of the image patches, then calculate the contribution degree of the image patches, and use the L1 norm to measure the weight size of each patch; S3: Construct a siamese neural network model, using a Transformer encoder with a pruning module as the feature extractor. Its inputs are the known and unknown rs-fMRI image data respectively. After being extracted by the same feature extractor, two sets of feature vectors with the same scale are obtained, and then the similarity of the two sets of vectors is measured by calculating the loss function; S4: Design a combined loss function that comprehensively considers structural similarity and label similarity. The combined loss function includes a contrastive loss function module L i , to measure the Euclidean distance of samples in the feature space, and an expected value loss function L(w, t); S5: After the model training is completed, input the test set samples one by one into the trained siamese network structure for prediction analysis. For each test sample to be predicted, the system will calculate its similarity with all samples in the training set in the feature space, take the average of the three most similar samples, and obtain the predicted brain age of each test data sample.

2. The brain age prediction model of the brain age prediction method based on the twin pruning attention neural network according to claim 1, wherein The steps of step S2 are as follows: Step S2.1: First, define the importance of each image patch. For each image patch, calculate its contribution degree, that is, through the formula: Among them, C1 represents the pruning function, M is a diagonal matrix composed of 0 and 1, where 1 indicates that the image block is retained and 0 indicates that the image block is pruned, D is the data set, B() is the result passed to the encoder by the pruned function, Z is the input feature vector, and W is the weight matrix. is the Hadamard product; Step S2.2: Then further reduce redundancy and use regularization to constrain the structural complexity of pruning, that is, through the formula: where C2 represents the regularization function, l represents the pruning layer index, M l represents the binary mask for pruning in the l-th layer, L is the total number of layers of the encoder, and ||·|| represents the L1 norm calculation; Step S2.3: Finally, combine the pruning function and the regularization function to obtain the final pruning model formula: where λ is the hyperparameter of regularization, which is used to control the trade-off between computational cost and loss performance, r l represents the final pruning ratio.

3. The brain age prediction model of the brain age prediction method based on the twin pruning attention neural network according to claim 1, wherein The steps of step S3 are as follows: Step S3.1: Assume that x is the sequence of processing units of the input rs-fMRI image. First, obtain the standard Transformer input mapping, through the formula: q, k, v = W q x, W k x, W v x (4) where \(x\in R\) N×D , \(N\) represents the length of the sequence, i.e., the number of processing units, \(D\) represents the dimension of each processing unit, and \(q\), \(k\), \(v\) are the query vector, key vector, and value vector respectively; \(W\) q , \(W\) k , \(W\) v are the corresponding learnable weight matrices; Step S3.2: Then calculate the attention weights, that is, through the formula: where qk T is the dot product similarity between the query and the key, divided by is for scaling to prevent gradient explosion, s is a scale attention vector, often used to improve the attention mechanism, Ψ is the Softmax operation, which normalizes the attention scores for all processing units, and finally the standard attention weight matrix A ∈ R N×N is obtained through the formula: x = Av (6) Apply the weighted attention A to the value vector to obtain the weighted output, and the representation of each processing unit is fused after being simply weighted by all processing units; Step S3.3: Then perform the pruning operation. First, calculate the number of processing units to be retained, that is, through the formula: K = N - (r l × N) (7) where N is the total number of processing units for the current input, r l represents the current pruning ratio, l represents the number of pruning stages, and K represents the number of processing units to be retained, that is, subtracting the number of processing units to be pruned from the original processing units; Step S3.4: Calculate the importance of each processing unit to the processing unit, that is, through the formula: Among them represents the attention score of the processing unit to all processing units in the h-th attention head, that is, taking the average of all attention heads. The processing unit represents the classification head, h represents the index of the attention head, c represents the index related to the calculation of the classification head, and A c,: ∈R N is the importance score of each processing unit finally; Step S3.5: Then sort the importance scores of the processing units obtained in the previous step from high to low and return the sorting index, that is, through the formula: index = argsort(A c,: ) (9) where argsort is the sorting function and index is the finally returned index value. Then select the top k processing units with the highest index, that is, the k most important processing units, that is, through the formula: Source index = index[...,:k] (10) Among them, Source index represents the k processing units with the highest contribution selected; Step S3.6: Finally, extract the corresponding top k most important processing units from the original processing unit sequence according to the selected index, that is, through the formula: x prune = gather(x,source index ) (11) where x prune is the sequence of processing units after pruning. The index extraction function during gathering extracts specified data according to the index.

4. The brain age prediction model of the brain age prediction method based on the twin pruning attention neural network according to claim 1, wherein The steps of step S4 are as follows: Step S4.1: Let G W (x prune ) and G W (y prune ) be the feature vectors extracted from the known brain age and the unknown brain age through the feature extractor with a pruning module respectively, and calculate the Euclidean distance between them, that is, through the formula: D w = || G W (x prune ) - G W (y prune ) || (12) Among them, D w represents the Euclidean distance between samples, which is used to initially measure the similarity of the feature vectors of two samples. x prune and y prune are the rs-fMRI images of unknown and known age labels respectively, and ||·|| represents the calculation of the Euclidean distance; Step S4.2: Using D w as the input, define the contrastive loss function, as shown in the following formula: L i = (1 - ζ) D 2 w + ζ(max(0, m - D 2 w )) (13) where L i is the contrastive loss function, m is the margin, which is used to control the separation degree of negative sample pairs, and ζ is the defined similarity coefficient, that is, when x prune belongs to y prune , ζ is 0, otherwise it is 1; Step S4.3: Obtain the expected loss by passing through the Euclidean distance and the contrastive loss function, that is, through the formula: L(w,τ) = E xi,τ [||G w (τ(x prune )) - μx prune ||] (14) Among them, \(L(w, \tau)\) is the expected loss, which represents the training loss under the current network parameters \(w\) and the augmentation process \(\tau\). The goal is to minimize this loss function, and \(E\) xi,τ is the expectation symbol, and \(G\) w (\(\tau(x\) prune )) is the feature vector after being extracted and augmented by the feature extractor. \(\mu\) is the mapping of the target category, which is set according to the actual label; Step S4.4: First, initialize w0 and μ0, where w0 and μ0 are the initial values of the weight and the learning rate respectively; Step S4.5: Then update the parameters by minimizing the expected loss and the contrastive loss function, and update them through the gradient descent direction propagation method, where w and μ are respectively through the formulas: Update, where θ is the parameter to be optimized, namely w or μ, η is the learning rate, is the gradient of the loss function with respect to θ. When updating and iterating w and μ, μ and w are kept as invariants respectively, and finally the best network model is obtained.

Citation Information

Cited By

  • Multi-head attention model conversion method and device, storage medium and electronic equipment

    CN121009924A