Classification method of spiking neural network based on 3D discrete wavelet transform
By introducing 3D discrete wavelet transformation and self-attention layer into the pulsed neural network, the problem of insufficient timing dependence and local feature extraction capabilities in the prior art is solved, and more efficient feature extraction and accurate classification results are achieved.
Patent Information
- Application Number
- CN202510040144.8
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-01-10
- Publication Date
- 2025-06-03
AI Technical Summary
Existing pulsed neural networks have shortcomings in capturing timing dependencies and local spatial features, resulting in poor performance in dynamic vision sensor data and computer vision tasks.
Using a pulse neural network model based on 3D discrete wavelet transformation, sample features are decoupled in the space-time dimension through multi-stage 3D discrete wavelet transformation, and global and local features are extracted in combination with the self-attention layer.
Effectively capture time-series dependencies and local detailed characteristics, improve the accuracy of classification results, reduce the amount of model parameters, and is suitable for resource-constrained environments.
Smart Images

Figure BDA0005236691350000101 
Figure BDA0005236691350000111 
Figure BDA0005236691350000112
Abstract
Description
Technical Field
[0001] The present invention relates to a classification method for spiking neural networks, in particular to a classification method for spiking neural networks based on 3D discrete wavelet transform. Background Art
[0002] Spiking Neural Networks (SNNs), also known as the third generation of neural networks, have attracted much attention due to their low energy consumption and event-driven characteristics. Different from artificial neural networks (ANNs) that transmit information through continuous real-valued numbers, the design inspiration of spiking neural networks comes from the biological nervous system, which transmits information by using binary pulses.
[0003] However, the existing spiking neural networks still face the following multiple challenges in applications: (1) insufficient ability to capture temporal dependencies. Most spiking neural networks mainly focus on the spatial information at a single time step and ignore the temporal dependencies between multiple time steps, which is particularly crucial for the data of dynamic vision sensors with rich action information; (2) the problem of the number of model parameters. The current spiking neural networks often have a large number of model parameters, making it difficult to deploy in resource-constrained environments. This not only increases the demand for computing resources but also makes it more difficult in practical applications; (3) insufficient local feature extraction ability. Local detailed features are crucial for computer vision tasks, but there are bottlenecks in the existing spiking neural networks when extracting these features. These challenges limit the performance of the current spiking neural networks in various downstream tasks, and new solutions are urgently needed to improve their efficiency and performance. Summary of the Invention
[0004] The technical problem to be solved by the present invention is to provide a spiking neural network that can better capture temporal dependencies and local spatial features.
[0005] The technical solution adopted by the present invention to solve the above technical problem is: a classification method for spiking neural networks based on 3D discrete wavelet transform, including the following steps:
[0006] Step 1): Select N training set data from the original dataset, and after preprocessing the N training set data, obtain N preprocessed training data;
[0007] Step 2): Construct a pulse neural network model to be trained. Randomly shuffle the N preprocessed training data and input it into the pulse neural network model to be trained. The pulse neural network model to be trained includes a convolutional projection layer, a first frequency domain layer based on 3D discrete wavelet transform, a second frequency domain layer based on 3D discrete wavelet transform, a first downsampling layer, a second downsampling layer, a third downsampling layer, a first self-attention layer based on pulse neural network, a second self-attention layer based on pulse neural network, and a classification layer. Randomly shuffle the N preprocessed training data and input it into the convolutional projection layer to obtain downsampled feature embeddings;
[0008] Step 3): Input the downsampled feature embeddings into the first frequency domain layer based on 3D discrete wavelet transform to extract the first feature including temporal dependence and local detail features. Input the first feature into the first downsampling layer for downsampling to reduce the spatial dimension and increase the channel dimension of the first feature, obtaining the first downsampled feature. Input the first downsampled feature into the second frequency domain layer based on 3D discrete wavelet transform to obtain the second feature including temporal dependence and local detail features. Input the second feature into the second downsampling layer for downsampling to reduce the spatial dimension and increase the channel dimension of the second feature, obtaining the second downsampled feature;
[0009] Step 4): Input the second downsampled feature into the first self-attention layer based on pulse neural network to further extract global features by combining the local detail features in the shallow layer of the model, obtaining the third feature. Input the third feature into the third downsampling layer for downsampling to obtain the third downsampled feature. Input the third downsampled feature into the second self-attention layer based on pulse neural network to further extract global features by combining the local detail features in the shallow layer of the model, obtaining the fourth feature;
[0010] Step 5): Input the fourth feature into the classification layer to obtain the predicted classification result of the preprocessed training data;
[0011] Step 6): Define the loss function of the pulse neural network model to be trained, obtain the value of the loss function using the predicted classification result of the preprocessed training data, and update the pulse neural network model to be trained through backpropagation. After training, obtain the trained pulse neural network model;
[0012] Step 7): Select M target samples from the target dataset for preprocessing and form a target dataset to be classified. Input the target dataset to be classified into the trained pulse neural network model to obtain the classification result of the target dataset to be classified, completing the classification process.
[0013] Compared with the prior art, the advantages of the present invention are as follows: The pulse neural network model to be trained mainly introduces multi-level 3D discrete wavelet transform into the pulse neural network, decouples the sample features into multi-scale and multi-directional spectral components in the spatio-temporal dimension, and finally inputs the global features combined with the local detail features extracted from the shallow layer of the model into the classification layer during the training process to obtain the predicted classification result. Subsequently, the pulse neural network model to be trained is updated through backpropagation according to the loss function to obtain the trained pulse neural network model; finally, the classification result of the target dataset to be classified is obtained; the advantage is that by deploying the frequency domain layer based on 3D discrete wavelet transform in the shallow layer of the network and the self-attention layer based on the pulse neural network in the deep layer of the network, the ability to fully extract local detail features and global modeling ability are utilized, making the classification result output by the trained pulse neural network model more accurate.
[0014] Specifically, in step 1), when the original dataset is a dynamic vision sensor dataset, the process of preprocessing the training set data is as follows: Integrate the files of the dynamic vision sensor dataset according to the time window to obtain 8 video frames, then the preprocessed training data includes video data of 8 time steps;
[0015] When the original dataset is a static image dataset, the process of preprocessing the training set data is as follows: Copy a single picture in the static image dataset four times and arrange the copied pictures in sequence to form a picture sequence, and use the formed picture sequence as the preprocessed training data of multiple time steps corresponding to the static image dataset.
[0016] Specifically, in step 3), the first frequency domain layer based on 3D discrete wavelet transform and the second frequency domain layer based on 3D discrete wavelet transform have the same composition. The process of extracting the first feature by embedding the downsampled features into the first frequency domain layer based on 3D discrete wavelet transform is as follows:
[0017] Step 3-1: The first frequency domain layer based on 3D discrete wavelet transform includes a 3D discrete wavelet transform layer, a first spiking neuron, a global feature fusion module, a multi-step feature fusion module, a 3D discrete wavelet inverse transform layer, and a first multi-layer perceptron. First, embed the downsampled features with dimensions of T×L×D into the 3D discrete wavelet transform layer, where T is the time step dimension, L is the spatial dimension, and D is the channel dimension. The 3D discrete wavelet transform layer performs frequency domain conversion on each downsampled feature embedding in the spatial and time dimensions through 3D discrete wavelet transform to obtain a low-frequency component with dimensions of and two high-frequency components of different scales. The dimension of the first high-frequency component is which is denoted as the first high-frequency component, and the dimension of the second high-frequency component is And it is denoted as the second high-frequency component, and R is the direction dimension of the high-frequency component;
[0018] Step 3-2: Input the first high-frequency component and the second high-frequency component into the first spiking neuron respectively, to obtain the first high-frequency component represented by binary pulses and the second high-frequency component represented by binary pulses. Input the low-frequency component into the first spiking neuron to obtain the low-frequency component represented by binary pulses with the same dimension;
[0019] Step 3-3: Input the low-frequency component represented by binary pulses into the global feature fusion module. The global feature fusion module performs element-wise dot multiplication on the low-frequency component represented by binary pulses and the first learnable parameter matrix with dimension to extract the temporal dependence relationship in the low-frequency component, and obtain the low-frequency component after global feature fusion with the same dimension;
[0020] Step 3-4: Input the first high-frequency component represented by binary pulses and the second high-frequency component represented by binary pulses into the multi-step feature fusion module respectively. The multi-step feature fusion module performs multi-dimensional feature update and channel feature fusion on the input first high-frequency component represented by binary pulses to obtain the first fused high-frequency component, and performs multi-dimensional feature update and channel feature fusion on the input second high-frequency component represented by binary pulses to obtain the second fused high-frequency component. The specific process is as follows:
[0021] Step 3-4-1: Multi-dimensional feature update: Perform element-wise dot multiplication on the first high-frequency component represented by binary pulses and the second learnable parameter matrix with dimension in the spatio-temporal dimension and the direction dimension to extract the temporal dependence relationship and local detail features in the high-frequency component, and obtain the first high-frequency component after multi-dimensional feature update;
[0022] Perform element-wise dot multiplication on the second high-frequency component represented by binary pulses and the third learnable parameter matrix with dimension in the spatio-temporal dimension and the direction dimension to extract the temporal dependence relationship and local detail features in the high-frequency component, and obtain the second high-frequency component after multi-dimensional feature update;
[0023] Step 3-4-2: Channel feature fusion: Use Einstein summation convention to perform matrix multiplication on the first high-frequency component after multi-dimensional feature update and the fourth learnable parameter matrix with dimension D×D and sum in the channel dimension to obtain the first fused high-frequency component with the same dimension;
[0024] Use Einstein summation convention to perform matrix multiplication on the second high-frequency component after multi-dimensional feature update and the fifth learnable parameter matrix with dimension D×D and sum in the channel dimension to obtain the second fused high-frequency component with the same dimension;
[0025] Step 3-5: The 3D inverse discrete wavelet transform layer performs 3D inverse discrete wavelet transform on the low-frequency component after global feature fusion, the first fused high-frequency component, and the second fused high-frequency component to obtain a fused feature with dimensions of T×L×D;
[0026] Step 3-6: Input the fused feature into the first multi-layer perceptron to obtain a first feature.
[0027] Specifically, in step 4), the structure of the first self-attention layer based on the spiking neural network is the same as that of the second self-attention layer based on the spiking neural network. Among them, the process of inputting the second downsampled feature into the first self-attention layer based on the spiking neural network and further extracting global features by combining the local detailed features of the shallow layer of the model to obtain the third feature is as follows:
[0028] Step 4-1: The first self-attention layer based on the spiking neural network includes three linear layers, a second spiking neuron, an attention calculation module, and a second multi-layer perceptron. Input the second downsampled feature into the three linear layers respectively to obtain a query vector, a key vector, and a value vector;
[0029] Step 4-2: Input the query vector, the key vector, and the value vector into the second spiking neuron respectively to obtain a query vector represented by binary pulses, a key vector represented by binary pulses, and a value vector represented by binary pulses;
[0030] Step 4-3: Input the query vector represented by binary pulses, the key vector represented by binary pulses, and the value vector represented by binary pulses into the attention calculation module for self-attention calculation to obtain the updated feature of the attention module;
[0031] Step 4-4: Input the updated feature of the attention module into the second multi-layer perceptron to obtain a third feature.
[0032] Specifically, the specific process of step 5) is as follows:
[0033] Step 5-1: The classification layer includes a pooling layer and a classification head. Pass the fourth feature through the pooling layer, and the pooling layer applies global average pooling to obtain a pooled feature vector;
[0034] Step 5-2: Denote the pooled feature vector corresponding to the i-th preprocessed training data as f i , 1 ≤ i ≤ N. Input f i into the classification head to obtain a classification result p i , where the classification result p i represents the predicted probability that the preprocessed training data belongs to each category in the original dataset, Softmax(·) represents the normalized exponential function, represents the weight matrix of the classification head, denotes the transpose of, denotes the bias of the classification head.
[0035] Specifically, the specific process of step 6) is as follows:
[0036] Step 6-1: Define the loss function of the pulse neural network model to be trained as L cls , represents the class label of the i-th preprocessed training data, 1 ≤ t ≤ C, where C represents the total number of classes contained in the N training set data. If the i-th preprocessed training data belongs to the t-th class, then otherwise represents the probability that the i-th training sample output by the pulse neural network model to be trained belongs to the t-th class;
[0037] Step 6-2: Set the maximum number of iterations. According to the loss function, use the AdamW optimization algorithm to iteratively optimize the pulse neural network model to be trained until the set maximum number of iterations is reached, and then stop the iteration process to obtain the trained pulse neural network model.
[0038] Specifically, when the original dataset is the CIFAR10 dataset or the CIFAR100 dataset, the maximum number of iterations is set to 400 times; when the original dataset is the UCF101-DVS dataset or the HMDB51-DVS dataset, the maximum number of iterations is set to 200 times; when the original dataset is the CIFAR10-DVS dataset or the DVS128-Gesture dataset, the maximum number of iterations is set to 100 times; when the original dataset is the Tiny-imagenet dataset, the maximum number of iterations is set to 300 times. Specific embodiments
[0039] The present invention will be further described in detail below in conjunction with embodiments.
[0040] A classification method for a pulse neural network based on 3D discrete wavelet transform, comprising the following steps:
[0041] Step 1): Select N training set data from the original dataset, and preprocess the N training set data to obtain N preprocessed training data; when the original dataset is a dynamic vision sensor dataset, the process of preprocessing the training set data is: Integrate the files of the dynamic vision sensor dataset according to the time window to obtain 8 video frames, then the preprocessed training data includes video data of 8 time steps;
[0042] When the original data set is a static image data set, the process of preprocessing the training set data is as follows: copy a single picture in the static image data set four times and arrange the copied pictures in order to form a picture sequence, and use the formed picture sequence as the preprocessed training data for multiple time steps corresponding to the static image data set.
[0043] Step 2): Construct a pulse neural network model to be trained, randomly shuffle N preprocessed training data and input them into the pulse neural network model to be trained. The pulse neural network model to be trained includes a convolutional projection layer, a first frequency domain layer based on 3D discrete wavelet transform, a second frequency domain layer based on 3D discrete wavelet transform, a first downsampling layer, a second downsampling layer, a third downsampling layer, a first self-attention layer based on pulse neural network, a second self-attention layer based on pulse neural network, and a classification layer. Randomly shuffle N preprocessed training data and input them into the convolutional projection layer to obtain downsampled feature embeddings.
[0044] Step 3): Input the downsampled feature embeddings into the first frequency domain layer based on 3D discrete wavelet transform to extract the first feature including temporal dependence relationship and local detail features. Input the first feature into the first downsampling layer for downsampling to reduce the spatial dimension and increase the channel dimension of the first feature, obtaining the first downsampled feature. Input the first downsampled feature into the second frequency domain layer based on 3D discrete wavelet transform to obtain the second feature including temporal dependence relationship and local detail features. Input the second feature into the second downsampling layer for downsampling to reduce the spatial dimension and increase the channel dimension of the second feature, obtaining the second downsampled feature; in Step 3), the first frequency domain layer based on 3D discrete wavelet transform and the second frequency domain layer based on 3D discrete wavelet transform have the same structure. The process of inputting the downsampled feature embeddings into the first frequency domain layer based on 3D discrete wavelet transform to extract the first feature is as follows:
[0045] Step 3-1: The first frequency domain layer based on 3D discrete wavelet transform includes a 3D discrete wavelet transform layer, a first pulse neuron, a global feature fusion module, a multi-step feature fusion module, a 3D discrete wavelet inverse transform layer, and a first multi-layer perceptron. First, input the downsampled feature embeddings with dimensions of T×L×D into the 3D discrete wavelet transform layer, where T is the time step dimension, L is the spatial dimension, and D is the channel dimension. The 3D discrete wavelet transform layer performs frequency domain conversion on each downsampled feature embedding in the spatial dimension and time dimension through 3D discrete wavelet transform, obtaining a low-frequency component with dimensions of and two high-frequency components with different scales. The dimension of the first high-frequency component is which is denoted as the first high-frequency component, and the dimension of the second high-frequency component is which is denoted as the second high-frequency component, and R is the direction dimension of the high-frequency component.
[0046] Step 3-2: Input the first high-frequency component and the second high-frequency component into the first spiking neuron respectively to obtain the first high-frequency component represented by binary pulses and the second high-frequency component represented by binary pulses. Input the low-frequency component into the first spiking neuron to obtain the low-frequency component represented by binary pulses with unchanged dimension.
[0047] Step 3-3: Input the low-frequency component represented by binary pulses into the global feature fusion module. The global feature fusion module performs element-wise dot multiplication on the low-frequency component represented by binary pulses and the first learnable parameter matrix with dimension to extract the temporal dependence relationship in the low-frequency component, and obtain the low-frequency component after global feature fusion with unchanged dimension.
[0048] Step 3-4: Input the first high-frequency component represented by binary pulses and the second high-frequency component represented by binary pulses into the multi-step feature fusion module respectively. The multi-step feature fusion module performs multi-dimensional feature update and channel feature fusion on the input first high-frequency component represented by binary pulses to obtain the first fused high-frequency component, and performs multi-dimensional feature update and channel feature fusion on the input second high-frequency component represented by binary pulses to obtain the second fused high-frequency component. The specific process is as follows:
[0049] Step 3-4-1: Multi-dimensional feature update: Perform element-wise dot multiplication on the first high-frequency component represented by binary pulses and the second learnable parameter matrix with dimension in the spatio-temporal dimension and the direction dimension to extract the temporal dependence relationship and local detail features in the high-frequency component, and obtain the first high-frequency component after multi-dimensional feature update;
[0050] Perform element-wise dot multiplication on the second high-frequency component represented by binary pulses and the third learnable parameter matrix with dimension in the spatio-temporal dimension and the direction dimension to extract the temporal dependence relationship and local detail features in the high-frequency component, and obtain the second high-frequency component after multi-dimensional feature update;
[0051] Step 3-4-2: Channel feature fusion: Use Einstein summation convention to perform matrix multiplication on the first high-frequency component after multi-dimensional feature update and the fourth learnable parameter matrix with dimension D×D and sum in the channel dimension to obtain the first fused high-frequency component with unchanged dimension;
[0052] Use Einstein summation convention to perform matrix multiplication on the second high-frequency component after multi-dimensional feature update and the fifth learnable parameter matrix with dimension D×D and sum in the channel dimension to obtain the second fused high-frequency component with unchanged dimension.
[0053] Step 3-5: The 3D inverse discrete wavelet transform layer performs 3D inverse discrete wavelet transform on the low-frequency component, the first fused high-frequency component, and the second fused high-frequency component after global feature fusion to obtain a fused feature with dimensions of T×L×D.
[0054] Step 3-6: Input the fused feature into the first multi-layer perceptron to obtain a first feature.
[0055] Step 4): Input the second downsampled feature into the first self-attention layer based on the spiking neural network, and further extract global features by combining the local detailed features of the shallow layer of the model to obtain a third feature. Input the third feature into the third downsampling layer for downsampling to obtain a third downsampled feature. Input the third downsampled feature into the second self-attention layer based on the spiking neural network, and further extract global features by combining the local detailed features of the shallow layer of the model to obtain a fourth feature.
[0056] In Step 4), the first self-attention layer based on the spiking neural network and the second self-attention layer based on the spiking neural network have the same structure. Among them, the specific process of inputting the second downsampled feature into the first self-attention layer based on the spiking neural network and further extracting global features by combining the local detailed features of the shallow layer of the model to obtain a third feature is as follows:
[0057] Step 4-1: The first self-attention layer based on the spiking neural network includes three linear layers, a second spiking neuron, an attention calculation module, and a second multi-layer perceptron. Input the second downsampled feature into the three linear layers respectively to obtain a query vector, a key vector, and a value vector.
[0058] Step 4-2: Input the query vector, the key vector, and the value vector into the second spiking neuron respectively to obtain a query vector represented by binary pulses, a key vector represented by binary pulses, and a value vector represented by binary pulses.
[0059] Step 4-3: Input the query vector represented by binary pulses, the key vector represented by binary pulses, and the value vector represented by binary pulses into the attention calculation module for self-attention calculation to obtain the updated feature of the attention module.
[0060] Step 4-4: Input the updated feature of the attention module into the second multi-layer perceptron to obtain a third feature.
[0061] Step 5): Input the fourth feature into the classification layer to obtain the predicted classification result of the preprocessed training data. The specific process is as follows:
[0062] Step 5-1: The classification layer includes a pooling layer and a classification head. Pass the fourth feature through the pooling layer, and the pooling layer applies global average pooling to obtain a pooled feature vector.
[0063] Step 5-2: Denote the pooled feature vector corresponding to the $i$-th preprocessed training data as $f$ i , where $1\leq i\leq N$. Input $f$ i into the classification head to obtain the classification result $p$ i , where the classification result $p$ i represents the predicted probabilities that the preprocessed training data belong to various classes in the original dataset, Softmax(·) represents the normalized exponential function, represents the weight matrix of the classification head, represents the transpose of, and represents the bias of the classification head.
[0064] Step 6): Define the loss function of the pulse neural network model to be trained. Use the predicted classification results of the preprocessed training data to obtain the value of the loss function, and update the pulse neural network model to be trained through backpropagation. After the training is completed, obtain the trained pulse neural network model. The specific process is as follows:
[0065] Step 6-1: Define the loss function of the pulse neural network model to be trained as $L$ cls , represents the class label of the $i$-th preprocessed training data, where $1\leq t\leq C$, and $C$ represents the total number of classes included in the $N$ training set data. If the $i$-th preprocessed training data belongs to the $t$-th class, then there is Otherwise represents the probability that the $i$-th training sample output by the pulse neural network model to be trained belongs to the $t$-th class; through the constraint of the classification loss, the pulse neural network model can effectively optimize the inter-class distance between different class samples and generate more discriminative features for samples of different classes.
[0066] Step 6-2: Set the maximum number of iterations. According to the loss function, use the AdamW optimization algorithm to iteratively optimize the pulse neural network model to be trained until the set maximum number of iterations is reached, and then stop the iterative process to obtain the trained pulse neural network model. When the original dataset uses the CIFAR10 dataset or the CIFAR100 dataset, the maximum number of iterations is set to 400 times. When the original dataset uses the UCF101-DVS dataset or the HMDB51-DVS dataset, the maximum number of iterations is set to 200 times. When the original dataset uses the CIFAR10-DVS dataset or the DVS128-Gesture dataset, the maximum number of iterations is set to 100 times. When the original dataset uses the Tiny-imagenet dataset, the maximum number of iterations is set to 300 times.
[0067] Step 7): Select M target samples from the target dataset for preprocessing and form the target dataset to be classified. Input the target dataset to be classified into the trained spiking neural network model to obtain the classification result of the target dataset to be classified, thus completing the classification process.
[0068] In the above embodiments, first, by introducing multi-level 3D discrete wavelet transform into the spiking neural network, the sample features are decoupled into multi-scale and multi-directional spectral components in the spatio-temporal dimension, separating the global information of the low-frequency components and the local detail information of the high-frequency components in the sample features, and obtaining the high-frequency components and low-frequency components corresponding to the sample features. According to the respective properties of the low-frequency components and high-frequency components, an independent feature fusion strategy is adopted. For the low-frequency components, global feature fusion is used, and a learnable parameter matrix with the same dimension size is multiplied element-wise with the low-frequency components to obtain the low-frequency components after global feature fusion. For the high-frequency components, step-by-step feature fusion is adopted. First, a learnable parameter matrix is multiplied element-wise with the high-frequency components to obtain the high-frequency components after feature update in multiple dimensions such as spatio-temporal and direction. Then, the Einstein summation convention is used with the learnable parameter matrix and the high-frequency components after multi-dimensional feature update to obtain the high-frequency components after feature fusion in the channel dimension. Each frequency domain layer based on 3D discrete wavelet transform effectively captures the temporal dependence relationship and local detail features, realizing efficient feature extraction of samples. By designing the encoder as a multi-level lightweight structure, deploying the frequency domain module based on 3D discrete wavelet transform in the shallow layer of the encoder network, and deploying the self-attention layer based on the spiking neural network in the deep layer of the encoder network, the ability to extract local detail features and global modeling ability of this structure are fully utilized to achieve more comprehensive feature extraction. At the same time, this multi-level structure combining each frequency domain layer and each self-attention layer can effectively reduce the number of model parameters.
[0069] The following shows the comparison of top-1 classification accuracies of different methods on the CIFAR10-DVS and DVS123-Gesture datasets in Table 1, where the classification method of the spiking neural network based on 3D discrete wavelet transform in this embodiment is abbreviated as this method.
[0070] Table 1
[0071]
[0072] As can be seen from Table 1, this method is better than other models and achieves the best performance on both datasets. It reaches the highest top-1 accuracy of 82.90% on the CIFAR10-DVS dataset, which is 1.5% higher than Spikingformer-CML. It reaches the highest top-1 accuracy of 99.30% on the DVS128-Gesture dataset, which is consistent with the accuracy of the current best method.
[0073] The following shows the comparison of the top-1 accuracy of different methods on the HMDB51-DVS and UCF101-DVS datasets through Table 2. Among the first three methods, t = 240, where t represents the duration of each sample input in milliseconds.
[0074] Table 2
[0075]
[0076] As can be seen from Table 2, this method is higher than the current model on both datasets. It achieved a top-1 accuracy of 59.1% on the HMDB51-DVS dataset, which is 3.5% higher than the current best method SpikePoint. On UCF101-DVS, it achieved a top-1 accuracy of 72.1%, which is 2.6% higher than the current best method SpikePoint.
[0077] The following shows the comparison of the top-1 accuracy of different methods on the CIFAR10 and CIFAR100 datasets through Table 3.
[0078] Table 3
[0079]
[0080] As can be seen from Table 3, this method is higher than the current model on both datasets. It achieved a top-1 accuracy of 96.50% on the CIFAR10 dataset and 81.28% on the CIFAR100 dataset, and the number of parameters is also smaller than other models, with only 6.28M parameter size. At the same time, reducing the model size of this method can also achieve good performance under the condition of a smaller parameter size.
[0081] The following shows the comparison of the top-1 accuracy of different methods on the Tiny-imagenet dataset through Table 4.
[0082] Table 4
[0083]
[0084] As can be seen from Table 4, this method is higher than the current model on the Tiny-imagenet dataset, achieving a top-1 accuracy of 68.26%, and the parameter size is also smaller than other models.
Claims
1. A classification method of spiking neural network based on 3D discrete wavelet transform, characterized in that The following steps are involved: Step 1): select N training set data from the original data set, and preprocess the N training set data to obtain N preprocessed training data; Step 2): construct a pulse neural network model to be trained, randomly shuffle N preprocessed training data and input them into the pulse neural network model to be trained, the pulse neural network model to be trained includes a convolutional projection layer, a first frequency domain layer based on 3D discrete wavelet transform, a second frequency domain layer based on 3D discrete wavelet transform, a first downsampling layer, a second downsampling layer, a third downsampling layer, a first self-attention layer based on a pulse neural network, a second self-attention layer based on a pulse neural network and a classification layer, randomly shuffle N preprocessed training data and input them into the convolutional projection layer to obtain downsampled feature embedding; Step 3): embed the downsampled features into the first frequency domain layer based on 3D discrete wavelet transform, extract the first features including temporal dependency and local detail features, input the first features into the first downsampling layer for downsampling, reduce the spatial dimension of the first features and increase the channel dimension, and obtain the first downsampled features, input the first downsampled features into the second frequency domain layer based on 3D discrete wavelet transform, obtain the second features including temporal dependency and local detail features, input the second features into the second downsampling layer for downsampling, reduce the spatial dimension of the second features and increase the channel dimension, and obtain the second downsampled features; Step 4): Input the second down-sampled feature into the first self-attention layer based on the pulse neural network, combine the local detail features of the shallow layer of the model to further extract the global feature, and obtain the third feature. Input the third feature into the third down-sampled layer for down-sampling to obtain the third down-sampled feature. Input the third down-sampled feature into the second self-attention layer based on the pulse neural network, and combine the local detail features of the shallow layer of the model to further extract the global feature to obtain the fourth feature. Step 5): Input the fourth feature into the classification layer to obtain the predicted classification result of the preprocessed training data; Step 6): define the loss function of the pulse neural network model to be trained, obtain the value of the loss function by using the predicted classification results of the preprocessed training data, update the pulse neural network model to be trained by back propagation, and obtain the trained pulse neural network model after the training is completed; Step 7): Select M target samples from the target data set for preprocessing and form a target data set to be classified, input the target data set to be classified into the trained spiking neural network model, obtain the classification result of the target data set to be classified, and complete the classification process.
2. The classification method of a 3D discrete wavelet transform-based spiking neural network according to claim 1, characterized in that In the step 1), when the original data set is a dynamic visual sensor data set, the process of preprocessing the training set data is: integrating the file of the dynamic visual sensor data set according to the time window to obtain 8 video frames, and the preprocessed training data includes video data of 8 time steps; When the original data set is a static image data set, the process of preprocessing the training set data is: copying a single picture in the static image data set four times and arranging the copied pictures in order to form a picture sequence, and using the formed picture sequence as the preprocessed training data for multiple time steps corresponding to the static image data set.
3. The classification method of a 3D discrete wavelet transform-based spiking neural network according to claim 1, characterized in that In the step 3), the first frequency domain layer based on 3D discrete wavelet transform and the second frequency domain layer based on 3D discrete wavelet transform have the same structure, wherein the downsampled features are embedded into the first frequency domain layer based on 3D discrete wavelet transform, and the process of extracting the first feature is as follows: Step 3-1: The first frequency domain layer based on 3D discrete wavelet transform includes a 3D discrete wavelet transform layer, a first pulse neuron, a global feature fusion module, a multi-step feature fusion module, a 3D discrete wavelet inverse transform layer and a first multilayer perceptron. First, the downsampled features of dimension T×L×D are embedded into the 3D discrete wavelet transform layer, where T is the time step dimension, L is the space dimension, and D is the channel dimension. The 3D discrete wavelet transform layer embeds each downsampled feature in the space dimension and the time dimension through 3D discrete wavelet transform and performs frequency domain conversion to obtain a feature with dimension The low-frequency component and two high-frequency components of different scales, the dimension of the first high-frequency component is And recorded as the first high-frequency component, the dimension of the second high-frequency component is And recorded as the second high-frequency component, R is the directional dimension of the high-frequency component; Step 3-2: input the first high-frequency component and the second high-frequency component into the first pulse neuron respectively to obtain the first high-frequency component represented by binary pulses and the second high-frequency component represented by binary pulses, and input the low-frequency component into the first pulse neuron to obtain the low-frequency component represented by binary pulses with unchanged dimension; Step 3-3: Input the low-frequency component represented by the binary pulse into the global feature fusion module, and the global feature fusion module combines the low-frequency component represented by the binary pulse with the dimension Perform element-by-element dot multiplication on the first learnable parameter matrix to extract the temporal dependency in the low-frequency component and obtain the low-frequency component after fusion of the global features with unchanged dimension; Step 3-4: The first high-frequency component represented by the binary pulse and the second high-frequency component represented by the binary pulse are respectively input into the multi-step feature fusion module. The multi-step feature fusion module performs multi-dimensional feature update and channel feature fusion on the first high-frequency component represented by the input binary pulse to obtain the first fused high-frequency component. The multi-step feature fusion module performs multi-dimensional feature update and channel feature fusion on the second high-frequency component represented by the input binary pulse to obtain the second fused high-frequency component. The specific process is as follows: Step 3-4-1: Multi-dimensional feature update: The first high-frequency component represented by the binary pulse is combined with the dimension The second learnable parameter matrix is element-by-element dot multiplication in the spatiotemporal dimension and the directional dimension to extract the temporal dependency and local detail features in the high-frequency component, and obtain the first high-frequency component after the multi-dimensional feature is updated; The second high-frequency component represented by the binary pulse is The third learnable parameter matrix performs element-by-element point multiplication in the spatiotemporal dimension and the directional dimension to extract the temporal dependency and local detail features in the high-frequency component, and obtains the second high-frequency component after the multi-dimensional feature is updated; Step 3-4-2: Channel feature fusion: Use the Einstein summation convention to perform matrix multiplication on the first high-frequency component after the multi-dimensional feature update and the fourth learnable parameter matrix of dimension D×D and sum them in the channel dimension to obtain the first fused high-frequency component with unchanged dimension; Use the Einstein summation convention to perform matrix multiplication on the second high-frequency component after the multi-dimensional feature update and the fifth learnable parameter matrix of dimension D×D and sum them in the channel dimension to obtain the second fused high-frequency component with unchanged dimension; Step 3-5: The 3D discrete wavelet inverse transform layer performs a 3D discrete wavelet inverse transform on the low-frequency component after global feature fusion, the first fused high-frequency component, and the second fused high-frequency component to obtain a fused feature with a dimension of T×L×D; Step 3-6: Input the fused features into the first multilayer perceptron to obtain the first feature.
4. The classification method of a 3D discrete wavelet transform-based pulse neural network according to claim 3 is characterized in that In the step 4), the first self-attention layer based on the spiking neural network and the second self-attention layer based on the spiking neural network have the same structure, wherein the second down-sampled feature is input into the first self-attention layer based on the spiking neural network, and the global feature is further extracted in combination with the local detail feature of the shallow layer of the model, and the specific process of obtaining the third feature is as follows: Step 4-1: The first self-attention layer based on the spiking neural network includes three linear layers, a second spiking neuron, an attention calculation module, and a second multilayer perceptron. The second down-sampled features are input into the three linear layers respectively to obtain a query vector, a key vector, and a value vector. Step 4-2: input the query vector, the key vector and the value vector into the second pulse neuron respectively, to obtain the query vector represented by binary pulses, the key vector represented by binary pulses and the value vector represented by binary pulses; Step 4-3: Input the query vector represented by the binary pulse, the key vector represented by the binary pulse, and the numerical vector represented by the binary pulse into the attention calculation module to perform self-attention calculation, and obtain the updated features of the attention module; Step 4-4: Input the features updated by the attention module into the second multi-layer perceptron to obtain the third features.
5. The classification method of a 3D discrete wavelet transform-based spiking neural network according to claim 4 is characterized in that The specific process of step 5) is as follows: Step 5-1: The classification layer includes a pooling layer and a classification head. The fourth feature is passed through the pooling layer, and the pooling layer applies global average pooling to obtain a pooled feature vector; Step 5-2: The pooled feature vector corresponding to the i-th preprocessed training data is recorded as f i , 1≤i≤N will f i Input the classification head to get the classification result p i , Among them, the classification result p i represents the predicted probability that the preprocessed training data belongs to each category in the original data set, Softmax(·) represents the normalized exponential function, represents the weight matrix of the classification head, express The transpose of Represents the bias of the classification head.
6. The classification method of a 3D discrete wavelet transform-based spiking neural network according to claim 5, characterized in that The specific process of step 6) is as follows: Step 6-1: Define the loss function of the pulse neural network model to be trained as L cls , represents the category label of the i-th preprocessed training data, 1≤t≤C, C represents the total number of categories contained in the N training set data. If the i-th preprocessed training data belongs to the t-th category, then otherwise Represents the probability that the i-th training sample output by the spiking neural network model to be trained belongs to the t-th class; Step 6-2: Set the maximum number of iterations, and use the AdamW optimization algorithm to iteratively optimize the pulse neural network model to be trained according to the loss function. When the maximum number of iterations is reached, stop the iteration process and obtain the trained pulse neural network model.
7. The classification method of a 3D discrete wavelet transform-based spiking neural network according to claim 6, characterized in that When the original dataset uses the CIFAR10 dataset or the CIFAR100 dataset, the maximum number of iterations is set to 400 times, when the original dataset uses the UCF101-DVS dataset or the HMDB51-DVS dataset, the maximum number of iterations is set to 200 times, when the original dataset uses the CIFAR10-DVS dataset or the DVS128-Gesture dataset, the maximum number of iterations is set to 100 times, and when the original dataset uses the Tiny-imagenet dataset, the maximum number of iterations is set to 300 times.