Explanatable model robust training method based on topological regular terms
By introducing topological regular terms and adaptive weight adjustment strategies into the interpretability model, the problems of insufficient robustness of model interpretation and fixed regular terms weight in the prior art are solved, and stronger anti-interference and generalization capabilities are achieved.
Patent Information
- Application Number
- CN202510168881.6
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-02-17
- Publication Date
- 2025-06-10
AI Technical Summary
The existing regularization method based on gradient alignment fails to effectively capture topological features in the interpretation image when improving the robustness of the model interpretation, and the regular term weight is fixed, making it difficult to dynamically adapt to multiple types of perturbations, affecting the generalization ability of the model.
A robust training method for interpretability model based on topological regular terms is adopted to generate perturbations through semantic-maintained perturbations, calculate gradient differences and topological differences, dynamically adjust the regular term weight, combine cross entropy loss to form a total loss, and backpropagation updates the model parameters.
It effectively improves the robustness of model interpretation, enhances the anti-interference ability of topological structures in complex scenarios, and improves the generalization ability of the model in a multi-type perturbation environment.
Smart Images

Figure CN120125933A_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to a robust training method for an interpretable model based on a topological regularization term, and belongs to the field of artificial intelligence security. Background Art
[0002] In the era of big data, the rapid development of deep learning technology has promoted the wide application of interpretable models in key fields such as medical diagnosis and financial decision-making that require a transparent decision-making process. However, when faced with input perturbations or adversarial attacks, the robustness of their interpretations is still insufficient. An attacker can interfere with or tamper with the interpretation output of the model to conceal the true decision-making logic of the model, thereby misleading users and causing them to make wrong judgments. Especially in high-risk fields such as healthcare and finance, incorrect interpretations may lead users to make inappropriate treatment decisions or investment judgments, thereby causing serious losses to life and property. Therefore, it is crucial to improve the robustness of model interpretations through the robust training of interpretable models. The main current method for improving the robustness of model interpretations is the regularization method based on gradient alignment.
[0003] The regularization method based on gradient alignment introduces a gradient difference regularization term into the loss function of the model, constrains the gradient difference between the original input and the perturbed input, and makes the gradients of the two consistent in magnitude and direction, so as to ensure that the model can maintain the stability of the interpretation result when faced with perturbations. However, this method only focuses on the gradient consistency before and after sample perturbation, and fails to effectively capture key topological features such as connectivity, holes, and circular structures in the interpretation image. Even if the local gradients are consistent, the perturbation may still change the spatial relationship or structural dependence between key regions to destroy these topological features, resulting in the breakage of connected regions or changes in holes, thereby affecting the stability of the interpretation result and weakening the anti-interference ability of the model interpretation; in addition, the existing regularization term weights usually remain fixed throughout the training process and cannot be dynamically adjusted according to the complexity of different input features or the type of perturbation. When faced with diverse inputs and complex perturbation environments, the model cannot flexibly adjust the impact of the regularization term on the total loss, resulting in overly strict or overly loose constraints on edge features or key regions, making it difficult to stably focus on important features, and thus affecting the generalization ability of the model.
[0004] In summary, the existing regularization method based on gradient alignment has the following problems when improving the robustness of model interpretations: (1) Only using gradient difference features for model training fails to effectively capture the topological features in the interpretation image, resulting in the topological structure of key regions being easily damaged in complex scenarios, thereby affecting the anti-interference ability of the model interpretation; (2) The regularization term weights remain fixed during the training process, resulting in the model being difficult to dynamically adapt to multiple types of perturbations, thereby affecting its generalization ability. Therefore, the present invention proposes a robust training method for an interpretable model based on a topological regularization term. Summary of the Invention
[0005] The object of the present invention is to address the problems that the existing gradient alignment-based regularization method only uses gradient difference features for model training, which affects its anti-interference ability, and the fixed regularization term weight is difficult to adapt to multiple types of perturbations, reducing the generalization ability of the model. A robust training method for an interpretable model based on a topological regularization term is proposed.
[0006] The design principle of the present invention is as follows: First, perform semantically-preserving perturbations on the original samples to obtain perturbed samples; second, input the original samples and the perturbed samples into the target model, calculate their gradient values respectively, and generate corresponding explanation images; then calculate the gradient differences before and after the sample perturbations based on the cosine distance and the Euclidean distance, and use the persistent homology method to extract the topological features of the explanation images before and after the sample perturbations to quantify the topological differences; finally, use the gradient differences and the topological differences as regularization terms, and jointly form the total loss with the cross-entropy loss. Dynamically adjust the regularization term weights according to the proportion of each difference in the total loss, and update the model parameters through backpropagation to improve the robustness of model interpretation.
[0007] The technical solution of the present invention is realized through the following steps:
[0008] Step 1, use a data augmentation strategy to expand sample diversity, and perform semantically-preserving perturbations on the original samples to obtain perturbed samples.
[0009] Step 1.1, introduce sample diversity using data augmentation strategies such as color perturbation and Gaussian blur.
[0010] Step 1.2, randomly sample from a uniform distribution, and perturb the original samples according to the perturbation intensity to obtain perturbed samples.
[0011] Step 2, input the original samples and the perturbed samples into the target model, calculate their gradient values respectively, and generate corresponding explanation images.
[0012] Step 2.1, input the original samples and their corresponding perturbed samples, obtain their respective outputs, and calculate the gradient values of each output with respect to its input respectively.
[0013] Step 2.2, generate the original explanation image and the corresponding perturbed explanation image based on the calculated gradient values.
[0014] Step 3, calculate the gradient differences before and after the sample perturbations based on the cosine distance and the Euclidean distance (l 2 distance), and use the persistent homology method to extract the topological features of the explanation images before and after the sample perturbations to quantify the topological differences.
[0015] Step 3.1, calculate the cosine distance and the l 2 distance between the gradients of the samples before and after the perturbation. The cosine distance measures the directional difference between the gradients before and after the perturbation; and the l2 The distance measures the magnitude difference between the gradients before and after perturbation.
[0016] Step 3.2: Conduct topological structure analysis on the original interpretation image and the perturbed interpretation image, and extract their topological features, including connected components, holes, and annular structures.
[0017] Step 3.3: Use the persistent homology method to generate a persistence diagram. Set a threshold based on the pixel values of the image, and gradually increase the threshold to layer-by-layer construct the topological structure of the image at different threshold scales. During this process, continuously track the topological features of the image, generate a persistence diagram based on the generation and disappearance times of the topological features at different scales, so as to capture the persistence of the features at different scales, and calculate the topological difference between the interpretation images before and after perturbation based on the persistence diagram.
[0018] Step 4: Use the gradient difference and the topological difference as regularization terms, and jointly form the total loss with the cross-entropy loss, and dynamically adjust the regularization term weights according to the proportion of each difference in the total loss.
[0019] Step 4.1: Construct a total loss function, including cross-entropy loss, gradient difference, and topological difference.
[0020] Step 4.2: Introduce an adaptive weight adjustment strategy. In each training iteration, calculate the proportion of each regularization term in the total loss, and dynamically adjust its weight according to the proportion. Update the model parameters through backpropagation to minimize the loss.
[0021] Beneficial effects
[0022] Compared with the traditional gradient alignment-based regularization method, the present invention combines topological consistency regularization and proposes to use the topological features of the interpretation image for model training to enhance the stability of the model interpretation in the face of perturbations; in addition, the present invention proposes an adaptive weight adjustment strategy, which can dynamically adjust the weights of each regularization term in the loss function, effectively cope with various types of perturbations, and further improve the robustness of the model interpretation in complex scenarios. Description of the drawings
[0023] Figure 1 It is a schematic diagram of the robust training method of the interpretable model based on topological regularization terms of the present invention. Detailed implementation manners
[0024] To better illustrate the purpose and advantages of the present invention, the following further details the implementation manners of the method of the present invention with examples.
[0025] The experimental data comes from the open-source dataset CIFAR-10, which covers color images of 10 categories, with 6,000 images in each category, for a total of 60,000 images. These 10 categories cover a wide range from transportation vehicles such as airplanes and cars to natural organisms such as birds and cats. The size of each image is 32×32 pixels and it is an RGB color image, that is, each pixel contains information about the three color channels of red, green, and blue.
[0026] The experiment uses the Random Perturbation Similarity (RPS) as the evaluation metric, and its calculation formula is shown in Equation (1).
[0027]
[0028] Among them, δ x represents the perturbation value of the input sample x, y represents the model output, D refers to the input dataset, is the expected perturbation term, S refers to the measurement method, ∈ refers to the perturbation intensity, N is the number of data tuples in D, U d represents the uniform distribution, and h(x) represents the interpretive image of the sample x. In the experiment, ∈ is set to 4, 8, and 16 respectively, and cosine similarity, Pearson correlation coefficient, and structural similarity are used as the similarity metric S.
[0029] Among them, the Pearson correlation coefficient r is an effective metric for evaluating the linear correlation between two variables, and its calculation formula is shown in Equation (2).
[0030]
[0031] Among them, x i and y i are two sample values, and are the means of these samples respectively, and n is the sample size.
[0032] Structural similarity measures the visual similarity between two images, which takes into account the structural information of the images, including brightness, contrast, and geometric structure. Its calculation formula is shown in Equation (3).
[0033]
[0034] Among them, μ x and μ y are the averages of images x and y, σ x and σ y are the standard deviations of x and y respectively, and are the variances of x and y respectively, σ xy is the covariance of x and y, C 1 、C 2 、C 3is a stability constant used to avoid a zero denominator, C i =(k i L) 2 and k i << 1.
[0035] The specific process of this experiment is as follows:
[0036] Step 1: Use data augmentation strategies to expand sample diversity and perform semantics-preserving perturbations on the original samples to obtain perturbed samples.
[0037] Step 1.1: Use color perturbation and Gaussian blur strategies during data preprocessing to enhance the training set as the original input samples.
[0038] Step 1.2: Generate a tensor with the same dimension as the input samples, where each element is sampled from a uniform distribution [-1, 1] and multiplied by a perturbation intensity of ∈ = 4 to obtain the perturbation.
[0039] Step 1.3: Apply the generated perturbation to the original input samples, that is, create a copy and add each element of the perturbation to the original input samples to obtain the perturbed input samples.
[0040] Step 2: Input the original samples and the perturbed samples into the target model, calculate their gradient values respectively, and generate corresponding explanation images.
[0041] Step 2.1: Replace all activation functions after the convolutional layer and fully connected layer in the neural network with the softplus function.
[0042] Step 2.2: In each round of training, the model inputs the original samples and their corresponding perturbed samples respectively, obtains their respective tensor outputs, converts them into scalar forms, calculates the gradients of the outputs of both with respect to their respective inputs, and obtains the original sample gradients and perturbed sample gradients. The calculation formula is shown in Equation (4).
[0043]
[0044] Among them, x is the input, f(x) is the corresponding output, is to calculate the gradient of x, and A grad is the calculated gradient value and also the attribution value.
[0045] Step 2.3: Normalize the attribution values of the original samples and the perturbed samples, and map the attribution values to between [0, 1], as shown in Equation (5).
[0046]
[0047] Among them, A is the attribution matrix, and A normis the normalized attribution matrix, where max(A) and min(A) are the maximum and minimum values in the attribution matrix A, respectively.
[0048] Step 2.4, convert the normalized attribution matrix into an explanation image through color mapping.
[0049] Step 3, calculate the gradient difference before and after sample perturbation based on the cosine distance and l 2 distance, and use the persistent homology method to extract the topological features of the explanation images before and after sample perturbation to quantify the topological difference.
[0050] Step 3.1, based on the original sample gradient and perturbed sample gradient obtained in Step 2.2, calculate the cosine distance and l 2 distance between them, where the cosine distance calculation formula is shown in Equation (6), and the l 2 distance calculation formula is shown in Equation (7).
[0051]
[0052] Among them, represents the cosine distance, represents the sample gradient, δ x represents the perturbation value of the sample x, and cossim represents the cosine similarity function. v and w respectively represent two variables when calculating the cosine similarity, v T represents the transpose of v, ‖v‖ 2 and ‖w‖ 2 respectively represent the l 2 distance of v and w.
[0053]
[0054] Among them, represents the l 2 distance, represents the sample gradient, δ x represents the perturbation value of the sample x.
[0055] Step 3.2, set the threshold t = 0.5, set the regions with pixel values greater than 0.5 to 1, and other regions to 0, and extract the regions that contribute the most to the prediction result in the explanation images before and after perturbation.
[0056] Step 3.3, connect the pixels in the binarized high - contribution regions through the Vietoris - Rips complex to obtain the topological features of this region, including connected components, holes, and loop structures.
[0057] Step 3.4, by adjusting the threshold t, gradually construct multiple different complexes, use persistent homology technology to calculate the appearance and disappearance times of connected components and loops in the complex, and generate a persistence diagram, represented as point pairs P = {(b i , d i )} and Q = {(b j , d j )}. Among them, (b i , d i ) and (b j , d j ) represent the appearance time and disappearance time of the i-th and j-th topological features respectively.
[0058] Step 3.5, construct the distance matrix D ij , calculate the Euclidean distance between all point pairs in P and Q, as shown in Equation (8), and find the optimal matching γ of the point pairs in P and Q based on the linear programming algorithm, so that the sum of the Euclidean distances of all points is the smallest, represented as the optimization problem in Equation (9).
[0059]
[0060] Among them, D ij represents the Euclidean distance between the i-th point of P and the j-th point of Q.
[0061]
[0062] Step 3.6, use the Wasserstein distance metric to measure the distance between the persistence diagrams corresponding to the interpreted images before and after perturbation, as shown in Equation (10). Take the calculated topological difference as the regularization term.
[0063]
[0064] Among them, W p refers to the Wasserstein distance of order p, and the order p takes 1.
[0065] Step 4, take the gradient difference and the topological difference as regularization terms, and jointly constitute the total loss with the cross-entropy loss, and dynamically adjust the regularization term weights according to the proportion of each difference in the total loss.
[0066] Step 4.1, construct the total loss function, including the cross-entropy loss, the gradient difference regularization term and the topological difference regularization term, in the form shown in Equation (11).
[0067]
[0068] Among them, x is the input, y is the output, is the expected perturbation term, L(x, y) is the total loss function, L CE(x, y) is the cross-entropy loss, and δ x is the perturbation sampled from a uniform distribution, and λ cos is the weight of the cosine distance regularization term, is the weight of the topo distance regularization term, and λ is the gradient calculated for the input x, and h(x) is the interpretive image of the output with respect to the input sample x. is the cosine distance, is the topo distance, and Γ
[0069] Step 4.2, set the initial weights
[0070] Step 4.3, in each training iteration, calculate the relative contribution of each regularization term to the total loss function, as shown in Equation (12).
[0071]
[0072] where L total represents the total loss, and L cos , and L topo respectively represent the cosine distance regularization term, the cos , and r topo represent the contribution of each regularization term to the total loss in the current iteration.
[0073] Step 4.4, set the target contribution ratio of each regularization term to the total loss, and set the target contribution ratios of the cosine distance regularization term, the distance regularization term, and the topological difference regularization term to be
[0074]
[0075] where α cos , and α topo are the adaptive adjustment factors for the cosine distance regularization term, the
[0076] distance regularization term, and the topological difference regularization term respectively.
[0077]
[0078] Among them, is the current weight of the regularization term, and α i is the adaptive adjustment factor for dynamically adjusting the weight. is the weight of the regularization term after update. When i takes cos, and topo, they respectively represent the cosine distance regularization term, distance regularization term and topological difference regularization term.
[0079] Step 4.6, update the backpropagation parameters until the model optimization training is completed.
[0080] Test results: The experiment uses the interpretable model robust training method based on topological regularization terms, and uses the test set in the CIFAR-10 dataset to conduct the robustness test of model interpretation. When the perturbation intensity ∈ takes 4, 8, and 16 respectively, the RPS values are 98.7%, 95.34%, and 83.91% respectively. Compared with the existing gradient alignment-based regularization method, it is increased by 7.76% on average, indicating that the method of the present invention can effectively improve the robustness of the model interpretation in the face of complex perturbations.
[0081] The above specific description further details the purpose, technical solution and beneficial effects of the invention. It should be understood that the above is only the specific embodiment of the present invention and is not used to limit the protection scope of the present invention. Any modification, equivalent replacement, improvement, etc. made within the spirit and principle of the present invention shall be included in the protection scope of the present invention.
Claims
1. A robust training method for interpretable models based on topological regularization terms, characterized by The method comprises the following steps: Step 1: Use data augmentation strategy to expand sample diversity, randomly sample from uniform distribution, generate perturbations according to perturbation intensity, apply them to original input samples, and obtain perturbation input samples; Step 2: Input the original sample and the perturbation sample into the target model, calculate the gradient of the output relative to the input, and generate the original interpretation image and the perturbation interpretation image based on the calculated gradient value; Step 3: First, the gradient difference of the sample before and after the disturbance is calculated based on the cosine distance and the Euclidean distance. Secondly, the topological structure of the interpreted image before and after the disturbance is analyzed to extract its topological features, including connected components, holes and ring structures. Then, the persistent homology method is used to generate a persistence graph to capture the topological features and persistence of the image at different scales. Finally, the topological difference of the interpreted image before and after the disturbance is calculated based on the persistence graph. Step 4: Take the gradient difference and topological difference as regularization terms, and together with the cross entropy loss, form the total loss function. An adaptive weight adjustment strategy is introduced. In each training iteration, the regularization term weight is dynamically adjusted according to the proportion of each difference in the total loss, and the model parameters are updated through back propagation.
2. The robust training method for interpretable models based on topological regularization terms according to claim 1, characterized in that: In step 3, the threshold t=0.5 is set, and the area with pixel value greater than 0.5 is set to 1, and the other areas are set to 0. The area that contributes most to the classification result in the interpreted image before and after the disturbance is extracted. The pixels in the binarized high-contribution area are connected through the Vietoris-Rips complex to obtain the topological features of the area, including connected components, holes and ring structures. By adjusting the threshold t, multiple different complexes are gradually constructed. The persistent homology technique is used to calculate the appearance and disappearance time of the connected components and rings in the complex to generate a persistence graph, which is represented by the point pair P={(b i ,d i )} and Q={(b j ,d j )}, where (b i ,d i ) and (b j ,d j ) represent the appearance time and disappearance time of the i-th and j-th topological features respectively. The Wasserstein distance is used to measure the distance between the persistence graphs corresponding to the explained images before and after the perturbation, and the calculated topological difference is used as the regularization term.
3. The robust training method for interpretable models based on topological regularization terms according to claim 1, characterized in that: In step 4, the topological difference is added as a regularization term to the total loss function, and the initial weights of each regularization term are set. where λ cos is the weight of the cosine distance regularization term, yes The weight of the distance regularization term, λ topo is the weight of the topological difference regularization term. In each round of training iteration, the relative contribution of each regularization term to the total loss function is calculated to obtain r cos =L cos / L total , r topo =L topo / L total , where L total represents the total loss, L cos , and L topo Respectively represent the cosine distance regularization term, Distance regularization term and topological difference regularization term, r cos , and r topo Indicates the contribution of each regularization term to the total loss in the current iteration, sets the target contribution ratio of each regularization term to the total loss, sets the cosine distance regularization term, The target contribution ratios of the distance regularization term and the topological difference regularization term are The adaptive adjustment factor is calculated by the ratio of the target contribution ratio to the current relative contribution Where i is cos, and topo, respectively, represent the cosine distance regularization term, The distance regularization term and the topological difference regularization term introduce a smooth update factor η = 0.05, and update the weight of each regularization term according to the adjustment factor, expressed as in is the current weight of the regularization term, is the weight after the regularization term is updated, and the back propagation parameters are updated until the model optimization training is completed.
Citation Information
Cited By
Ginger stem and leaf disease and insect pest recognition optimization method based on image recognition
CN121982705A
An image recognition-based ginger stem and leaf disease and pest identification optimization method
CN121982705B