Cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention
By using a mask-guided causal intervention method to generate a counterfactual mask set and combining it with a dual-branch convolutional neural network and a collaborative domain alignment module, the problems of false correlation between classes and labels and sample scarcity in cross-domain hyperspectral image classification are solved, thereby improving classification accuracy and model generalization ability.
Patent Information
- Application Number
- CN202510793415.7
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-06-13
- Publication Date
- 2025-09-16
AI Technical Summary
In cross-domain hyperspectral image classification, the spurious correlation between classes and labels and the scarcity of labeled samples in the target domain lead to insufficient generalization ability of the model in unseen classes, hindering the applicability of traditional methods.
A cross-domain few-shot hyperspectral image classification method based on mask-guided causal intervention is adopted. A counterfactual mask set is generated through a random mask image mixing module. Combined with a two-branch convolutional neural network and a collaborative domain alignment module, model training is performed to extract representative class-discriminative causal features and reduce the false correlation between classes and labels.
The classification accuracy of hyperspectral images is improved, and the application value of the model in the fine classification of cross-domain hyperspectral images is enhanced, which has important theoretical significance.
Smart Images

Figure CN120656062A_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of hyperspectral image classification, and in particular to a cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention. Background Art
[0002] As one of the core tasks of remote sensing information processing, hyperspectral image classification plays an important role in environmental monitoring, precision agriculture, national defense and security, etc. With the rapid development of hyperspectral remote sensing technology and the continuous increase in application demand, hyperspectral image classification technology has made great progress.
[0003] However, the complex spectral characteristics of cross-domain hyperspectral images lead to false correlations between classes and labels, which weakens the generalization ability of the model in unseen classes. At the same time, the scarcity of labeled samples in the target domain also hinders the extraction of reliable category-specific features. These problems further restrict the applicability of traditional cross-domain hyperspectral image classification methods. Summary of the Invention
[0004] The present invention provides a cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention to overcome the above technical problems.
[0005] In order to achieve the above object, the technical solution of the present invention is:
[0006] A cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention specifically includes the following steps:
[0007] S1: Obtain a target domain dataset of hyperspectral scene images and a source domain dataset that has been labeled with categories;
[0008] And the target domain dataset includes unlabeled data samples and labeled data samples;
[0009] Using unlabeled data samples as a test set, using the source domain dataset and labeled data samples as a training set, and randomly selecting a support set and a query set from the training set;
[0010] S2: Preprocess the support set and query set to obtain the optimized support set and optimized query set respectively;
[0011] S3: Randomly divide the optimized support set into a training support set and a validation support set;
[0012] Randomly divide the optimized query set into a training query set and a validation query set;
[0013] S4: Based on the training support set and the training query set, the constructed hyperspectral image classification model is trained to obtain the optimal hyperspectral image classification model;
[0014] The hyperspectral image classification model consists of a sequentially connected random mask image mixing module, a two-branch convolutional neural network module with multi-scale spectral convolution blocks, a linear classifier, and a collaborative domain alignment module;
[0015] The model training is specifically as follows:
[0016] S41: Through the random mask image mixing module, mask-guided causal intervention is performed on the training query set to obtain the counterfactual mask set;
[0017] S42: Perform pixel feature extraction on the training query set, the training support set, and the counterfactual mask set through a dual-branch convolutional neural network module to obtain support set features, query set features, and mask set features; and obtain a class prototype of the support set features based on the support set features;
[0018] S43: Using the Euclidean distance metric method, according to the distance between the class prototypes of the query set features and the support set features, and using a linear classifier to obtain the classification probability of the class prediction result, the class prediction result of the query sample is obtained, and then the trained hyperspectral image classification model is obtained;
[0019] S44: Through the collaborative domain alignment module, based on the total model loss function constructed by the support set features, the query set features, and the mask set features, and based on the validation support set and the validation query set, the trained hyperspectral image classification model is validated to determine whether the output of the trained hyperspectral image classification model converges;
[0020] If yes, then confirm that the trained hyperspectral image classification model is the optimal hyperspectral image classification model; otherwise, adaptively adjust the weight parameters of the trained hyperspectral image classification model based on the back propagation method, and repeat step S41;
[0021] S5: Based on the optimal hyperspectral image classification model, hyperspectral image classification of unlabeled data samples in the test set is realized.
[0022] Furthermore, the pre-processing method in S2 specifically includes:
[0023] Performing Gaussian noise processing on the hyperspectral scene image with sample deviation caused by scarcity of labeled samples to obtain hyperspectral scene image samples;
[0024] And the expression for Gaussian noise processing is
[0025]
[0026] Where: represents a hyperspectral scene image sample; x represents a hyperspectral scene image with sample bias caused by scarcity of labeled samples; α represents a random sampling factor; β represents an empirical value used to control the noise intensity to match the signal-to-noise ratio of the hyperspectral data;
[0027] Perform dimension unification processing on hyperspectral scene image samples.
[0028] Furthermore, the method for obtaining the counterfactual mask set in S41 is specifically as follows:
[0029] S411: traverse each image sample in the training query set in sequence, and take any image sample in the training query set as a target sample;
[0030] The remaining image samples in the training query set except the target samples are used as interference samples;
[0031] The interference sample and the target sample are data samples that belong to the same source domain dataset or target domain dataset but have different image categories;
[0032] S412: Grid-dividing the target sample to obtain an n*n grid image;
[0033] Define the central pixel grid area b*b of the grid image as a constant pixel value, and perform a random mask operation on the other grid areas except the central pixel grid area to obtain a mask matrix image;
[0034] S413: Perform an inverse mask operation on the interference sample to obtain an inverse mask matrix image;
[0035] And obtain the counterfactual mask set according to the mask matrix image and the inverse mask matrix image;
[0036] And the formula for obtaining the counterfactual mask set is
[0037] x m =x q ⊙m+x i ⊙(1-m)
[0038]
[0039] Where: x m represents the counterfactual image sample; x q Indicates that each query sample in the query set is the target sample; x i Represents interference samples of the same domain as the input samples but of different categories; ⊙ represents the element-by-element product operator; m represents a random mask matrix generated by a preset random mask generator; i, j represent the position coordinates in the random mask matrix; center represents a fixed central pixel grid area; r represents a value randomly sampled from a uniform distribution of [0,1]; δ represents the pixel interference ratio.
[0040] Furthermore, the dual-branch convolutional neural network module in S42 includes an input layer, a spectral feature extraction branch, a spatial feature extraction branch, a channel splicing module, and an output layer;
[0041] The input layer is used to transmit the training query set or the training support set or the counterfactual mask set to the spectral feature extraction branch and the spatial feature extraction branch respectively;
[0042] The spectral feature extraction branch includes a multi-scale spectral convolution module and a global spectral attention module connected in sequence. The multi-scale spectral convolution module is used to extract the multi-scale spectral features of the input data and obtain a spectral feature map.
[0043] The global spectral attention module is used to extract the global spectral attention features of the spectral feature map and obtain the global spectral attention feature map;
[0044] The spatial feature extraction branch includes a 3D convolution module, a spatial branch residual block, and a global spatial attention module connected in sequence. The 3D convolution module is used to perform dimensionality compression on the data from the input layer. The spatial branch residual block includes two stacked blocks for performing pixel spatial residual feature extraction on the output of the 3D convolution module to obtain a spatial residual feature map.
[0045] The global spatial attention module is used to extract the global spatial attention features of the spatial residual feature map and obtain the global spatial attention feature map;
[0046] The channel splicing module is used to perform channel splicing on the global spectral attention feature map and the global spatial attention feature map to obtain the spatial-spectral embedding feature map;
[0047] The output layer is used to output the spatial-spectral embedding feature map to obtain support set features, query set features, or mask set features.
[0048] Furthermore, the multi-scale spectral convolution module includes a first 3D convolution layer with residual connection, a channel splicing block, and a first convolution block, a second convolution block, a third convolution block, and a fourth convolution block with different convolution kernel parameters;
[0049] The first convolution block, the second convolution block, the third convolution block and the fourth convolution block are all provided with a second 3D convolution layer, a normalization layer and a nonlinear activation layer;
[0050] The second 3D convolutional layer is used to extract multi-scale spectral features from its input;
[0051] The normalization layer is used to normalize the output of the second 3D convolutional layer;
[0052] The nonlinear activation layer is used to perform nonlinear activation operations on the output of the normalization layer;
[0053] The first 3D convolution layer is used to perform a 3D convolution operation on the data from the input layer;
[0054] The first convolution block is used to extract spectral features from the data from the input layer to obtain a first feature map;
[0055] The second convolution block is used to extract spectral features from the first feature map to obtain a second feature map;
[0056] The third convolution block is used to extract spectral features from the second feature map to obtain a third feature map;
[0057] The channel splicing block is used to perform channel splicing on the output of the first 3D convolution layer, the output of the first convolution block, the output of the second convolution block, and the output of the third convolution block to obtain a spliced feature map;
[0058] The fourth convolution block is used to extract spectral features from the spliced feature map to obtain a spectral feature map.
[0059] Furthermore, in S44, the total loss function of the model is constructed based on the support set features, query set features and mask set features through the collaborative domain alignment module, specifically including
[0060] S441: Construct FSL functions based on class prototypes of support set features;
[0061] And the FSL function L fsl The construction formula is
[0062]
[0063] Where: Represents the support set prototype of category k in the hyperspectral image dataset and the query set Q m The loss of matching classification prediction with mask features; k represents the category in the hyperspectral impact dataset; S k represents the support set of category k in the hyperspectral image dataset and S k ={(x1,y1),...,(x i ,y i ),...,(x N ,y N )};x i represents the samples that make up the support set; y i Represents x i corresponding categories; N represents the number of samples in the support set; d(·) represents the Euclidean metric symbol; c k represents the class prototype of the kth class in the support set; c trepresents the number of classes in each support set; G(·) represents a two-branch convolutional neural network module; x j represents the mask set sample; y j,pred Represents x j The predicted label of y true Represents x j The true label of ;γ represents the design parameter;
[0064] S442: Construct a contrast loss function based on the support set features and the mask set features;
[0065] And the contrast loss function L con The expression is
[0066]
[0067] Where: N represents the number of samples in the current batch; z represents the support set features and mask set features Feature samples obtained by splicing according to the 0th dimension; axis represents the splicing dimension; z i represents the normalized feature of the i-th feature sample; z j Represents the normalized features of the jth feature sample; represents the set of samples that share the same category label as the i-th feature sample; t represents the temperature parameter; T represents transposition;
[0068] S443: Constructing a conditional adversarial domain adaptation loss function based on support set features and query set features;
[0069] And the conditional adversarial domain adaptation loss function The construction formula is
[0070]
[0071] Where: L adv represents the loss value of the conditional adversarial domain adaptation loss function; D, G, C represent the domain discriminator, feature extractor, and linear classifier respectively; P s (x),P t (x) represents the data distribution of the source domain and the target domain; D(,) represents the predicted probability that x comes from the source domain; 1-D(,) represents the predicted probability that x comes from the target domain; s i Represents the support set features of the i-th source domain dataset after extraction and mask set features The feature samples obtained by splicing according to the 0th dimension and s i ∈s; G(s i ) indicates that s is trained through a dual-branch convolutional neural network module. i Sample features for spectral feature extraction; Represents the sample features of the source domain dataset s i The classification probability of t j Represents the jth support set feature of the target domain dataset extracted from the target domain and mask set features The feature samples obtained by splicing according to the 0th dimension and t j ∈t; G(t j ) indicates that through the dual-branch convolutional neural network module, t j Sample features for spectral feature extraction; Represents the sample feature t of the target domain dataset j The classification probability of Represents the data distribution P from the source domain s The sample s drawn from (x) i The mathematical expectation of log loss; Represents the data distribution P from the target domain t The sample t drawn from (x) j The mathematical expectation of log loss;
[0072] S444: Based on steps S441 to S443, a total loss function L of the model is constructed, which is expressed as
[0073] L=λ1L fsl +λ2L adv +λ3L con
[0074] Where: λ1, λ2, λ3 represent the weighting coefficients used to balance the contribution of the classification task and the domain alignment task.
[0075] Beneficial effects: The present invention provides a cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention, and performs cross-domain few-sample hyperspectral image classification based on a mask-guided causal intervention meta-learning network. By selecting a support set and a query set, a large number of labeled samples are avoided, and a new mask set is generated by intervening on each sample of the query set through a constructed random mask image mixing module; the support set, the query set and the mask set are combined to perform model training on a hyperspectral image classification model constructed based on causal meta-learning training, and more representative class discrimination causal features are obtained through network learning, and the false correlation between classes and labels is reduced, thereby greatly improving the classification accuracy of hyperspectral images, and the optimal hyperspectral image classification model obtained through training has important application value in aspects such as fine classification of cross-domain hyperspectral image objects, and its use of the technology of mask-guided causal intervention meta-learning network to classify hyperspectral images has important theoretical significance. BRIEF DESCRIPTION OF THE DRAWINGS
[0076] In order to more clearly illustrate the embodiments of the present invention or the technical solutions in the prior art, the following is a brief introduction to the drawings required for use in the embodiments or the description of the prior art. Obviously, the drawings described below are some embodiments of the present invention. For ordinary technicians in this field, other drawings can be obtained based on these drawings without paying any creative labor.
[0077] Figure 1 This is a flow chart of the cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention of the present invention;
[0078] Figure 2 Flowchart of the random mask image mixing module in this embodiment;
[0079] Figure 3 Schematic diagram of a dual-branch convolutional neural network with multi-scale spectral convolution in this embodiment;
[0080] Figure 4 Schematic diagram of the structure of the multi-scale spectral convolution block in this embodiment;
[0081] Figure 5 is a pseudo-color image of the test data set in this embodiment;
[0082] Figure 6 This is a simulation diagram of the classification results of the test data set in this embodiment. DETAILED DESCRIPTION
[0083] To make the objectives, technical solutions, and advantages of the embodiments of the present invention more clear, the technical solutions in the embodiments of the present invention will be clearly and completely described below in conjunction with the accompanying drawings of the embodiments of the present invention. Obviously, the described embodiments are only part of the embodiments of the present invention, not all of the embodiments. All other embodiments obtained by ordinary technicians in this field based on the embodiments of the present invention without making any creative efforts shall fall within the scope of protection of the present invention.
[0084] This embodiment provides a cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention, such as Figure 1 As shown, the specific steps include:
[0085] S1: Obtain a target domain dataset of hyperspectral scene images and a source domain dataset that has been labeled with categories;
[0086] And the target domain dataset includes unlabeled data samples and labeled data samples;
[0087] Using unlabeled data samples as a test set, using the source domain dataset and labeled data samples as a training set, and randomly selecting a support set and a query set from the training set;
[0088] Specifically, the hyperspectral source scene image in this embodiment is from the Chikusei dataset captured by the Headwall Hyperspec VNIR-C imaging sensor, and after processing, the Chikusei dataset has 128 bands. The hyperspectral target scene image in this embodiment is from the Pavia University (UP) dataset captured by the Reflection Optical System Imaging Spectrometer (ROSIS) sensor, and after processing, the UP dataset has 103 bands. In order to reduce the amount of calculation and retain the maximum spatial-spectral characteristics, the spectral channels of the dataset are reduced to 100 bands by using a preset dimension unification module. The Chikusei dataset contains a total of 19 determined ground object categories, as shown in Table 1; the UP dataset contains a total of 9 determined ground object categories, as shown in Table 2:
[0089] Table 1. Number of samples of target object categories in the Chikusei dataset
[0090]
[0091]
[0092] Table 2. Number of samples of target object categories in the UP dataset
[0093]
[0094] This embodiment adopts a meta-learning training strategy: that is, samples of target object categories in the dataset are selected from the hyperspectral scene image to form a training set. Each training run randomly selects the same number of categories as the target domain categories, randomly selects 5 samples from each category to form a support set, and randomly selects 5 samples from the remaining samples to form a query set.
[0095] S2: Preprocess the support set and query set to obtain the optimized support set and optimized query set respectively;
[0096] Specifically, the pretreatment method includes:
[0097] Performing Gaussian noise processing on the hyperspectral scene image with sample deviation caused by scarcity of labeled samples to obtain hyperspectral scene image samples;
[0098] And the expression for Gaussian noise processing is
[0099]
[0100] Where: represents a hyperspectral scene image sample; x represents a hyperspectral scene image with sample bias caused by scarcity of labeled samples; α represents a random sampling factor; β represents an empirical value used to control the noise intensity to match the signal-to-noise ratio of the hyperspectral data;
[0101] Perform dimensionality unification processing on hyperspectral scene image samples;
[0102] S3: Randomly divide the optimized support set into a training support set and a validation support set;
[0103] Randomly divide the optimized query set into a training query set and a validation query set;
[0104] S4: Based on the training support set and the training query set, the constructed hyperspectral image classification model is trained to obtain the optimal hyperspectral image classification model;
[0105] The hyperspectral image classification model consists of a sequentially connected random mask image mixing module, a two-branch convolutional neural network module with multi-scale spectral convolution blocks, a linear classifier, and a collaborative domain alignment module;
[0106] The model training is specifically as follows:
[0107] S41: Through the random mask image mixing module, mask-guided causal intervention is performed on the training query set to obtain the counterfactual mask set;
[0108] In a specific embodiment, Figure 2 As shown, the method for obtaining the counterfactual mask set is as follows:
[0109] S411: traverse each image sample in the training query set in sequence, and take any image sample in the training query set as a target sample;
[0110] The remaining image samples in the training query set except the target samples are used as interference samples;
[0111] The interference sample and the target sample are data samples that belong to the same source domain dataset or target domain dataset but have different image categories;
[0112] S412: Divide the target sample into a grid to obtain an n*n grid image, preferably a 9*9 grid image;
[0113] Define the central pixel grid area b*b of the grid image as a constant pixel value, preferably the central pixel grid area b*b is 3*3 and the central 3*3 area is defined as always 1 (indicating no interference), and perform a random mask operation on the other grid areas except the central pixel grid area to obtain a mask matrix image;
[0114] Specifically, the random mask image mixing module obtains a random mask matrix based on a preset random mask generator.
[0115] S413: Perform an inverse mask operation on the interference sample to obtain an inverse mask matrix image;
[0116] And obtain the counterfactual mask set according to the mask matrix image and the inverse mask matrix image;
[0117] And the formula for obtaining the counterfactual mask set is
[0118] x m =x q ⊙m+x i ⊙(1-m)
[0119]
[0120] Where: x m represents the counterfactual image sample; x q Indicates that each query sample in the query set is the target sample; x i Represents interference samples of the same domain as the input sample but of different categories; ⊙ represents the element-by-element product operator; m represents a random mask matrix generated by a preset random mask generator; i, j represent the position coordinates in the random mask matrix; center represents a fixed central pixel grid area (e.g., the central 3*3 area); r represents a value randomly sampled from a uniform distribution in [0,1]; δ represents the pixel interference ratio;
[0121] S42: Using a dual-branch convolutional neural network module, pixel features are extracted from the training query set, the training support set, and the counterfactual mask set to obtain support set features, query set features, and mask set features; and a class prototype of the support set features is obtained based on the support set features.
[0122] And the formula for obtaining the class prototype of the support set feature is
[0123]
[0124] Where: k represents the category in the hyperspectral impact dataset; S k represents the support set of category k in the hyperspectral image dataset and S k ={(x1,y1),...,(x i ,y i ),...,(x N ,y N )};x i represents the samples that make up the support set; y i Represents x i The corresponding category; N represents the number of support set samples; G(·) represents the two-branch convolutional neural network module;
[0125] Specifically, if Figure 3 The dual-branch convolutional neural network module shown includes an input layer, a spectral feature extraction branch, a spatial feature extraction branch, a channel splicing module, and an output layer;
[0126] The input layer is used to transmit the training query set or the training support set or the counterfactual mask set to the spectral feature extraction branch and the spatial feature extraction branch respectively;
[0127] The spectral feature extraction branch includes a multi-scale spectral convolution module and a global spectral attention module connected in sequence. The multi-scale spectral convolution module is used to extract the multi-scale spectral features of the input data and obtain the spectral feature map; among them, Figure 4 The multi-scale spectral convolution module shown includes a first 3D convolution layer with residual connection, a channel splicing block, and a first convolution block, a second convolution block, a third convolution block, and a fourth convolution block with different convolution kernel parameters;
[0128] The first convolution block, the second convolution block, the third convolution block and the fourth convolution block are all provided with a second 3D convolution layer, a normalization layer and a nonlinear activation layer;
[0129] The second 3D convolutional layer is used to extract multi-scale spectral features from its input;
[0130] The normalization layer is used to normalize the output of the second 3D convolutional layer;
[0131] The nonlinear activation layer is used to perform nonlinear activation operations on the output of the normalization layer;
[0132] The first 3D convolution layer is used to perform a 3D convolution operation on the data from the input layer;
[0133] The first convolution block is used to extract spectral features from the data from the input layer to obtain a first feature map;
[0134] The second convolution block is used to extract spectral features from the first feature map to obtain a second feature map;
[0135] The third convolution block is used to extract spectral features from the second feature map to obtain a third feature map;
[0136] The channel splicing block is used to perform channel splicing on the output of the first 3D convolution layer, the output of the first convolution block, the output of the second convolution block, and the output of the third convolution block to obtain a spliced feature map;
[0137] The fourth convolution block is used to extract spectral features from the spliced feature map to obtain a multi-scale spectral feature map;
[0138] The global spectral attention module is used to extract the global spectral attention features of the spectral feature map and obtain the global spectral attention feature map;
[0139] The spatial feature extraction branch includes a 3D convolution module, a spatial branch residual block, and a global spatial attention module connected in sequence. The 3D convolution module is used to perform dimensionality compression on the data from the input layer. The spatial branch residual block includes two stacked blocks for performing pixel spatial residual feature extraction on the output of the 3D convolution module to obtain a spatial residual feature map.
[0140] The global spatial attention module is used to extract the global spatial attention features of the spatial residual feature map and obtain the global spatial attention feature map;
[0141] The channel splicing module is used to perform channel splicing on the global spectral attention feature map and the global spatial attention feature map to obtain the spatial-spectral embedding feature map;
[0142] The output layer is used to output the spatial-spectral embedding feature map to obtain support set features, query set features, or mask set features;
[0143] The network structure adopted by the dual-branch convolutional neural network module described in this embodiment consists of a spatial branch and a spectral branch, wherein the spectral branch contains a multi-scale spectral convolution block and a global spectral attention block, and the spatial branch contains a residual block and a global spatial attention block. Table 3 shows the specific network structure. The network structure model is as follows: Figure 3 As shown;
[0144] Table 3. Network structure of dual-branch convolutional neural network
[0145]
[0146]
[0147] S43: Using the Euclidean distance metric method, according to the distance between the class prototypes of the query set features and the support set features, and using a linear classifier to obtain the classification probability of the class prediction result, the class prediction result of the query sample is obtained, and then the trained hyperspectral image classification model is obtained;
[0148] S44: The total loss function of the model is constructed based on the support set features, query set features, and mask set features through the collaborative domain alignment module, and the trained hyperspectral image classification model is validated based on the validation support set and validation query set to determine whether the output of the trained hyperspectral image classification model converges.
[0149] Specifically, through the collaborative domain alignment module, based on the total loss function of the model constructed by the support set features, query set features and mask set features, it specifically includes
[0150] S441: Construct FSL functions based on class prototypes of support set features;
[0151] And the FSL function L fsl The construction formula is
[0152]
[0153] Where: Represents the support set class prototype and mask set Q of category k in the hyperspectral image dataset m The loss of matching classification prediction with mask features; d(·) represents the Euclidean metric; c k represents the class prototype of the kth class in the support set; c t represents the number of classes in each support set; G(·) represents a two-branch convolutional neural network module; x j represents the mask set sample; y j,pred Represents x j The predicted label of y true Represents x j The true label of ;γ represents the design parameter;
[0154] S442: In the domain alignment stage, the collaborative domain alignment module aims to strengthen the importance of causal factors in class discrimination and alleviate the domain shift problem by combining conditional adversarial domain adaptation and contrastive learning. This module adopts a two-branch strategy, applying contrastive learning between the support set and the mask set to reduce the distance between positive and negative pairs and increase the distance between positive and negative pairs, thereby increasing the importance of class discriminative features for classification, while applying conditional adversarial domain adaptation between the support set and the query set to alleviate the domain shift problem. The collaborative domain alignment module adopts a two-branch strategy, namely, by applying contrastive learning between the support set and the mask set, and applying conditional adversarial domain adaptation between the support set and the query set;
[0155] Among them, the contrast loss function is constructed based on the support set features and the mask set features;
[0156] And the contrast loss function L con The expression is
[0157]
[0158] Where: N represents the number of samples in the current batch; z represents the support set features and mask set features Feature samples obtained by splicing according to the 0th dimension; axis represents the splicing dimension; z i represents the normalized feature of the i-th feature sample; z j Represents the normalized features of the jth feature sample; represents the set of samples that share the same category label as the i-th feature sample; t represents the temperature parameter; T represents transposition;
[0159] S443: Constructing a conditional adversarial domain adaptation loss function based on support set features and query set features;
[0160] And the conditional adversarial domain adaptation loss function The construction formula is
[0161]
[0162] Where: L adv represents the loss value of the conditional adversarial domain adaptation loss function; D, G, C represent the domain discriminator, feature extractor, and linear classifier respectively; P s (x),P t (x) represents the data distribution of the source domain and the target domain; D(,) represents the predicted probability of x coming from the source domain; 1-D(,) represents the predicted probability of x coming from the target domain. Through this loss function, the domain identification error function can be minimized, and on the feature extractor G and the classifier C, L adv maximized;s i Represents the support set features of the i-th source domain dataset after extraction and mask set features The feature samples obtained by splicing according to the 0th dimension and s i ∈s; G(s i ) indicates that s is trained through a dual-branch convolutional neural network module. i Sample features for spectral feature extraction; Represents the sample features of the source domain dataset s i The classification probability of t j Represents the jth support set feature of the target domain dataset extracted from the target domain With mask set features The feature samples obtained by splicing according to the 0th dimension and t j ∈t; G(t j ) indicates that through the dual-branch convolutional neural network module, t j Sample features for spectral feature extraction; Represents the sample feature t of the target domain dataset j The classification probability of Represents the data distribution P from the source domain s The sample s drawn from (x) i The mathematical expectation of log loss; Represents the data distribution P from the target domain t The sample t drawn from (x) j The mathematical expectation of log loss;
[0163] S444: Based on steps S441 to S443, a total loss function L of the model is constructed, which is expressed as
[0164] L=λ1L fsl +λ2L adv +λ3L con
[0165] Where: λ1, λ2, λ3 represent the weighting coefficients used to balance the contribution of the classification task and the domain alignment task;
[0166] If yes, then confirm that the trained hyperspectral image classification model is the optimal hyperspectral image classification model; otherwise, adaptively adjust the weight parameters of the trained hyperspectral image classification model based on the back propagation method, and repeat step S41;
[0167] S5: Based on the optimal hyperspectral image classification model, hyperspectral image classification of unlabeled data samples in the test set is implemented; this embodiment also includes classification results of hyperspectral image classification of the test set based on the optimal hyperspectral image classification model, and calculating the overall accuracy, average accuracy and Kappa coefficient of the prediction based on the classification results of the hyperspectral image to prove the model performance of the optimal hyperspectral image classification model;
[0168] The expression of the overall accuracy predicted by the calculation model is:
[0169]
[0170] Where: N i,i represents the number of pixels correctly classified into category i; N represents the total number of pixels in the hyperspectral image in the test set; C represents the total number of categories; OA represents the overall accuracy;
[0171] The expression of the average classification accuracy predicted by the calculation model is:
[0172]
[0173] Where: N i+ represents the total number of true pixels of category i, and AA represents the average classification accuracy, which can better reflect the performance of the model on minority categories;
[0174] The expression for calculating the Kappa coefficient is:
[0175]
[0176] Where: N +i represents the total number of pixels predicted to be of category i; and a higher Kappa coefficient value indicates a better classification result. In this embodiment, an example experiment was conducted on the UP dataset using the method described in this embodiment, and the experimental results are shown in Table 4:
[0177] Table 4. UP classification accuracy (%)
[0178]
[0179]
[0180] Among them, OA (Overall Accuracy) represents the overall classification accuracy, AA (Average accuracy) represents the average classification accuracy, and Kappa represents the Kappa coefficient. The Kappa coefficient refers to a multivariate discrete method for evaluating the classification accuracy and error matrix of remote sensing images. It also considers various missed and misclassified pixels outside the diagonal, and can punish the bias of the model, thereby more comprehensively evaluating the classification effect. Figures 5 and 6 The figure shows the pseudo-color image and classification result image of the test dataset.
[0181] In order to more objectively evaluate the role of each step in the mask-guided causal intervention meta-learning network for cross-domain few-shot hyperspectral image classification in this embodiment, an existing ablation experiment is added for illustration. On the basis of the ordinary prototype network, a single module or a combination of different modules is added to compare the experimental results. The specific experimental results are shown in Table 5:
[0182] Table 5. Classification accuracy of different modules (%)
[0183]
[0184] The following conclusions can be drawn from the above ablation experiments:
[0185] (1) The experimental results in Table 5 show that the proposed mask-guided causal intervention meta-learning network for cross-domain few-sample hyperspectral image classification has a good classification effect, which proves that the method described in this embodiment has excellent performance in cross-domain small-sample classification.
[0186] (2) The ablation experiment data in Table 5 show that the classification results of adding the mask-guided causal intervention module are significantly better than the classification results of using only the ordinary meta-learning network, which proves that this module promotes the effectiveness of the network in extracting class discriminant features and shows more robust performance.
[0187] (3) The ablation experiment data in Table 5 show that the classification results of adding the collaborative domain alignment module are significantly better than the results of jointly using meta-learning and mask-guided causal intervention modules, which proves that the collaborative domain alignment module further enhances the importance of class discriminant features in classification, while alleviating domain shift and strengthening the cross-domain generalization of the model, thereby improving the overall classification effect of the model.
[0188] Finally, it should be noted that the above embodiments are only used to illustrate the technical solutions of the present invention, rather than to limit it. Although the present invention has been described in detail with reference to the above embodiments, those skilled in the art should understand that they can still modify the technical solutions described in the above embodiments, or replace some or all of the technical features therein with equivalents. However, these modifications or replacements do not cause the essence of the corresponding technical solutions to deviate from the scope of the technical solutions of the embodiments of the present invention.
Claims
1. A cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention, characterized by: The specific steps include: S1: Obtain a target domain dataset of hyperspectral scene images and a source domain dataset that has been labeled with categories; And the target domain dataset includes unlabeled data samples and labeled data samples; Using unlabeled data samples as a test set, using the source domain dataset and labeled data samples as a training set, and randomly selecting a support set and a query set from the training set; S2: Preprocess the support set and query set to obtain the optimized support set and optimized query set respectively; S3: Randomly divide the optimized support set into a training support set and a validation support set; Randomly divide the optimized query set into a training query set and a validation query set; S4: Based on the training support set and the training query set, the constructed hyperspectral image classification model is trained to obtain the optimal hyperspectral image classification model; and the constructed hyperspectral image classification model includes a random mask image mixing module connected in sequence, a two-branch convolutional neural network module with a multi-scale spectral convolution block, a linear classifier, and a collaborative domain alignment module; The model training is specifically as follows: S41: Through the random mask image mixing module, mask-guided causal intervention is performed on the training query set to obtain the counterfactual mask set; S42: Using a dual-branch convolutional neural network module, pixel features are extracted from the training query set, the training support set, and the counterfactual mask set to obtain support set features, query set features, and mask set features. And obtain the class prototype of the support set features according to the support set features; S43: Using the Euclidean distance metric method, according to the distance between the class prototypes of the query set features and the support set features, and using a linear classifier to obtain the classification probability of the class prediction result, the class prediction result of the query sample is obtained, and then the trained hyperspectral image classification model is obtained; S44: Through the collaborative domain alignment module, based on the total model loss function constructed by the support set features, the query set features, and the mask set features, and based on the validation support set and the validation query set, the trained hyperspectral image classification model is validated to determine whether the output of the trained hyperspectral image classification model converges; If yes, then confirm that the trained hyperspectral image classification model is the optimal hyperspectral image classification model; otherwise, adaptively adjust the weight parameters of the trained hyperspectral image classification model based on the back propagation method, and repeat step S41; S5: Based on the optimal hyperspectral image classification model, hyperspectral image classification of unlabeled data samples in the test set is realized.
2. The cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention according to claim 1 is characterized in that: The pre-processing method in S2 specifically includes: Performing Gaussian noise processing on the hyperspectral scene image with sample deviation caused by scarcity of labeled samples to obtain hyperspectral scene image samples; And the expression for Gaussian noise processing is Where: represents a hyperspectral scene image sample; x represents a hyperspectral scene image with sample bias caused by scarcity of labeled samples; α represents a random sampling factor; β represents an empirical value used to control the noise intensity to match the signal-to-noise ratio of the hyperspectral data; Perform dimension unification processing on hyperspectral scene image samples.
3. The cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention according to claim 2 is characterized in that: The method for obtaining the counterfactual mask set in S41 is specifically as follows: S411: traverse each image sample in the training query set in sequence, and take any image sample in the training query set as a target sample; The remaining image samples in the training query set except the target samples are used as interference samples; The interference sample and the target sample are data samples that belong to the same source domain dataset or target domain dataset but have different image categories; S412: Grid-dividing the target sample to obtain an n*n grid image; Define the central pixel grid area b*b of the grid image as a constant pixel value, and perform a random mask operation on the other grid areas except the central pixel grid area to obtain a mask matrix image; S413: Perform an inverse mask operation on the interference sample to obtain an inverse mask matrix image; And obtain the counterfactual mask set according to the mask matrix image and the inverse mask matrix image; And the formula for obtaining the counterfactual mask set is x m =x q ⊙m+x i ⊙(1-m) Where: x m represents the counterfactual image sample; x q Indicates that each query sample in the query set is the target sample; x i Represents interference samples of the same domain as the input samples but of different categories; ⊙ represents the element-by-element product operator; m represents a random mask matrix generated by a preset random mask generator; i, j represent the position coordinates in the random mask matrix; center represents a fixed central pixel grid area; r represents a value randomly sampled from a uniform distribution of [0,1]; δ represents the pixel interference ratio.
4. The cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention according to claim 3 is characterized in that: The dual-branch convolutional neural network module in S42 includes an input layer, a spectral feature extraction branch, a spatial feature extraction branch, a channel splicing module, and an output layer; The input layer is used to transmit the training query set or the training support set or the counterfactual mask set to the spectral feature extraction branch and the spatial feature extraction branch respectively; The spectral feature extraction branch includes a multi-scale spectral convolution module and a global spectral attention module connected in sequence. The multi-scale spectral convolution module is used to extract the multi-scale spectral features of the input data and obtain a spectral feature map. The global spectral attention module is used to extract the global spectral attention features of the spectral feature map and obtain the global spectral attention feature map; The spatial feature extraction branch includes a 3D convolution module, a spatial branch residual block, and a global spatial attention module connected in sequence. The 3D convolution module is used to perform dimensionality compression on the data from the input layer. The spatial branch residual block includes two stacked blocks for performing pixel spatial residual feature extraction on the output of the 3D convolution module to obtain a spatial residual feature map. The global spatial attention module is used to extract the global spatial attention features of the spatial residual feature map and obtain the global spatial attention feature map; The channel splicing module is used to perform channel splicing on the global spectral attention feature map and the global spatial attention feature map to obtain the spatial-spectral embedding feature map; The output layer is used to output the spatial-spectral embedding feature map to obtain support set features, query set features, or mask set features.
5. The cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention according to claim 4 is characterized in that: The multi-scale spectral convolution module includes a first 3D convolution layer with residual connection, a channel splicing block, and a first convolution block, a second convolution block, a third convolution block, and a fourth convolution block with different convolution kernel parameters; The first convolution block, the second convolution block, the third convolution block and the fourth convolution block are all provided with a second 3D convolution layer, a normalization layer and a nonlinear activation layer; The second 3D convolutional layer is used to extract multi-scale spectral features from its input; The normalization layer is used to normalize the output of the second 3D convolutional layer; The nonlinear activation layer is used to perform nonlinear activation operations on the output of the normalization layer; The first 3D convolution layer is used to perform a 3D convolution operation on the data from the input layer; The first convolution block is used to extract spectral features from the data from the input layer to obtain a first feature map; The second convolution block is used to extract spectral features from the first feature map to obtain a second feature map; The third convolution block is used to extract spectral features from the second feature map to obtain a third feature map; The channel splicing block is used to perform channel splicing on the output of the first 3D convolution layer, the output of the first convolution block, the output of the second convolution block, and the output of the third convolution block to obtain a spliced feature map; The fourth convolution block is used to extract spectral features from the spliced feature map to obtain a spectral feature map.
6. The cross-domain few-sample hyperspectral image classification method based on mask-guided causal intervention according to claim 5 is characterized in that: In S44, the total loss function of the model is constructed based on the support set features, query set features, and mask set features through the collaborative domain alignment module. Specifically, S441: Construct FSL functions based on class prototypes of support set features; And the FSL function L fsl The construction formula is Where: Represents the support set class prototype and mask set Q of category k in the hyperspectral image dataset m The loss of matching classification prediction with mask features; k represents the category in the hyperspectral impact dataset; S k represents the support set of category k in the hyperspectral image dataset and S k ={(x1,y1),...,(x i ,y i ),...,(x N ,y N )};x i represents the samples that make up the support set; y i Represents x i corresponding categories; N represents the number of samples in the support set; d(·) represents the Euclidean metric symbol; c k represents the class prototype of the kth class in the support set; c t represents the number of classes in each support set; G(·) represents a two-branch convolutional neural network module; x j represents the mask set sample; y j,pred Represents x j The predicted label of y true Represents x j The true label of ;γ represents the design parameter; S442: Construct a contrast loss function based on the support set features and the mask set features; And the contrast loss function L con The expression is Where: N represents the number of samples in the current batch; z represents the support set features and mask set features Feature samples obtained by splicing according to the 0th dimension; axis represents the splicing dimension; z i represents the normalized feature of the i-th feature sample; z j Represents the normalized features of the jth feature sample; represents the set of samples that share the same category label as the i-th feature sample; t represents the temperature parameter; T represents transposition; S443: Constructing a conditional adversarial domain adaptation loss function based on support set features and query set features; And the conditional adversarial domain adaptation loss function The construction formula is Where: L adv represents the loss value of the conditional adversarial domain adaptation loss function; D, G, C represent the domain discriminator, feature extractor, and linear classifier respectively; P s (x),P t (x) represents the data distribution of the source domain and the target domain; D(,) represents the predicted probability that x comes from the source domain; 1-D(,) represents the predicted probability that x comes from the target domain; s i Represents the support set features of the i-th source domain dataset after extraction and mask set features The feature samples obtained by splicing according to the 0th dimension and s i ∈s; G(s i ) indicates that s is trained through a dual-branch convolutional neural network module. i Sample features for spectral feature extraction; Represents the sample features of the source domain dataset s i The classification probability of t j Represents the jth support set feature of the target domain dataset extracted from the target domain and mask set features The feature samples obtained by splicing according to the 0th dimension and t j ∈t; G(t j ) indicates that through the dual-branch convolutional neural network module, t j Sample features for spectral feature extraction; Represents the sample feature t of the target domain dataset j The classification probability of Represents the data distribution P from the source domain s The sample s drawn from (x) i The mathematical expectation of log loss; Represents the data distribution P from the target domain t The sample t drawn from (x) j The mathematical expectation of log loss; S444: Based on steps S441 to S443, a total loss function L of the model is constructed, which is expressed as L=λ1L fsl +λ2L adv +λ3L con Where: λ1, λ2, λ3 represent the weighting coefficients used to balance the contribution of the classification task and the domain alignment task.