Multi-modal multi-task fusion attention model for Alzheimer's disease classification

Through the multimodal and multi-task fusion attention model, the problem of insufficient fusion of imaging and clinical data was solved, high-precision Alzheimer's disease classification was achieved, and the accuracy and interpretability of diagnosis were improved.

CN120708930APending Publication Date: 2025-09-26CHANGCHUN UNIV OF SCI & TECH
View PDF 0 Cites 1 Cited by

Patent Information

Application Number
CN202510807990.8
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-06-17
Publication Date
2025-09-26

AI Technical Summary

Technical Problem

In existing Alzheimer's disease analysis methods, image data and clinical data are not fully utilized, and the multimodal data fusion effect is poor, resulting in unsatisfactory classification accuracy.

Method used

A multimodal multi-task fusion attention model is adopted, including a feature extraction module, a clinically guided attention module, a proxy attention module, a dynamic gated feature fusion module and a multi-task loss function module. Through multi-plane and multi-scale feature extraction, clinical data guidance, dynamic feature mixing and multi-task collaborative training, efficient fusion of imaging and clinical data is achieved.

Benefits of technology

The accuracy of Alzheimer's disease classification has been improved, with an AD/CN binary classification accuracy of 100%, an AD/MCI binary classification accuracy of 95.16%, and an AD/CN/MCI triple classification accuracy of 92.47%. It also provides a high-precision and explainable multimodal fusion framework.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120708930A_ABST
    Figure CN120708930A_ABST
Patent Text Reader

Abstract

The invention relates to a multi-modal multi-task fusion attention model for Alzheimer's disease classification, and belongs to the technical field of attention model construction, and the multi-modal multi-task fusion attention model comprises a feature extraction module, a clinical guidance attention module, a proxy attention mechanism module, a dynamic gating feature fusion module and a multi-task loss function module which are connected in sequence. The agent attention mechanism module is also connected with the deep cross network module and the contrast learning projection module. According to the model, efficient fusion of image-clinical data is realized through a clinical attention guiding module, a dynamic gating feature fusion module and a multi-task comparison decoupling three-stage collaborative mechanism. The multi-modal multi-task fusion attention model is trained and verified based on 912 multi-center data of an ADNI database, the AD / CN binary classification accuracy is 100%, the AD / MCI binary classification accuracy is 95.16%, the AD / CN / MCI ternary classification accuracy is 92.47%, and the advanced level is achieved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technical field of attention model construction, and in particular to a multimodal multi-task fusion attention model for Alzheimer's disease classification. Background Art

[0002] Alzheimer's disease (AD) is a neurodegenerative disorder that causes neuronal cell death and brain tissue atrophy, leading to a significant decline in cognitive function and ability to function independently. The disease is characterized by irreversible progression, and currently, no cure has been found. Only through early identification and appropriate medical intervention can the progression of the disease be slowed and the patient's quality of life improved. Faced with an increasingly aging population, the number of people suffering from AD continues to climb, making it one of the fastest-growing mortality rates worldwide. It is projected that by 2050, global cases of dementia will reach 153 million, and in the United States, AD is already among the top ten leading causes of death. China has one of the highest numbers of Alzheimer's patients in the world, with an estimated 9.83 million patients in 2023. This large population places tremendous pressure on family care, social security, and medical resources.

[0003] With the rapid development of contemporary medical technology, particularly breakthroughs in neuroimaging and artificial intelligence algorithms, computer-assisted diagnosis (CAD) has opened up new avenues for precision medicine in Alzheimer's disease. Deep learning (DL), with its remarkable performance in image recognition, has made analyzing and processing MRI imaging data a key approach to improving the accuracy of Alzheimer's disease diagnoses. Compared to traditional machine learning methods, deep learning demonstrates superior performance, significantly improving the accuracy and efficiency of disease diagnosis while also enabling faster model training. Deep learning algorithms automatically perform feature selection during model construction and loss function optimization, eliminating the need to rely on the prior knowledge of domain experts.

[0004] Current Alzheimer's disease (AD) analysis methods based on medical imaging can be divided into four categories based on feature extraction strategies: Voxel-based methods: These methods directly associate disease-related microstructural features through voxel-by-voxel analysis and are widely used in AD diagnosis and MCI conversion prediction. However, this method has high computational complexity due to high-dimensional features and is difficult to effectively model local texture correlation (FURFaisal, G.-R. Kwon Automated detection of Alzheimer's disease and mild cognitive impairment using whole brain MRI IEEE Access, 10 (2022), pp. 65055-65066.; D. Jin, et al. Attention-based 3D convolutional network for Alzheimer's disease diagnosis and biomarkers exploration 2019 IEEE 16th International Symposium on Biomedical Imaging (ISBI 2019) (Apr. 2019), pp. 1047-1051). Slice-based methods: They extract two-dimensional slice features from three-dimensional images to reduce the number of network parameters and significantly alleviate the pressure on computing resources. However, the two-dimensional processing method breaks the spatial continuity between adjacent slices, resulting in partial loss of three-dimensional structural information and easily leading to data leakage problems (Kumar S. Sambath, Nandhini M., Initia Alzheimers Dis Neuroimaging Entropy slicing extraction and transfer le arning classification for early diagnosis of Alzheimer diseases with sMRI AC M Trans. Multimed. Comput. Commun. Appl., 17(2)(2021), 10.1145 / 3383749.;Uttam Khatri, Goo-Rak Kwon, Diagnosis of Alzheimer's disease via optimized lightweight convolution-attention and structural MRI, Computers in Biology and Medicine, Volume 171, 2024, 108116, ISSN 0010-4825.). Region-based (ROI-based) methods: Based on prior anatomical knowledge, they locate specific brain regions (such as the hippocampus and entorhinal cortex) and extract regional features, which can improve the specificity of detecting AD-related pathologies. Its limitation is that potential pathological areas outside the ROI may be ignored, reducing the sensitivity of whole-brain scale analysis (M. Liu, F. Li, H. Yan, K. Wang, Y. Ma, L. Shen, M. Xu, A multi-model deep convolutional neural network for automatic hippocampus segmentation and classification in Alzheimer's disease, Neuroimage (2020) 208.; M. A. Ebrahimighahnavieh, S. Luo, R. Chiong, Deep learning to detect Alzheimer's disease from neuroimaging: a systematic review, Comput. Methods Progr. Biomed. 187 (2020) 105242, https: / / doi.org / 10.1016 / j.cmpb.2019.105242.). 3D-patch based method: It extracts 3D contextual features through a local cube sampling strategy, enhancing sensitivity to subtle pathological changes while preserving spatial dependencies. The block analysis mechanism of this method can effectively alleviate the computational burden brought by high-dimensional data (Shangran Qiu, Prajakta S Joshi, Matthew I Miller, Chonghua Xue, Xiao Zhou, Cody Karjadi, Gary H. Chan, Anant S. Joshi, Brigid Dwyer, Shuhan Zhu, Michelle Kaku, Yan Zhou, Yazan J. Alderazi, Arun Swaminathan, Sachin Kedar, Marie-Hélène Saint-Hilaire, Sanford H. Auerbach, Jing Yuan, E. Alton Sartor, Rhoda Au, Vijaya B. Kolachalama, Development and validation of an interp retable deep learning framework for Alzheimer's disease classification, Brain, Volume 143, Issue 6, June 2020, Pages 1920–1933, https: / / doi.org / 10.1093 / brain / awaa137.-Shangran Qiu, Prajakta S Joshi, Matthew I Miller, Chonghua Xue, Xiao Zhou, Cody Karjadi, Gary H Chang, Anant S Joshi, Brigid Dwyer, Shuh An Zhu, Michelle Kaku, Yan Zhou, Yazan J Alderazi, Arun Swaminathan, Sac hin Kedar, Marie-Helene Saint-Hilaire,Sanford H Auerbach,JingYuan,E Alto n Sartor,Rhoda Au,Vijaya B Kolachalama,Development and validation of an interpretabledeep learning framework for Alzheimer's diseaseclassification,Bra in,Volume 143,Issue 6,June 2020,Pages 1920–1933, https: / / doi.org / 10.1093 / brain / awaa137.).

[0005] Compared with single-modality approaches, multimodal analysis can provide complementary information and enhance AD ​​detection. To overcome the limitations of traditional single-modality analysis, current research paradigms are gradually shifting towards multimodal artificial intelligence technologies, integrating multiple sources of information, such as neuroimaging, cognitive assessment scales, and biomarkers, through heterogeneous data fusion strategies. This cross-modal fusion framework, through a multi-dimensional feature complementarity mechanism, can systematically analyze the correlation between pathophysiological changes and cognitive impairment in Alzheimer's patients, thereby improving the efficacy of disease diagnosis.In the critical stage of early disease identification - the process of mild cognitive impairment (MCI) progressing to Alzheimer's disease, multimodal analysis technology provides important technical support for accurately identifying early biomarkers and implementing precise interventions (S. Qiu, M.I. Miller, P.S. Joshi, J.C. Lee, C. Xue, Y. Ni, Y. Wang, I. De Anda-Duran, P.H. Wang, J.A. Cramer, B.C. Dwyer, H. Hao, M.C. Kaku, S. Kedar, P.H. Lee, A.Z. Mian, D.L. Murman, S. O'Shea, A.B. Paul, M.H. Saint-Hilaire, S.E. Alton, A.R. Saxena, L.C. Shih, J.E. Mall, M.J. Smith, A. Swaminathan, C.E. Takahashi, O. Taraschenko, H. You, J. Yuan, Y. Zhou, S. Zhu, M.L. Alosco, J. Mez, T.D. Stein, K.L. Poston, R. Au, V.B. Kolachalama, Multimodal deep learning for Alzheimer's disease dementia assessment, Nat. Commun. 13 (1) (2022) 3404.; Fei Liu, Huabin Wang, Shiuan-Ni Liang, Zhe Ji n, Shicheng Wei, Xuejun Li, MPS-FFA: A multiplane and multiscale feature fus ion attention network for Alzheimer's disease prediction with structural MRI, Computers in Biology and Medicine, Volume 157, 2023, 106790, ISSN 0010-4825, https: / / doi.org / 10.1016 / j.compbiomed.2023.106790.).

[0006] While Alzheimer's disease analysis methods based on medical imaging have advanced, they still face challenges in applying multimodal data. Current methods generally suffer from insufficient data fusion depth, particularly the lack of effective fusion strategies when integrating imaging information with clinical features, which directly affects the accuracy of classification diagnosis. Summary of the Invention

[0007] In response to the technical problems existing in existing Alzheimer's disease analysis methods - insufficient utilization of image data and clinical data, poor multimodal data fusion effect and unsatisfactory classification accuracy, this paper proposes a multimodal multi-task fusion attention model for Alzheimer's disease classification.

[0008] In order to solve the above technical problems, the technical solutions of the present invention are as follows:

[0009] A multimodal, multi-task fusion attention model for Alzheimer's disease classification, comprising a feature extraction module, a clinically guided attention module, a proxy attention module, a dynamic gated feature fusion module, and a multi-task loss function module, connected in sequence; the proxy attention module is further connected to a deep cross network module and a contrastive learning projection module;

[0010] The feature extraction module is used to construct multi-plane and multi-scale feature extraction of images, including: two three-dimensional convolution layers, a main path, an auxiliary path and a fusion layer;

[0011] The main path includes four layer encapsulation modules connected in sequence, the middle two layer encapsulation modules include a residual basic block, a multi-plane multi-scale feature extraction module and a downsampling module, and the remaining two layer encapsulation modules include a residual basic block and a downsampling module;

[0012] The auxiliary path includes two combined downsampling modules, each of which includes a residual basis block, a multi-plane attention fusion module and a downsampling module;

[0013] Among them, the residual basic block is used to construct the basic unit of the deep network; the multi-plane multi-scale feature extraction module realizes multi-dimensional collaborative enhancement of complex features through the series connection of multi-plane attention fusion and multi-scale attention fusion; the downsampling module is used to reduce the spatial dimension of the feature map; the multi-plane attention fusion module is used to realize multi-dimensional feature interaction;

[0014] The 3D convolutional layer is used to extract basic image features. The main path extracts deep features layer by layer and retains the intermediate output. The auxiliary path performs multi-plane attention downsampling on the intermediate features of the main path. Finally, the fusion layer splices the terminal features of the main path and the aggregated features of the auxiliary path along the channel dimension to perform multi-scale feature compression and output an enhanced feature map.

[0015] The clinically guided attention module is used to implement clinical data guided visual feature enhancement;

[0016] The proxy attention module is used to realize feature space interaction through a learnable proxy vector;

[0017] The dynamic gated feature fusion module is used to achieve dynamic mixing of image and clinical features;

[0018] The multi-task loss function module is used to realize unified loss calculation of multi-objective collaborative training;

[0019] The deep cross network module is used to achieve regression prediction of high-order interaction of features;

[0020] The contrastive learning projection module is used to construct a cross-modal contrastive learning space.

[0021] In the above technical solution, among the two three-dimensional convolutional layers, the first three-dimensional convolutional layer uses a 4x4x4 convolution kernel, which expands the input 1 channel to 32 channels and performs downsampling with a stride of 2. The second three-dimensional convolutional layer increases the number of channels to 64 through a 3x3x3 convolution kernel.

[0022] In the above technical solution, the residual basic block includes a sequence of two convolutional layers self.conv: the first layer expands the input channel to twice the number of output channels through a 1x1x1 convolution kernel, and the second layer compresses the number of channels to the target output dimension through a 3x3x3 convolution kernel.

[0023] In the above technical solution, the multi-plane and multi-scale feature extraction module realizes multi-dimensional collaborative enhancement of complex features through the series connection of multi-plane attention fusion and multi-scale attention fusion;

[0024] Among them, the multi-plane attention fusion module creates three cross-dimensional convolution branches. The first branch directly performs conventional three-dimensional convolution and spatial attention calculation. The second branch adjusts the input to the (B, C, Y, Z, X) format through dimensional permutation, and restores the original dimensional arrangement after convolution and spatial attention processing. The third branch adjusts the input to the (B, C, X, Z, Y) format through dimensional permutation and also restores the original dimension after processing. The outputs of the three branches are spliced ​​along the channel dimension, and feature fusion is performed through the channel compression layer. The final output is a feature map that maintains the same spatial dimension and number of channels as the input tensor;

[0025] The multi-scale attention fusion module creates a dynamic number of parallel convolution branches. Each branch executes the convolution sequence in sequence to obtain the feature map, generates a weight mask through the corresponding local attention module, and multiplies the original feature map with the attention map element by element to achieve feature reweighting; collects the feature maps processed by all branches and splices them along the channel dimension to form a fused feature; finally, features are aggregated through the channel compression layer, and the output and input tensors maintain the enhanced features of the same spatial dimension and number of channels.

[0026] In the above technical solution, the clinically guided attention module constructs a clinical feature projection layer self.clinical_proj to map low-dimensional clinical data to a high-dimensional image feature space, and uses an attention generator self.attn_conv composed of two 3D convolutions. The first layer compresses the number of channels to 1 / 8, and the second layer generates a single-channel attention map; in the forward propagation method forward(), the clinical features are linearly transformed and expanded into a five-dimensional tensor that matches the image features. Modal interaction is achieved through feature addition, and finally spatial attention weights are generated and multiplied element-by-element with the original image features to output an enhanced feature map guided by clinical information.

[0027] In the above technical solution, the proxy attention module constructs a feature transformation network self.fc to realize inter-channel information fusion, and adopts deep separable convolution self.dwc to enhance spatial feature locality; in the forward propagation method forward(), the spatial correlation between the proxy vector and the input feature is first calculated by scaling the dot product attention, and the proxy vector is used as the query vector to obtain the global information of the feature space; then the proxy vector information is broadcasted to the original feature space through the attention mechanism in reverse to realize proxy-guided feature reconstruction; finally, the feature map with dual spatial and channel enhancement is output through deep separable convolution and residual connection to maintain the consistency of input and output dimensions.

[0028] In the above technical solution, the dynamic gated feature fusion module learns the fusion weights through the gated network self.gate. The first-layer full connection reduces the dimension of the spliced ​​multimodal features to 32 dimensions, and the second layer is mapped to a single-dimensional gated value and constrained to the [0,1] interval through the Sigmoid function. In the forward propagation, the clinical features are expanded and replicated according to the number of channels required, and the gated values ​​are used to perform weighted summation on the image features and the extended clinical features to achieve dynamic interpolation and fusion of the feature space.

[0029] In the above technical solution, the deep cross network module integrates the cross network and deep network dual paths. The cross network performs explicit feature crossover through multi-layer linear transformation to capture the combination relationship between features; the deep network mines deep nonlinear implicit abstract features through a three-layer MLP containing batch normalization and GELU activation; finally, the explicit combination features of the cross network are spliced ​​with the implicit abstract features of the deep network, and the predicted value is output through the regression layer.

[0030] In the above technical solution, the contrastive learning projection module includes a parallel image projection head self.img_proj and a clinical projection head self.clinical_proj, which map heterogeneous features to a unified 128-dimensional space through an independent fully connected layer; in the forward propagation, L2 normalization is performed on the projection results of the two modalities respectively to ensure that the feature vectors are distributed on the unit hypersphere when the contrast loss is calculated, thereby enhancing the stability of the cross-modal similarity measurement.

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

[0032] The multimodal multitask fusion attention model (MMA) for Alzheimer's disease classification of the present invention achieves efficient fusion of imaging and clinical data through a three-stage collaborative mechanism. Specifically:

[0033] 1. Clinical Guided Attention Module: The Clinical Guided Attention module encodes MMSE / CDRSB scores into a three-dimensional spatial attention map, dynamically enhancing the image feature expression of AD-sensitive areas such as the hippocampus and temporal cortex;

[0034] 2. Dynamic Gated Feature Fusion Module: Adaptively balances the contribution of imaging and clinical features through a learnable gating unit (Dynamic Gated Fusion);

[0035] 3. Multi-task contrast decoupling: Construct regression tasks, contrastive learning tasks, and a multi-task loss function (classification: regression: contrast = 0.7:0.2:0.1) to improve feature discriminability;

[0036] 4. The model of the present invention was trained and verified based on 912 multicenter data from the ADNI database: the accuracy of AD / CN two-classification was 100% (AUC=1.00), the accuracy of AD / MCI two-classification was 95.16% (AUC=0.9938), and the accuracy of AD / CN / MCI three-classification was 92.47% (AUC=0.9882), reaching the advanced level.

[0037] 5. The model of this invention also captures pathological features through a multi-plane and multi-scale feature extraction module. The dynamic proxy attention mechanism models cross-modal feature interactions to alleviate the distribution differences between imaging and clinical data. Combined with the DCNRegressor, it explicitly models the nonlinear association between imaging and scores, achieving fine-grained alignment of pathological and clinical features. This provides a highly accurate and interpretable multimodal fusion framework for AD diagnosis. BRIEF DESCRIPTION OF THE DRAWINGS

[0038] The present invention will be further described in detail below with reference to the accompanying drawings and specific embodiments.

[0039] Figure 1This is a data preprocessing flowchart of the multimodal multi-task fusion attention model for Alzheimer's disease classification of the present invention.

[0040] Figure 2 This is a comparison chart before and after data processing.

[0041] Figure 3 This is a structural diagram of the multimodal multi-task fusion attention model for Alzheimer's disease classification of the present invention. DETAILED DESCRIPTION

[0042] The multimodal multi-task fusion attention model for Alzheimer's disease classification of the present invention is described in detail below with reference to the accompanying drawings.

[0043] The multimodal multitask fusion attention model (MMA) for Alzheimer's disease classification of the present invention is shown in the following figure: Figure 3, including: including: Feature Extraction Module, Clinical Guided Attention Module, Agent Attention Module, Dynamic Gating Feature Fusion Module and Multitask Contrastive Loss Module connected in sequence; Agent Attention Module is also connected to Deep Cross Network Module respectively. network, DCN) and a contrastive learning projection module (ContrastiveProjection); the feature extraction module is used to construct a multi-plane and multi-scale feature extraction of an image, including: two three-dimensional convolutional layers (Conv3d), a main path, an auxiliary path and a fusion layer; the main path includes four layer encapsulation modules (LayerWarp) connected in sequence, the two middle layer encapsulation modules include a residual basic block (BasicBlock), a multi-plane and multi-scale feature extraction module and a downsampling module (DownSample), and the remaining two layer encapsulation modules include a residual basic block and a downsampling module; the auxiliary path includes two combined downsampling modules (LayerDownSample), each combined downsampling module (LayerDownSample) includes a residual basic block, a multi-plane attention fusion module and a downsampling module; wherein the residual basic block is used to construct the basic unit of the deep network; the multi-plane and multi-scale feature extraction module is connected through a multi-plane attention fusion module. The series connection of intention fusion and multi-scale attention fusion realizes multi-dimensional collaborative enhancement of complex features; the downsampling module is used to reduce the spatial dimension of the feature map; the multi-plane attention fusion module is used to realize multi-dimensional feature interaction; the three-dimensional convolution layer is used to extract basic image features, the main path extracts deep features layer by layer and retains the intermediate output, the auxiliary path performs multi-plane attention downsampling on the intermediate features of the main path, and finally the terminal features of the main path and the aggregated features of the auxiliary path are spliced ​​along the channel dimension through the fusion layer to perform multi-scale feature compression and output the enhanced feature map; the clinical guided attention module is used to realize visual feature enhancement guided by clinical data; the proxy attention module is used to realize feature space interaction through learnable proxy vectors; the dynamic gated feature fusion module is used to realize dynamic mixing of image and clinical features; the multi-task loss function module is used to realize unified loss calculation for multi-target collaborative training; the deep cross network module is used to realize regression prediction of high-order feature interactions; the contrastive learning projection module is used to construct a cross-modal contrastive learning space.

[0044] The following is a more specific and detailed introduction to each module of the multimodal multi-task fusion attention model for Alzheimer's disease classification of the present invention.

[0045] The feature extraction module is used to construct multi-plane and multi-scale feature extraction of images. The module sets a four-stage processing configuration in the initialization method __init__(), and each stage is defined by a tuple (number of blocks, attention enable flag, input channel, output channel); creates an initial fast downsampling layer self.first_layer, which contains two three-dimensional convolution layers ConvLayer (Conv3d): the first layer uses a 4x4x4 convolution kernel to expand the input 1 channel to 32 channels and performs downsampling with a step size of 2, and the second layer uses a 3x3x3 convolution kernel to increase the number of channels to 64; constructs the main path self.layers and the auxiliary path self.layers_m, and dynamically generates two processing modules based on the stage configuration: the main path uses the layer encapsulation module (LayerWarp) to implement feature extraction (the parameters include the current stage index and attention flag), and the auxiliary path is realized by group The combined downsampling module (LayerDownSample) implements multi-plane attention downsampling (taking the multi-plane attention fusion module as the attention module parameter); defines the final feature fusion layer self.reduce_ch (i.e., the multi-plane attention fusion module), which contains three cascaded convolutional layers (3x3x3 convolution reduces the dimension to 512→128 channels, and 1x1x1 convolution compresses to 64 channels); in the forward propagation method forward(self,x), the input tensor first passes through the initial convolution layer to obtain the basic features, and then undergoes four-stage hierarchical processing: the main path extracts deep features layer by layer and retains the intermediate output, and the auxiliary path performs multi-plane attention downsampling on the intermediate features of the main path; finally, the terminal features of the main path and the aggregated features of the auxiliary path are spliced ​​along the channel dimension, and multi-scale feature compression is performed through the fusion layer (i.e., the multi-plane attention fusion module), and the output is an enhanced feature map with 64 channels.The layer encapsulation module (LayerWarp) is used to implement feature extraction of multiple basic block stacks. The module accepts the input channel number ch_in, output channel number ch_out, stacking count, attention enable flag att, layer number layer_no and debug flag is_print parameters in the initialization method __init__(); dynamically creates a specified number of basic block sequences: when the att flag is activated, the branch configuration is selected according to the ratio of the current block sequence number to the total number of blocks (the first half of the blocks use four branches (1, 3, 5, 7), and the second half of the blocks use two branches (1, 3)), and embeds the attention mechanism through AttBlock; each basic block uses the residual basic block (BasicBlock) structure to keep the number of input and output channels consistent; the downsampling module (DownSample) is integrated at the end to achieve spatial dimension compression; in the forward propagation method forward(self, x), the input tensor is enhanced through all basic blocks in sequence. When the debug mode is turned on, the intermediate features are saved in numpy format to the specified path; finally, the feature map after channel expansion is output through the downsampling module (DownSample). The multi-plane and multi-scale feature extraction module achieves multi-dimensional collaborative enhancement of complex features through the cascade of multi-plane attention fusion and multi-scale attention fusion. Through the interaction of features at different levels (planes) and scales (receptive fields), the model's ability to capture target details and global context is improved.Among them, the multi-plane attention fusion module is used to realize multi-dimensional feature interaction. The module accepts the input channel number in_channel parameter in the initialization method __init__(); creates three cross-dimensional convolution branches (corresponding to different dimensional arrangements), each branch contains two 3D convolution layers, batch normalization layers and ReLU activation functions: the first layer uses a 3x3x3 convolution kernel to double the number of channels and maintain the spatial dimension, and the second layer compresses the number of channels back to the original value through a 1x1x1 convolution kernel; defines a channel compression layer self.reduce_channel, and uses a 1x1x1 convolution kernel to reduce the tripled number of channels after the three branches are merged to the original input channel number; equips each branch with an independent space The spatial attention modules self.sp_xyz, self.sp_yzx and self.sp_xzy; in the forward propagation method forward(self,x), the input tensor x is processed in multiple paths: the first branch directly performs conventional three-dimensional convolution and spatial attention calculation; the second branch adjusts the input to the (B, C, Y, Z, X) format through dimension permutation, and restores the original dimension arrangement after convolution and spatial attention processing; the third branch adjusts the input to the (B, C, X, Z, Y) format through dimension permutation, and also restores the original dimension after processing; the outputs of the three branches are spliced ​​along the channel dimension, and feature fusion is performed through the channel compression layer, and the final output is a feature map that maintains the same spatial dimension and number of channels as the input tensor.The multi-scale attention fusion module is used to realize multi-scale feature fusion. The module accepts the input channel number in_channel and the branch configuration parameter branch (default is (1,3,5)) in the initialization method __init__(); creates a dynamic number of parallel convolution branches, each branch corresponds to a different convolution kernel size: contains two 3D convolution layer structures, the first layer uses the convolution kernel of the specified size (1x1x1, 3x3x3 or 5x5x5) for feature transformation, keeps the spatial dimension unchanged by automatically calculating the padding value, and expands the number of channels to twice; the second layer compresses the number of channels back to the original value through 1x1x1 convolution; creates a local attention submodule for each branch, learns the attention weight through 1x1x1 kernel convolution, and uses The sigmoid function generates an attention map in the range of [0,1]; defines a channel compression layer self.reduce_channel, and reduces the number of feature channels after merging multiple branches to the number of original input channels through a 1x1x1 convolution kernel; in the forward propagation method forward(self,x), the input tensor x is processed in parallel with multiple branches: each branch executes the convolution sequence in turn to obtain a feature map, generates a weight mask through the corresponding local attention module, and multiplies the original feature map and the attention map element-by-element to achieve feature reweighting; collects the feature maps processed by all branches and splices them along the channel dimension to form a fused feature; finally, performs feature aggregation through a channel compression layer, and outputs enhanced features that maintain the same spatial dimension and number of channels as the input tensor. The SpatialAttention module involved is used to enhance the spatial information of three-dimensional feature maps. The module's core components are defined in the initialization method __init__(): a 3D convolutional layer self.conv with 2 input channels and 1 output channel, using a 3x3x3 convolution kernel and setting padding of 1 to maintain the spatial dimension. Normalized attention weights are generated using the Sigmoid activation function self.sigmoid. In the forward propagation method forward(self,x), multi-step spatial feature aggregation is performed: first, the mean avgout (preserving the spatial dimension) and maximum maxout of the input tensor are calculated along the channel dimension, and the two are concatenated along the channel dimension to form a two-channel feature map. The spatial attention weight distribution is learned through the convolutional layer, and the weights are normalized to the range [0,1] using the Sigmoid function. Finally, the generated attention map is element-wise multiplied with the original input tensor to achieve feature reweighting in the spatial dimension, outputting enhanced features with the same shape as the input tensor.

[0046] A combined downsampling module (LayerDownSample) is used to integrate feature processing and spatial compression. The module receives the input channel number ch_in, the output channel number ch_out, and the optional attention module att parameter in the initialization method __init__(); constructs a serialized processing flow: the first residual basic block (BasicBlock) performs feature enhancement through the optional attention mechanism, followed by the downsampling module (DownSample) to perform channel expansion and spatial downsampling; the forward propagation directly performs feature enhancement and downsampling operations sequentially, and the output channel number is the compressed feature of ch_out. The residual basic block (BasicBlock) is used to construct the basic unit of a deep network. The module accepts the input channel number ch_in, the output channel number ch_out, and the optional attention module att_net parameters in the initialization method __init__(). It adopts a bottleneck structure design to create a sequence self.conv containing two convolutional layers: the first layer expands the input channels to twice the number of output channels through a 1x1x1 convolution kernel, and the second layer compresses the number of channels to the target output dimension through a 3x3x3 convolution kernel, both of which maintain the same spatial size. The optional integrated attention module self.att_net implements feature reweighting. In the forward propagation method forward(), the input tensor is first feature-transformed through the convolution sequence and residually connected with the original input. If an attention module is present, the convolution output is enhanced with attention and then residually fused with the original input to form a dual residual structure, effectively alleviating the gradient vanishing problem while enhancing feature expression capabilities.

[0047] The agent attention module (AgentAttention) is used to realize feature space interaction through learnable agent vectors. The module defines n_agents trainable agent vectors self.agents and bias parameters in the initialization method __init__(); constructs a feature transformation network self.fc to realize inter-channel information fusion, and uses deep separable convolution self.dwc to enhance spatial feature locality; in the forward propagation method forward(), first calculates the spatial association between the agent vector and the input feature through scaled dot product attention, and uses the agent vector as the query vector to obtain the global information of the feature space; then, the agent vector information is broadcasted to the original feature space through the attention mechanism in reverse, realizing agent-guided feature reconstruction; finally, the spatial-channel dual enhanced feature map is output through deep separable convolution and residual connection to maintain the consistency of input and output dimensions.

[0048] The clinical guided attention module (ClinicalGuidedAttention) is used to achieve visual feature enhancement guided by clinical data. The module receives the image feature channel number in_channels and clinical feature dimension clinical_dim parameters in the initialization method __init__(); constructs a clinical feature projection layer self.clinical_proj to map low-dimensional clinical data to a high-dimensional image feature space, and uses an attention generator self.attn_conv composed of two 3D convolutions. The first layer compresses the number of channels to 1 / 8, and the second layer generates a single-channel attention map; in the forward propagation method forward(), the clinical features are linearly transformed and expanded into a five-dimensional tensor matching the image features. Modal interaction is achieved by feature addition, and finally spatial attention weights are generated and multiplied element-by-element with the original image features to output an enhanced feature map guided by clinical information.

[0049] The dynamic gating feature fusion module (Dynamic gating feature fusion), also known as the classification task, is used to achieve dynamic mixing of image and clinical features. The module learns fusion weights through the gated network self.gate: the first layer of full connection reduces the dimensionality of the spliced ​​multimodal features to 32 dimensions, and the second layer maps them to single-dimensional gating values ​​and constrains them in the interval [0,1] through the Sigmoid function; in the forward propagation, the clinical features are expanded and replicated according to the number of channels required (for example, 64 channels require 2-dimensional clinical features to be repeated 32 times), and the gating values ​​are used to perform weighted summation of image features and expanded clinical features to achieve dynamic interpolation and fusion of the feature space.

[0050] The ContrastiveProjection module, also known as the contrastive learning task, is used to construct a cross-modal contrastive learning space. The module contains a parallel image projection head (self.img_proj) and a clinical projection head (self.clinical_proj), which map heterogeneous features to a unified 128-dimensional space through independent fully connected layers. During forward propagation, L2 normalization is performed on the projection results of the two modalities to ensure that the feature vectors are distributed on the unit hypersphere when calculating the contrast loss, thereby enhancing the stability of the cross-modal similarity measurement.

[0051] The deep cross network module (DCN), also known as the regression task, is used to implement regression prediction of high-order feature interactions. The module integrates the cross network and deep network dual paths: the cross network performs explicit feature crossover through multi-layer linear transformations (the calculation formula for each layer is x⊙(Wx+b)+x) to capture the combination relationship between features; the deep network mines deep nonlinear features through a three-layer MLP with batch normalization and GELU activation; finally, the explicit combination features of the cross network are spliced ​​with the implicit abstract features of the deep network, and the predicted values ​​are output through the regression layer, which has the ability to model feature interactions and extract deep semantics.

[0052] The multi-task loss function module (MultitaskContrastiveLoss) is used to implement unified loss calculation for multi-objective collaborative training, weighting the aforementioned classification tasks, contrastive learning tasks, and regression tasks to obtain results. The multi-task loss function module accepts the mode flag mode (only supports fixed weight mode), the initial weight dictionary init_weights (default {'cls': 0.7, 'reg': 0.2, 'contrast': 0.1}), and the contrast temperature parameter temperature (default 0.07) in the initialization method __init__(); constructs the basic loss function: the classification task, that is, the dynamic gated feature fusion module, uses the cross entropy loss self.cls_criterion; the regression task, that is, the deep cross network module, uses the mean squared error loss self.reg_criterion; the contrast task, that is, the contrastive learning projection module, uses the improved NT-Xent loss through the custom contrastive_loss method. The forward propagation method forward() performs a multi-stage loss calculation: 1) Calculates the classification loss cls_loss, the sum of the MMSE and CDRSB regression losses reg_loss, and the cross-modal contrast loss contrast_loss; 2) Linearly weights the three losses according to preset weights to generate a total loss total_loss; 3) Returns a dictionary containing the individual loss terms, weight parameters, and the total loss. Each loss term is blocked from backpropagation via the detach() method to prevent interference with weight assignment. The core process of the contrast loss calculation: L2 normalization is performed on the image and clinical embedding vectors, the cross-modal similarity matrix is ​​calculated using the Einstein sum convention, diagonally aligned pseudo-labels are constructed after temperature scaling, and finally cross-modal instance discrimination is achieved using the cross-entropy function.

[0053] In summary, the multimodal multitask fusion attention model for Alzheimer's disease classification of the present invention can be defined as a multimodal multitask fusion attention model definition unit, which is used to realize the collaborative diagnosis of medical images and clinical data. The model constructs a multi-component collaborative system in the initialization method __init__(): the feature extraction module self.backbone is integrated, and the multi-level spatial-semantic feature extraction of the image is realized through the BackBone class; the dual-path attention mechanism self.clinical_guided_att and self.agent_att are deployed, and the spatial attention enhancement guided by the clinical scale is realized through the ClinicalGuidedAttention class, and the AgentAttention class is realized. Channel interaction guided by proxy vector; configure the two-branch deep cross regressor self.mmse_regressor and self.cdrsb_regressor, use the DCNRegressor class for scale score prediction, support feature explicit cross and deep implicit modeling; build a dynamic gated fusion module self.fusion, and use the DynamicGatedFusion class to achieve parameterized mixing of global image features and clinical features; set the contrastive learning projection head self.contrast_proj, and generate the unit sphere embedding of cross-modal contrastive learning through the ContrastiveProjection class; finally, realize disease classification through the two-stage fully connected network self.classifier.

[0054] The multimodal multitask fusion attention model for Alzheimer's disease classification of the present invention may further include: a model operation subunit and an early stopping mechanism module (EarlyStopping).

[0055] The model execution subunit is used to define the forward function of the multimodal, multitask fusion attention model of the present invention. In the forward propagation method forward(), the input 3D medical image x and the clinical scales mmse and cdrsb are processed in multiple stages: 1) The image is passed through the backbone network (feature extraction module) to extract a 64-channel 4x5x4 spatial feature map; 2) The clinical scales are spliced ​​into a 2D feature vector, which is then reweighted by the clinical-guided attention module to achieve spatial feature reweighting, and then the channel interaction is enhanced by the proxy attention module; 3) The enhanced feature map is flattened into a 5120-dimensional vector and input into a deep cross-regressor to predict the scale score; 4) Global pooling is performed on the image features to obtain a 64-dimensional semantic vector, which is dynamically fused with the clinical features through a gating mechanism; 5) The flattened image features and clinical features are respectively projected into a 128-dimensional contrast space and L2 normalized; the final output is a dictionary containing the classification result, scale prediction, and contrast embedding, where the classification result is generated by a two-layer MLP that fused the features, and contrast learning returns independent embedding pairs of the image and clinical for loss calculation.

[0056] The early stopping mechanism module (EarlyStopping) is used to monitor the validation loss during the model training process of the present invention and prevent overfitting. The module accepts three parameters in the initialization method __init__(): patience represents the maximum number of consecutive epochs allowed without performance improvement (default is 5), min_delta defines the minimum threshold for loss improvement (default is 0.0), and save_path specifies the optimal model save path (default is 'best_model.pth'); the internal maintenance state variable best_loss records the best validation loss (initialized to positive infinity), counter records the number of consecutive times without improvement, and the early_stop flag indicates whether the stopping condition is triggered.

[0057] In the __call__() execution method, the current validation loss (current_loss) and model parameter update mechanism are used: if the current loss is lower than the historical best loss by more than the min_delta threshold, the best loss value is updated, the counter is reset, and the internal method _save_model is called to save the current optimal model. Otherwise, the no-improvement counter is incremented, and the early stopping flag is activated when the number of consecutive no-improvement times reaches the patience threshold. This module continuously tracks the dynamic changes in validation loss, enabling intelligent interruption control during training and persistent storage of the optimal model.

[0058] The data preprocessing flow chart of the multimodal multitask fusion attention model for Alzheimer's disease classification of the present invention can be found in Figure 1Data were retrieved from the Alzheimer's Disease Neuroimaging Initiative (ADNI) database. All sMRI scans were performed at 3T resolution, with T1-weighted images consisting of magnetization-prepared rapid-acquisition gradient-echo sequences. To prevent data leakage, only one structural MRI image was selected for each individual. A total of 1057 patients were included, including 290 with AD, 373 with CN, and 392 with MCI. To ensure data balance, 912 structural MRI images of patients with AD, CN, and MCI were randomly selected, including 290 with AD, 311 with CN, and 311 with MCI. Corresponding clinical data, namely the Mini-Mental State Examination (MMSE) and the Clinical Dementia Rating Scale (CDRSB), were also used. The MMSE is a standardized cognitive screening tool that assesses orientation, memory, and language skills for early identification of cognitive impairment. The CDRSB provides a multidimensional assessment based on the patient's daily life performance (such as social activities and self-care) and the degree of cognitive impairment, quantifying dementia stage progression. The two complement each other through quantitative indicators and behavioral dimensions, breaking through the single modality limitations of imaging / genetic data.

[0059] The data preprocessing process is based on the integration of SPM12 and CAT12 toolboxes on the Matlab platform. The core steps include:

[0060] 1. AC-PC joint calibration

[0061] The original image was loaded through the Display module of SPM12, and the anterior commissure (AC) and posterior commissure (PC) anatomical landmarks were located in the sagittal, coronal, and horizontal planes.

[0062] Adjust the three-dimensional coordinate system so that the AC-PC line is horizontally aligned, complete the head motion correction (eliminate the influence of the subject's body position offset), and save the redirected image.

[0063] 2. Skull removal and brain tissue segmentation

[0064] The Segment Data module of CAT12 was called to perform nonlinear registration and spatial normalization based on the MNI152 template.

[0065] A mixed Gaussian model was used to segment gray matter (mwp1), white matter (mwp2), and whole brain tissue (wm), and the skull and non-brain tissue were removed to generate a size-standardized (121×145×121 voxel) whole-brain image.

[0066] The preprocessing results were visualized and verified in the report folder, and finally a dataset containing whole-brain sMRI images was constructed for subsequent model training and verification.

[0067] The spatial resolution of the obtained image was adjusted to 113×137×113 after downsampling. Figure 2 .

[0068] The model of the present invention was trained and verified based on 912 multi-center data from the ADNI database: the accuracy of AD / CN two-classification was 100% (AUC=1.00), the accuracy of AD / MCI two-classification was 95.16% (AUC=0.9938), and the accuracy of AD / CN / MCI three-classification was 92.47% (AUC=0.9882), reaching an advanced level.

[0069] Under the same training settings, the present invention conducted an ablation experiment. By changing certain modules of the model while keeping other conditions unchanged, the model classification task performance was improved, proving the effectiveness of the module. See the table below. An early stopping mechanism module was set to stop training when the verification loss did not decrease within 15 rounds of model training.

[0070]

[0071] Obviously, the above embodiments are merely examples for clarity of explanation and are not intended to limit the implementation methods. Those skilled in the art will readily appreciate that other variations or modifications based on the above descriptions are possible. It is not necessary and impossible to enumerate all implementation methods here. Obvious variations or modifications arising therefrom remain within the scope of protection of the present invention.

Claims

1. A multimodal multitask fusion attention model for Alzheimer's disease classification, characterized by: It includes the following connected in sequence: feature extraction module, clinical guided attention module, proxy attention module, dynamic gated feature fusion module and multi-task loss function module; the proxy attention module is also connected to the deep cross network module and contrastive learning projection module respectively; The feature extraction module is used to construct multi-plane and multi-scale feature extraction of images, including: two three-dimensional convolution layers, a main path, an auxiliary path and a fusion layer; The main path includes four layer encapsulation modules connected in sequence, the middle two layer encapsulation modules include a residual basic block, a multi-plane multi-scale feature extraction module and a downsampling module, and the remaining two layer encapsulation modules include a residual basic block and a downsampling module; The auxiliary path includes two combined downsampling modules, each of which includes a residual basis block, a multi-plane attention fusion module and a downsampling module; Among them, the residual basic block is used to construct the basic unit of the deep network; the multi-plane multi-scale feature extraction module realizes multi-dimensional collaborative enhancement of complex features through the series connection of multi-plane attention fusion and multi-scale attention fusion; the downsampling module is used to reduce the spatial dimension of the feature map; the multi-plane attention fusion module is used to realize multi-dimensional feature interaction; The 3D convolutional layer is used to extract basic image features. The main path extracts deep features layer by layer and retains the intermediate output. The auxiliary path performs multi-plane attention downsampling on the intermediate features of the main path. Finally, the fusion layer splices the terminal features of the main path and the aggregated features of the auxiliary path along the channel dimension to perform multi-scale feature compression and output an enhanced feature map. The clinically guided attention module is used to implement clinical data guided visual feature enhancement; The proxy attention module is used to realize feature space interaction through a learnable proxy vector; The dynamic gated feature fusion module is used to achieve dynamic mixing of image and clinical features; The multi-task loss function module is used to realize unified loss calculation of multi-objective collaborative training; The deep cross network module is used to achieve regression prediction of high-order interaction of features; The contrastive learning projection module is used to construct a cross-modal contrastive learning space.

2. The multimodal multi-task fusion attention model according to claim 1, characterized in that In the two 3D convolutional layers, the first 3D convolutional layer uses a 4x4x4 convolution kernel, which expands the input 1 channel to 32 channels and performs downsampling with a stride of 2. The second 3D convolutional layer increases the number of channels to 64 through a 3x3x3 convolution kernel.

3. The multimodal multitask fusion attention model according to claim 1, characterized in that The residual basic block contains a sequence of two convolutional layers self.conv: the first layer expands the input channels to twice the number of output channels through a 1x1x1 convolution kernel, and the second layer compresses the number of channels to the target output dimension through a 3x3x3 convolution kernel.

4. The multimodal multitask fusion attention model according to claim 1, characterized in that The multi-plane and multi-scale feature extraction module achieves multi-dimensional collaborative enhancement of complex features through the series connection of multi-plane attention fusion and multi-scale attention fusion; Among them, the multi-plane attention fusion module creates three cross-dimensional convolution branches. The first branch directly performs conventional three-dimensional convolution and spatial attention calculation. The second branch adjusts the input to the (B, C, Y, Z, X) format through dimensional permutation, and restores the original dimensional arrangement after convolution and spatial attention processing. The third branch adjusts the input to the (B, C, X, Z, Y) format through dimensional permutation and also restores the original dimension after processing. The outputs of the three branches are spliced ​​along the channel dimension, and feature fusion is performed through the channel compression layer. The final output is a feature map that maintains the same spatial dimension and number of channels as the input tensor; The multi-scale attention fusion module creates a dynamic number of parallel convolution branches. Each branch executes the convolution sequence in sequence to obtain the feature map, generates a weight mask through the corresponding local attention module, and multiplies the original feature map with the attention map element by element to achieve feature reweighting; collects the feature maps processed by all branches and splices them along the channel dimension to form a fused feature; finally, features are aggregated through the channel compression layer, and the output and input tensors maintain the enhanced features of the same spatial dimension and number of channels.

5. The multimodal multitask fusion attention model according to claim 1, characterized in that The clinically guided attention module constructs a clinical feature projection layer self.clinical_proj to map low-dimensional clinical data to a high-dimensional image feature space, and uses an attention generator self.attn_conv composed of two 3D convolutions. The first layer compresses the number of channels to 1 / 8, and the second layer generates a single-channel attention map. In the forward propagation method forward(), the clinical features are linearly transformed and expanded into a five-dimensional tensor that matches the image features. Modal interaction is achieved through feature addition, and finally spatial attention weights are generated and multiplied element-by-element with the original image features to output an enhanced feature map guided by clinical information.

6. The multimodal multitask fusion attention model according to claim 1, characterized in that The proxy attention module constructs a feature transformation network self.fc to achieve inter-channel information fusion, and uses deep separable convolution self.dwc to enhance spatial feature locality. In the forward propagation method forward(), the spatial association between the proxy vector and the input feature is first calculated through scaled dot product attention, and the proxy vector is used as the query vector to obtain the global information of the feature space. Then, the proxy vector information is broadcasted to the original feature space through the attention mechanism to achieve proxy-guided feature reconstruction. Finally, the feature map with dual spatial and channel enhancement is output through deep separable convolution and residual connection to maintain the consistency of input and output dimensions.

7. The multimodal multitask fusion attention model according to claim 1, characterized in that The dynamic gated feature fusion module learns fusion weights through the gated network self.gate. The first layer is fully connected to reduce the dimension of the spliced ​​multimodal features to 32 dimensions. The second layer is mapped to a single-dimensional gated value and constrained to the interval [0, 1] through the Sigmoid function. In the forward propagation, the clinical features are expanded and replicated according to the number of channels required, and the gated value is used to perform weighted summation on the image features and the extended clinical features to achieve dynamic interpolation and fusion of the feature space.

8. The multimodal multitask fusion attention model according to claim 1, characterized in that The deep cross network module integrates the cross network and deep network dual paths. The cross network performs explicit feature crossover through multi-layer linear transformation to capture the combination relationship between features; the deep network mines deep nonlinear implicit abstract features through a three-layer MLP containing batch normalization and GELU activation; finally, the explicit features of the cross network are spliced ​​with the implicit abstract features of the deep network, and the predicted values ​​are output through the regression layer.

9. The multimodal multitask fusion attention model according to claim 1, characterized in that The contrastive learning projection module includes a parallel image projection head self.img_proj and a clinical projection head self.clinical_proj, which map heterogeneous features to a unified 128-dimensional space through independent fully connected layers. During the forward propagation, L2 normalization is performed on the projection results of the two modalities to ensure that the feature vectors are distributed on the unit hypersphere when calculating the contrast loss, thereby enhancing the stability of the cross-modal similarity measurement.

Citation Information

Cited By

  • Method for accurately predicting clinical prognosis of colorectal cancer patient

    CN121329977A