Model based on multi-modal cross attention and uncertainty integral gradient and application
By employing a multimodal cross-attention and uncertainty integral gradient model, the problems of difficulty in multimodal data fusion and weak model generalization ability are solved, achieving high-precision drug sensitivity prediction and target identification, and improving the interpretability and robustness of the model.
Patent Information
- Application Number
- CN202511416979.5
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-09-30
- Publication Date
- 2026-01-16
AI Technical Summary
Existing drug sensitivity prediction methods have difficulties in multimodal data fusion, have weak model generalization ability, lack interpretability, and are difficult to provide accurate and interpretable predictions under small sample conditions.
We employ a model based on multimodal cross-attention and uncertainty integral gradient. The cross-attention mechanism captures the correlation between modes, Monte Carlo dropout is combined to achieve uncertainty estimation, and the integral gradient algorithm is used to analyze feature importance, thus achieving a balance between high-precision prediction and biological interpretability.
It improves the efficiency of multimodal data fusion, enhances the interpretability and predictive performance of the model, and can provide high-precision drug sensitivity prediction and target identification under small sample conditions, making it suitable for real-world scenarios where clinical data quality varies.
Smart Images

Figure FT_1 
Figure SMS_1 
Figure SMS_2
Abstract
Description
TECHNICAL FIELD
[0001] The present application relates to the technical field of bioinformatics, in particular to a model based on multi-modal cross-attention and uncertainty integral gradient and application thereof. BACKGROUND
[0002] In the field of precision medicine for cancer, the accuracy of drug sensitivity prediction directly determines whether patients can benefit from targeted therapy or immunotherapy. In the past decade, with the public release of large cohorts such as TCGA and GDSC, researchers have tried to use classic algorithms such as random forest, support vector machine or gradient boosting tree to establish the mapping relationship between "gene features-drug response". These models can often achieve an AUC of about 0.7 on public benchmarks, but when entering real clinical scenarios, the prediction performance often drops to below 0.6. The fundamental reason is that the complexity of tumor molecular mechanisms far exceeds what can be described by a single modality of data: mRNA expression can only reflect abnormalities at the transcription level, miRNA regulatory networks can amplify or suppress signals rapidly after transcription, and the chemical characteristics of drugs themselves determine their binding kinetics with target proteins. The traditional approach simply concatenates or weighted averages the three types of features, resulting in not only the loss of high-order collaborative information between modalities, but also the amplification of overfitting risk in the high-dimensional small sample scenario. More difficult is that when the model gives a "sensitive" or "drug-resistant" conclusion, clinicians cannot know which key genes or drug substructures drive this judgment, nor can they assess the uncertainty of the prediction, leading to misjudgment of high-risk patients.
[0003] In recent years, the success of deep learning in multi-modal tasks such as images and texts has brought hope, but direct migration to drug sensitivity prediction still faces two major obstacles. First, the sample size of public datasets is generally only a few hundred, much lower than the scale of natural images which are often in the millions. Second, the mRNA, miRNA and drug features differ greatly in numerical distribution and physical meaning, and existing contrast learning or denoising autoencoders are mainly designed for homogeneous modalities, making it difficult to be directly applicable. Therefore, how to effectively integrate heterogeneous omics data under small sample conditions while providing interpretable and uncertainty-quantified predictions remains a technical problem to be solved. SUMMARY
[0004] To address the aforementioned issues, this invention provides a model based on multimodal cross-attention and uncertainty integral gradient. This model, constructed using multimodal cross-attention and uncertainty integral gradient, can solve problems such as difficulty in multimodal data fusion, weak model generalization ability, and lack of interpretability in existing drug sensitivity prediction methods. The model captures intermodal correlations through a cross-attention mechanism, combines Monte Carlo dropout (MC Dropout) to achieve uncertainty estimation, and utilizes the integral gradient algorithm to analyze feature importance, ultimately achieving a balance between high-precision prediction and biological interpretability.
[0005] This invention provides a model based on multimodal cross-attention and uncertainty integral gradient, including a prediction module and an uncertainty integral gradient module. The prediction module includes a cross-attention module and a multi-task learning module. The cross-attention module includes a linear projection layer and a cross-attention fusion module. The linear projection layer is used to input multimodal features and align the dimensions and distributions of different modal features. The cross-attention fusion module is used to realize cross-modal information interaction. The multi-task learning module is used to introduce uncertainty and output latent features. The uncertainty integral gradient module is used to quantify the uncertainty of the prediction results and analyze the contribution of each feature to the prediction results.
[0006] In one embodiment, the cross-attention fusion module includes a cross-attention mechanism and a fusion layer. The cross-attention mechanism is used to input the output data of the linear projection layer. The output data of the cross-attention mechanism is concatenated to obtain high-dimensional features, which are then input into the fusion layer. The fusion layer compresses and reconstructs the high-dimensional features through a fully connected network and nonlinear transformation, and outputs fused features. The multi-task learning module includes a variational autoencoder, which is used to encode the input fused features. The output data of the variational autoencoder is reparameterized, sampled from the latent distribution, and output latent features.
[0007] In one embodiment, the cross-attention mechanism includes a self-attention submodule and a cross-attention submodule. The self-attention submodule is used to perform self-attention calculation, and the cross-attention submodule is used to construct multiple sets of cross-modal attention. The feature concatenation includes integrating the output data of the self-attention submodule and the output data of the cross-attention submodule. In one embodiment, the uncertainty integral gradient module includes an MC Dropout uncertainty estimation module and an integral gradient module, wherein the MC Dropout uncertainty estimation module includes a forward function; The uncertainty integral gradient module is used to input the output data of the variational autoencoder, construct a baseline, input the forward function, calculate the gradient integral, perform smoothing, take the average, generate a feature importance matrix, perform dynamic weight balancing, select the top - k key features, and output the prediction result.
[0008] The present invention also provides a method for predicting drug sensitivity and / or target identification, comprising the following steps: inputting the sample data to be evaluated into the model and calculating the prediction result.
[0009] In one embodiment, the sample data includes: gene expression data and drug characteristics; The prediction results include drug sensitivity prediction results, target identification results, and / or drug property results.
[0010] In one embodiment, the gene expression data includes mRNA sequencing data and miRNA sequencing data; the drug characteristics include chemical structure, target information, and drug sensitivity tags.
[0011] The present invention also provides a system for predicting drug sensitivity and / or target identification, comprising: A data storage module for the sample data to be evaluated and the module itself; The data analysis module is used to perform analysis according to the method described; and The data display module is used to output and display the prediction results.
[0012] Compared with the prior art, the present invention has the following beneficial effects: This invention relates to a model and its application based on multimodal cross-attention and uncertainty integral gradient. The model is constructed based on multimodal cross-attention and uncertainty integral gradient, which can solve the problems of difficulty in multimodal data fusion, weak model generalization ability, and lack of interpretability in existing drug sensitivity prediction methods. The model captures the correlation between modes through the cross-attention mechanism, combines Monte Carlo dropout (MC Dropout) to achieve uncertainty estimation, and uses the integral gradient algorithm to analyze feature importance, ultimately achieving a unity of high-precision prediction and biological interpretability. Attached Figure Description
[0013] Figure 1 This is a flowchart of the workflow of the model of the present invention. Detailed Implementation
[0014] To facilitate understanding of the present invention, a more complete description will be given below with reference to the accompanying drawings. Preferred embodiments of the invention are shown in the drawings. However, the invention can be implemented in many different forms and is not limited to the embodiments described herein. Rather, these embodiments are provided to provide a thorough and complete understanding of the disclosure of the invention.
[0015] Unless otherwise defined, all technical and scientific terms used herein have the same meaning as commonly understood by one of ordinary skill in the art to which this invention pertains. The terminology used herein in the description of the invention is for the purpose of describing particular embodiments only and is not intended to be limiting of the invention. The term "and / or" as used herein includes any and all combinations of one or more of the associated listed items.
[0016] Unless otherwise specified, all reagents, materials, and equipment used in this embodiment are commercially available; unless otherwise specified, all test methods are conventional test methods in this field.
[0017] Example I. Overview of the overall framework of the model.
[0018] The core framework of this invention includes a cross-attention module (for multimodal feature encoding and cross-attention fusion), a multi-task learning module (for drug sensitivity prediction, target prediction, and feature reconstruction), and an uncertainty integral gradient module (for uncertainty quantification and feature importance analysis). It extracts key features for each modality through modality-specific encoding, utilizes a cross-attention mechanism to achieve cross-modal information interaction, optimizes model parameters using a multi-task loss function, and finally achieves prediction uncertainty estimation and key target identification through MCDropout and integral gradient algorithms.
[0019] 1. Construction of the cross-attention mechanism.
[0020] The cross-attention mechanism is a core module for achieving multimodal feature fusion. Its design goal is to capture the complex relationships between mRNA, miRNA, and drug characteristics, and to enhance the transmission of important information through dynamic weight allocation. The specific construction process is as follows: S1: Feature projection layer design.
[0021] S11: The 256-dimensional features output by the encoder for the three modalities of mRNA, miRNA and drug are mapped to a 128-dimensional unified feature space through a linear projection layer (weight matrix dimension of 256×128) to ensure that the features of different modalities have comparable dimensions and distributions.
[0022] S12: Apply batch normalization to each modality feature, with the formula: Projected feature = Batch normalization (linear transformation (original feature)). The linear transformation compresses the feature dimension through the learned weight matrix, and the batch normalization eliminates the dimensional differences between modalities through standardization (mean is 0, variance is 1).
[0023] S2: Construction of self-attention submodules.
[0024] To enhance the expression of key features within the same modality, self-attention submodules were designed for mRNA and miRNA features respectively.
[0025] S21: Input the projected mRNA features (dimension is sample number × 1 × 128) and miRNA features (dimension is sample number × 1 × 128), where "1" represents the single sequence length (because the modality features are global encoding results).
[0026] S22: An 8-head attention mechanism is used to split the 128-dimensional features according to the number of heads (16 dimensions per head).
[0027] S23: Calculate the query, key, and value matrices respectively (all generated from the input features through linear transformation).
[0028] S24: Calculate the attention weight, which is obtained by using the dot product similarity between the query and the key.
[0029] S25: Concatenate the outputs of the 8 heads of attention.
[0030] S26: Compression via linear transformation (128×128).
[0031] S27: Perform a residual connection between the compressed output and the input features (input features + attention output), and apply layer normalization (LayerNorm). The formula is: Self-attention output = Layer normalization (input features + linear transformation (multi-head attention output)).
[0032] S3: Construction of cross-attention submodule.
[0033] S31: To achieve cross-modal information interaction, five sets of cross-attention submodules are designed to capture the associations between mRNA and miRNA, mRNA and drug, miRNA and drug, drug and mRNA, and drug and miRNA, respectively.
[0034] S32: Define the input. Taking the cross-attention between "mRNA and drug" as an example, the query is the mRNA projection feature, and the key and value are the drug projection features.
[0035] S33: Multi-head attention computation: An 8-head attention mechanism is adopted to split the Query (128 dimensions) and Key (128 dimensions) into 8 heads.
[0036] S34: Calculate cross-modal attention weights:
[0037] S35: Concatenates the outputs of 8 attention heads.
[0038] S36: Compress the spliced output through linear transformation.
[0039] S37: Perform a residual connection between the compressed output and the Query (i.e., mRNA projection features) (mRNA projection features + cross-modal attention output).
[0040] S38: Apply layer normalization to the residual connection results to obtain the final output.
[0041] S4: Multimodal feature fusion.
[0042] S41: Collect the outputs of the self-attention submodule (mRNA self-attention, miRNA self-attention) and the cross-attention submodule (5 groups of cross-modal attention), for a total of 7 groups of 128-dimensional features.
[0043] S42: Concatenate the 7 sets of features into 7×128=896 dimensional features.
[0044] S43: The concatenated high-dimensional features are compressed to 128 dimensions through a fusion layer (which includes linear transformation (896×128), batch normalization, ReLU activation function and dropout (probability 0.1)) to obtain the final fused features.
[0045] 2. Construction of the uncertain integral gradient algorithm based on Dropout.
[0046] This algorithm aims to quantify the uncertainty of prediction results and analyze the contribution of each feature to the prediction, providing a basis for target identification. The specific construction process is as follows: S5: MC Dropout Uncertainty Estimation Module.
[0047] S51: Maintain the activation state of the dropout layer in the model (encoder and classifier) during the prediction phase (probability 0.3-0.5).
[0048] S52: For the same input sample, perform 10 consecutive forward propagations (each time dropout randomly occludes some neurons) to obtain 10 drug sensitivity prediction probabilities. ).
[0049] S53: Calculate the mean of the 10 predicted probabilities As the final prediction result, standard deviation As a measure of uncertainty The larger the value, the lower the reliability of the prediction result. S6: Calculation of the importance of integral gradient features.
[0050] Based on the ensemble gradient algorithm and combined with noise tunneling processing, the contribution of each feature to the prediction result is quantified. Specific steps include: S61: Baseline selection: Use the mean of each modality feature as the baseline input (mRNA baseline = mean of sample mRNA features, miRNA baseline and drug baseline are similar) to ensure that the baseline and sample input have the same distribution characteristics.
[0051] S62: Constructs a linear path from the baseline to the sample input, containing 50 interpolation points.
[0052] S63: Calculate the gradient of the model with respect to the input features at each interpolation point (the effect of small changes in features on the prediction probability).
[0053] S64: Obtain the feature attribution value through integration and summation:
[0054] S65: Noise Tunneling: To reduce the randomness of gradient estimation, 30 random Gaussian noises (standard deviation 0.05) are added to the sample input. S66: For each sample after noise perturbation, repeat steps S62-S64 to calculate the feature attribution value.
[0055] S67: The average of the feature attribution values calculated from 30 noise disturbances is taken as the final feature importance score, using the following formula:
[0056] S7: Balancing the Importance of Multimodal Features S71: Min-max normalization (mapped to the 0-1 range) is performed on the importance scores of mRNA, miRNA and drug features respectively.
[0057] S72: Normalization formula: Normalized importance = (Feature importance - Minimum feature importance) / (Maximum feature importance - Minimum feature importance).
[0058] 3. Model training and multi-task loss function.
[0059] Multi-task design: Three tasks are set up: drug sensitivity binary classification (main task), target prediction (auxiliary task), and feature reconstruction (auxiliary task).
[0060] Loss function construction.
[0061] The total loss function is the weighted sum of the losses from each task, and the formula is:
[0062] Classification loss: The difference between the predicted drug sensitivity probability and the true label (sensitive = 1, insensitive = 0) is calculated using binary cross-entropy (BCE).
[0063] Target loss: The difference between the predicted probability of multiple targets and the true target label is calculated using binary cross-entropy (BCE).
[0064] Reconstruction loss: Mean squared error (MSE) was used to calculate the difference between the original features and the reconstructed features with a weight of mRNA:miRNA:drug = 0.4:0.4:0.2.
[0065] KLD loss: Constrained variational autoencoder (VAE) latent variables follow a standard normal distribution, improving feature robustness. Training strategy: The AdamW optimizer is used (learning rate 1e-4, weight decay 1e-5); the training and validation sets are split using 5-fold cross-validation. Each training session lasts 60 epochs. The optimal model parameters are saved using the validation set AUC (area under the ROC curve) as the metric.
[0066] II. Working principles of each component of the model.
[0067] 1. Drug sensitivity prediction module.
[0068] 1.1 Multimodal Feature Input: The module takes mRNA features, miRNA features, and drug features as inputs, covering key information dimensions of gene expression and drug action. mRNA and miRNA features reflect the gene regulatory state of tumor cells, while drug features characterize the chemical structure and action properties of drugs, laying the foundation for subsequent feature interaction.
[0069] 1.2 Linear Projection Layer: As a pre-processing step for multimodal feature fusion, the core function of the linear projection layer is to align the dimensions and distributions of features from different modalities. Through independent linear transformations, the original features of mRNA, miRNA, and drugs are mapped to a unified dimensional space (e.g., 256-dimensional), eliminating dimensional differences between modalities and providing a "comparable" feature representation for cross-modal interactions of the cross-attention mechanism, ensuring the effectiveness of subsequent attention calculations.
[0070] 1.3 Cross-attention mechanism: The core engine for achieving deep fusion of multimodal features, aiming to uncover the complex relationships between mRNA, miRNA, and drug characteristics. 1.3.1 Self-Attention Submodule: Performs self-attention calculations on mRNA and miRNA features separately. Taking mRNA features as an example, by constructing a query, key, and value matrix, the correlation strength of elements within a feature is measured, dynamically strengthening key gene expression patterns (such as the feature weights of tumor driver genes), thus highlighting important information within the modality.
[0071] 1.3.2 Cross-Attention Submodule: Constructs multiple sets of cross-modal attention (such as mRNA-drug, miRNA-drug, etc.). Taking mRNA-drug attention as an example, the query generated by mRNA features is used to calculate the association weight with the key generated by drug features. Based on this, information that has a significant impact on the current mRNA regulation is screened from the values of drug features, realizing precise interaction of cross-modal features and capturing potential associations such as "gene-drug target".
[0072] 1.4 Feature Assembly and Fusion Layer: The feature assembly stage integrates the outputs of self-attention (mRNA, miRNA) and cross-attention (multiple cross-modal groups) to form a feature set containing multi-dimensional correlation information. The fusion layer then compresses and reconstructs the assembled high-dimensional features through a fully connected network and nonlinear transformation, extracting 128-dimensional fusion features as the core input for subsequent multi-task learning, ensuring that key correlation information is preserved during feature compression.
[0073] 1.5 VAE Encoding and Reparameterization: Fusion Features Entered into Variational Autoencoder (VAE) During the encoding stage, the mean ( ) and log-variance The reparameterization technique introduces Gaussian noise (…). ), to achieve from latent distribution Mid-sampling, outputting stable latent features This process preserves the distribution characteristics of features, introduces uncertainty into the model to improve feature robustness, and provides potential representations for multi-task outputs.
[0074] 1.6 Multi-task Output: Drug Sensitivity Prediction: Based on latent features, a binary classification network (e.g., fully connected layer + sigmoid activation) outputs drug sensitivity probabilities to determine the sensitivity of tumor cells to drugs. Target Prediction: Using a multi-classification network, for a pre-defined set of targets (e.g., 20 key targets), the activation probability of each target is output to identify potential biological targets for drug action. Feature Reconstruction: The VAE decoder, based on latent features, reverse-engineers the original features of mRNA, miRNA, and drugs. By calculating reconstruction loss (e.g., mean squared error), the model is constrained to learn "meaningful" feature representations, enhancing the accuracy of feature fusion.
[0075] 2. Integral gradient module (i.e., uncertain integral gradient module).
[0076] After the data is input into the Dropout-based uncertainty integral gradient module, a baseline is constructed, which is then fed into MC Dropout, processed by the forward function, and finally output to the gradient algorithm for smoothing and importance score calculation. The details are as follows: 2.1 Baseline Construction: The mean of each modality feature is used as the baseline to simulate a reference state with "no feature differences". For example, the mRNA baseline is the mean of the mRNA features of all samples. A "benchmark feature set" with the same dimension as the input features is constructed to provide a basis for comparison in subsequent gradient calculations and to quantify the impact of feature changes on model prediction.
[0077] 2.2 Forward Function and Gradient Calculation: Along the linear path of "baseline-sample features" (e.g., 50 interpolation points), the gradient of the model with respect to the input features is calculated point by point. By integrating these gradients, the contribution of the feature's change from the baseline to the true value to the prediction result is measured, i.e., the feature attribution value. This process transforms "feature importance" into gradient integrals, accurately analyzing the degree of influence of each feature on drug sensitivity prediction and target identification.
[0078] 2.3 Smoothing Gradients and Importance Matrix: Gaussian noise (e.g., 30 perturbations) is introduced to smooth the gradient calculation, reducing the randomness of a single gradient calculation. The gradient integrals after multiple perturbations are averaged to generate a stable feature importance matrix, covering the contributions of mRNA, miRNA, and drug features, providing a reliable basis for feature analysis.
[0079] 2.4 Dynamic Weight Balancing and Top-k Screening: The importance scores of mRNA, miRNA, and drug features are normalized to eliminate the influence of differences in the number and distribution of original features between modalities, ensuring the comparability of feature importance across different modalities. Based on the normalized scores, Top-k key features (e.g., Top 15) are selected, focusing on the gene and drug attributes that have the greatest impact on drug sensitivity and target prediction, thus aiding in the analysis of biological mechanisms.
[0080] III. Specific workflow of the module.
[0081] The specific workflow of the module of this invention is as follows: Figure 1 As shown.
[0082] 1. Multimodal feature processing flow (1) Data acquisition and preprocessing.
[0083] Integrate gene expression data (mRNA, miRNA sequencing results) from public databases (such as TCGA) with drug characteristics (chemical structure, target information) from drug databases (such as DrugBank), while also incorporating drug sensitivity labels from clinical trials.
[0084] Gene expression data preprocessing: missing values were filled using the median, and dimensions were eliminated by Z-score normalization.
[0085] Drug characterization preprocessing: Extracting structured representations such as molecular fingerprints and pharmacophores.
[0086] (2) Linear projection and dimension alignment.
[0087] Design independent linear transformation layers to process mRNA, miRNA, and drug characteristics.
[0088] Taking mRNA features as an example, the original high-dimensional gene expression vector (e.g., 10,000 dimensions) is mapped to 256 dimensions through matrix multiplication and bias adjustment.
[0089] 2. Implementation of cross-attention mechanism.
[0090] (1) Self-attention construction.
[0091] Step 1: Calculate self-attention weights for mRNA features (or miRNA features) based on dot product similarity.
[0092] Step 2: Treat the feature matrix as a sequence, and generate a Query, Key, and Value for each element. Use the Softmax function to normalize the attention weights.
[0093] Step 3: Employ an 8-head attention mechanism to split the 256-dimensional features into 16 dimensions per head, and compute multiple sets of attention in parallel. Concatenate the outputs of the 8-head attention mechanism and compress them back to 256 dimensions using a linear transformation.
[0094] Step 4: Perform residual connections and layer normalization.
[0095] (2) Cross-attention interaction.
[0096] First, five cross-modal attention mechanisms were constructed (mRNA and miRNA, mRNA and drug, miRNA and drug, drug and mRNA, and drug and miRNA). Taking mRNA-drug attention as an example, mRNA features were used to generate queries, and drug features were used to generate keys and values. Second, the similarity between the mRNA query and the drug key was calculated to obtain the cross-modal attention weights. The drug value information was then fused using the attention weights. Then, the weighted fusion result was residually connected to the mRNA query and layer normalized. The outputs of the self-attention (mRNA, miRNA) and the five cross-attention mechanisms were concatenated to form a 7-channel 128-dimensional feature set. Finally, the feature set was compressed to 128 dimensions through a fusion layer (linear transformation (896×128), batch normalization, ReLU activation, dropout (0.1)).
[0097] 3. VAE Encoding and Multi-task Learning.
[0098] (1) VAE encoding and reparameterization.
[0099] The fused feature input VAE encoder generates the mean through a multilayer perceptron. The log-variance (logvar) describes the distribution parameters of the latent features. Using a reparameterization technique, Gaussian noise is introduced to generate latent features z, enabling the model to learn robust feature representations and improve generalization ability. A VAE loss is constructed, including reconstruction loss (mean squared error, measuring the difference between the original and reconstructed features) and KL divergence loss (constraining the latent distribution to approximate a standard normal distribution), balancing the quality of feature reconstruction with the reasonableness of the distribution.
[0100] (2) Multi-task loss optimization.
[0101] A loss function is constructed, and the total loss is composed of the drug sensitivity prediction loss (binary cross-entropy, matching binary classification tasks), the target prediction loss (multi-label cross-entropy, adapting to the multi-classification requirements of 20 targets), the VAE reconstruction loss (mean squared error), and the KL divergence loss weighted together. By dynamically adjusting the weights (e.g., classification loss weight 1.0, reconstruction loss weight 0.4), the multi-task objectives are optimized in a coordinated manner.
[0102] The Adam optimizer was used, with a learning rate of 1e - 4 and weight decay of 1e - 5. The dataset was split using 5-fold cross-validation, and the model was trained for 60 epochs per fold. The model with the best AUC on the validation set was saved to ensure model stability and generalization.
[0103] Integral gradient and interpretability analysis.
[0104] Baseline and gradient integral.
[0105] (1) Calculate the mean of each modality feature as the baseline (e.g., the mRNA baseline is the mean vector of all sample mRNA features) and construct a benchmark representation with the same dimension as the input features.
[0106] (2) Construct a linear interpolation path (50 equidistant points) from the baseline to the sample features, calculate the gradient of the model to the input features point by point, accumulate the gradient values by trapezoidal integral method, quantify the contribution of feature changes to the prediction results, and generate preliminary attribution values.
[0107] (3) Noise smoothing and importance analysis: Gaussian noise (standard deviation 0.05) is added to the sample features 30 times, the gradient integral is repeatedly calculated, and the mean is taken as the smoothing attribution value to reduce the randomness error of gradient calculation.
[0108] (4) The attribution values of mRNA, miRNA, and drug features are normalized by min-max to ensure that the importance of different modal features is comparable. Based on the normalization scores, the top-k key features are selected, and a list of genes and drug attributes that have the greatest influence on drug sensitivity prediction and target identification is output.
[0109] IV. Experimental verification and result analysis.
[0110] 1. Model training and evaluation.
[0111] (1) Training setup: On multimodal datasets (such as THCA, MESO, LGG tumor data), 5-fold cross-validation was used to train the model, with AUC and accuracy (ACC) as the core indicators to monitor the training process and model convergence.
[0112] (2) Performance Comparison: To verify the superiority of the method of this invention, five classic models (Random Forest, SVM, Logistic Regression, XGBoost, Simple DNN) were selected and 5-fold cross-validation was carried out on four real / simulated datasets (THCA, MESO, LGG, and simulation data). The AUC (classification performance) and ACC (accuracy) were used as the core indicators for comparison. The results are shown in the table below: Table 1. Performance comparison of each model (mean of 5-fold cross-validation)
[0113] Overall performance advantages: The method of this invention has higher average AUC (0.8170) and average ACC (0.7496) on all datasets than other comparative models.
[0114] On real tumor datasets (THCA, MESO, LGG), the average AUC reached 0.8943, which is 1.14% higher than the second-best Simple DNN (0.8078), fully demonstrating the value of multimodal fusion and cross-attention mechanisms in mining feature associations and improving prediction performance. Dataset adaptability: In the simulation dataset (with high-dimensional heterogeneity and small sample characteristics), the AUC and ACC of the method of this invention are still the highest, indicating that the integral gradient and VAE module can effectively cope with noise interference and small sample challenges, and ensure the robustness and generalization ability of the model. Compared with traditional models: When processing multimodal high-dimensional data, traditional machine learning models such as Random Forest and SVM have difficulty capturing complex nonlinear relationships and cross-modal information, resulting in significantly lower performance than the deep learning-based method of this invention. Although Simple DNN is a deep learning model, it is inferior to the method of this invention in terms of the comprehensiveness of feature utilization and model stability because it does not employ cross-attention fusion and interpretability enhancement design.
[0115] 2. Uncertainty and interpretability verification.
[0116] (1) Uncertainty quantification: By using MC Dropout to keep the Dropout layer active during the prediction phase, the same sample is forward-propagated multiple times (e.g., 10 times), and the mean and standard deviation of the prediction probability are calculated. The results show that the proportion of samples with a standard deviation > 0.15 is < 5%, and their prediction accuracy is significantly lower than the overall rate, effectively identifying high-risk predictions.
[0117] (2) Verification of feature importance: Among the top-k key features, the proportions of mRNA, miRNA, and drug features are consistent with the biological prior (gene regulation dominates drug action). For example, the proportion of mRNA features is 40%, which verifies the rationality of the integral gradient module analysis results and assists in the discovery of drug targets and the study of mechanisms of action.
[0118] V. Summary.
[0119] The model of the present invention has the following advantages: 1. Cross-attention mechanism improves multimodal fusion efficiency.
[0120] Compared to traditional splicing or weighted fusion methods, cross-attention significantly improves the complementarity of multimodal features by dynamically learning the association weights between modalities. Experiments show that on real datasets such as THCA and LGG, the AUC of the model that integrates multimodal features (0.9057, 0.9062) is significantly higher than that of the single modality (mRNA modality AUC is about 0.87-0.89, and drug modality AUC is about 0.85-0.88). 2. Uncertainty integral gradient enhances model interpretability.
[0121] Uncertainty quantification achieved through MC Dropout can effectively identify high-risk prediction samples (such as samples with a standard deviation > 0.15); feature importance scores calculated by integral gradient can accurately screen out key targets related to drug sensitivity (such as the top 5 importance scores in mRNA features all > 0.8), providing a clear direction for biological experimental validation. 3. Its predictive performance is superior to that of traditional models.
[0122] Compared with Random Forest, SVM, XGBoost and Simple DNN, the present invention achieves an average AUC of 0.8170 and an average accuracy of 0.7496 in 5-fold cross-validation, which are 1.14% and 1.31% higher than the best comparison model (Simple DNN), respectively, with significant advantages, especially in high-dimensional small sample scenarios.
[0123] 4. It exhibits outstanding robustness and generalization ability.
[0124] The synergistic effect of the multi-task loss function and the VAE module enables the model to maintain stable performance even with missing data or noise interference (e.g., the AUC decreases by less than 2% after adding 10% noise), making it suitable for real-world scenarios where clinical data quality varies greatly. In summary, this invention, through a carefully designed cross-attention mechanism and uncertainty integral gradient algorithm, achieves efficient fusion and interpretable prediction of multimodal biomedical data, providing important technical support for drug selection and target validation in precision oncology.
[0125] This model achieves accurate prediction of drug sensitivity and effective target identification through a collaborative design of multimodal feature fusion, cross-attention interaction, multi-task learning, and interpretability analysis, providing technical support and biological mechanism analysis tools for precision cancer treatment.
Claims
1. A model based on multi-modal cross-attention and uncertainty integrated gradients, characterized in that, The model comprises a prediction module and an uncertainty integral gradient module, the prediction module comprises a cross-attention module and a multi-task learning module; the cross-attention module comprises a linear projection layer and a cross-attention fusion module, the linear projection layer is used for inputting multi-modal features and aligning the dimensions and distributions of different modal features; the cross-attention fusion module is used for realizing cross-modal information interaction, and the multi-task learning module is used for introducing uncertainty and outputting latent features; The uncertainty integral gradient module is used for quantifying the uncertainty of the prediction result and analyzing the contribution of each feature to the prediction result.
2. The model of claim 1, wherein, The cross-attention fusion module comprises a cross-attention mechanism and a fusion layer, the cross-attention mechanism is used for inputting the output data of the linear projection layer, the output data of the cross-attention mechanism is spliced to obtain high-dimensional features, and the high-dimensional features are input into the fusion layer; the fusion layer compresses and reconstructs the high-dimensional features through a full connection network and a nonlinear transformation, and outputs fusion features; The multi-task learning module comprises a variational autoencoder, the variational autoencoder is used for encoding the fusion features, the output data of the variational autoencoder is reparameterized, sampled from a latent distribution, and latent features are output.
3. The model of claim 2, wherein, The cross-attention mechanism comprises a self-attention sub-module and a cross-attention sub-module, the self-attention sub-module is used for self-attention calculation, and the cross-attention sub-module is used for constructing multiple groups of cross-modal attention; The feature splicing comprises: integrating the output data of the self-attention sub-module and the output data of the cross-attention sub-module.
4. The model of claim 1, wherein, The uncertainty integral gradient module comprises an MC Dropout uncertainty estimation module and an integral gradient module, the MC Dropout uncertainty estimation module comprises a forward function; The uncertainty integral gradient module is used for inputting the output data of the variational autoencoder, constructing a baseline, inputting the forward function, calculating a gradient integral, smoothing, taking an average, generating a feature importance matrix, performing dynamic weight balancing, screening top-k key features, and outputting a prediction result.
5. A method for predicting drug sensitivity and / or target identification, characterized in that, The method comprises the following steps: Inputting sample data to be evaluated into the model in any one of claims 1-4 to obtain a prediction result.
6. The method of claim 5, wherein, The sample data comprises gene expression data and drug features; The prediction result comprises a drug sensitivity prediction result, a target recognition result and / or a drug attribute result.
7. The method of claim 6, wherein, The gene expression data comprises mRNA sequencing data and miRNA sequencing data, and the drug features comprise chemical structures, target information and drug sensitivity labels.
8. A system for predicting drug sensitivity and / or target identification, characterized in that, The method comprises: a data storage module for sample data to be evaluated and the module in any one of claims 1-4; a data analysis module for performing analysis according to any one of claims 5-7; and a data display module for outputting and displaying the prediction result.
Citation Information
Cited By
Drug target prediction method based on cross-modal attention and uncertainty evaluation
CN122050487A