Cross-domain line-of-sight prediction method and device based on minimization of multi-branch network rotation variance, electronic equipment and storage medium
By performing illumination transformation and pseudo-label training on the cross-domain line of sight estimation method, combined with minimizing the rotation variance of multi-branch networks, the problems of low precision and large network burden in cross-domain line of sight estimation in the prior art are solved, and more efficient and robust line of sight prediction is achieved.
Patent Information
- Application Number
- CN202510050507.6
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-01-13
- Publication Date
- 2025-05-30
- Estimated Expiration
- 2045-01-13
AI Technical Summary
Existing cross-domain line-of-sight estimation methods exhibit low precision in the target domain, and existing methods require additional rotation enhancement models, increasing network burden, and uncertainty reduction methods are only applied in the target domain, ignoring the importance of the source domain.
A cross-domain line of sight prediction method based on minimizing the rotation variance of multi-branch networks is adopted. By performing illumination transformation on the source domain image, domain differences are reduced, and pseudo-label training is used in the target domain, combining variance minimization and pseudo-label supervision losses, the line of sight prediction network is optimized.
The accuracy of cross-domain line of sight estimation is improved, the predicted line of sight error is reduced, the adaptability and robustness of the model is enhanced, and the uncertainty of the network is reduced.
Smart Images

Figure CN120071422A_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of line-of-sight prediction, and in particular to a cross-domain line-of-sight prediction method, device, electronic device and storage medium based on minimizing the rotation variance of a multi-branch network. Background Art
[0002] Humans mainly understand the world by obtaining visual information. The human eye can perceive the surrounding environment, express a person's attention, convey personal emotions, etc., and has important significance in the natural interaction between humans and the outside world. Therefore, the line-of-sight direction of the human eye can also reflect potential brain cognitive processes and can be widely applied to various human-computer interaction fields, with good commercial value. In recent years, with the rapid development of deep learning technology and the accelerated iteration of high-performance hardware devices, many advances have been made in line-of-sight estimation methods based on deep learning. However, it should be noted that in practical applications, line-of-sight estimation networks are usually trained on data collected under controlled conditions, and then the line-of-sight estimation networks will be applied to an environment with different domains and uncontrolled conditions. Therefore, the accuracy of cross-domain line-of-sight estimation methods will be very low, so the breakthrough of cross-domain line-of-sight estimation is of great significance.
[0003] Existing face line-of-sight estimation methods can be divided into three steps: First, preprocess the image, including performing operations such as face detection, alignment, enhancement, and normalization on the collected single-face image or video frame in sequence; then use a convolutional network or Transformer to effectively extract high-level line-of-sight features from the high-dimensional image; finally, learn the non-linear mapping from the line-of-sight features to the line of sight through a multi-layer perceptron.
[0004] Dr. Xucong Zhang et al. proposed a full-face gaze estimation method based on the attention mechanism in the literature "X. Zhang, Y. Sugano, and M. Fritz. It's written all over your face: Full-face appearance-based gaze estimation. in Proceedings of IEEE Conference on Computer Vision and Pattern Recognition Workshops, 2017, 51-60". By a specific branch, it learns the weights of different regions of the face, thereby enhancing the weight of the eye region while suppressing the weights of other regions unrelated to the gaze direction. Bao et al. discovered the rotation consistency property in gaze estimation and introduced it as a pseudo-label in unsupervised domain adaptation in the literature "Y. Bao, Y. Liu, and H. Wang. Generalizing gaze estimation with rotation consistency. in Proceedings of IEEE Conference on Computer Vision and Pattern Recognition, 2022, 4207-4216". By training a rotation augmentation model in the source domain and then adapting the model to the target domain using synthetic images with rotation-consistent gaze directions. Cai et al. proposed an uncertainty reduction gaze estimation framework in the literature "X. Cai, J. Zeng, S. Shan, and X. Chen. Source-Free Adaptive Gaze Estimation by Uncertainty Reduction. in Proceedings of IEEE Conference on Computer Vision and Pattern Recognition, 2023, 22035-22045", which reduces the uncertainty of cross-domain gaze estimation by applying variance minimization and pseudo-label mechanisms in the target domain.
[0005] In the process of performing face gaze estimation, the above methods need to train an additional rotation augmentation model in the source domain to synthesize rotation-consistent images, adding an extra difficult task for gaze estimation and increasing the network burden. At the same time, uncertainty reduction is only applied in the target domain, ignoring the source domain. In fact, the performance of the network depends on a large number of labeled images in the source domain, rather than the target domain with only a small number of unlabeled images. Summary of the Invention
[0006] To solve the above problems existing in the prior art, the present invention provides a cross-domain line-of-sight prediction method, device, electronic device and storage medium based on minimizing the rotation variance of a multi-branch network. The technical problems to be solved by the present invention are realized through the following technical solutions:
[0007] In the first aspect of the embodiments of the present invention, a cross-domain line-of-sight prediction method based on minimizing the rotation variance of a multi-branch network is provided, including the following steps:
[0008] Perform illumination transformation on the source domain image to obtain a source domain input image set;
[0009] Input the source domain input image set into the initial line-of-sight prediction network to obtain an output result, and minimize the root mean square error between the output result and the line-of-sight true label to update the parameters of the initial line-of-sight prediction network, so as to obtain a pre-trained line-of-sight prediction network; wherein, the line-of-sight prediction network includes: a backbone network, an original line-of-sight angle prediction branch network and a plurality of rotation prediction branch networks arranged in parallel at the output end of the backbone network;
[0010] Input the illumination transformation image after performing illumination transformation on the target domain image into the pre-trained line-of-sight prediction network, output a first pre-trained result, and determine the target domain pseudo label according to the first pre-trained result, the rotation angle of the original line-of-sight angle prediction branch network and the plurality of rotation prediction branch networks;
[0011] Input the original image corresponding to the illumination transformation of the target domain image into the pre-trained line-of-sight prediction network, and train the pre-trained line-of-sight prediction network in combination with the target domain pseudo label to obtain a prediction target line-of-sight prediction network;
[0012] Input the target domain image into the prediction target line-of-sight prediction network, and determine the prediction target line of sight according to the output prediction result.
[0013] In an embodiment of the present invention, the step of inputting the source domain input image set into the initial line-of-sight prediction network to obtain an output result, and minimizing the root mean square error between the output result and the line-of-sight true label to update the parameters of the initial line-of-sight prediction network to obtain a pre-trained line-of-sight prediction network includes:
[0014] Input the source domain input image set into the backbone network of the initial line-of-sight prediction network to output line-of-sight features;
[0015] Input the line-of-sight features into the original line-of-sight angle prediction branch network of the initial line-of-sight prediction network to obtain a first output result;
[0016] Minimize the root mean square error between the first output result and the line-of-sight true label to update the parameters of the backbone network and the original line-of-sight angle prediction branch network of the initial line-of-sight prediction network;
[0017] Input the line-of-sight features into multiple rotation prediction branch networks of an initial line-of-sight prediction network to obtain a second output result;
[0018] Minimize the root mean square error between the second output result and the line-of-sight true label to update the parameters of the multiple rotation prediction branch networks of the initial line-of-sight prediction network;
[0019] After the update is completed, a pre-trained line-of-sight prediction network is obtained.
[0020] In an embodiment of the present invention, input the illumination-transformed image of the target domain image into the pre-trained line-of-sight prediction network, output a first pre-training result, and determine a target domain pseudo-label according to the first pre-training result, the rotation angles of the original line-of-sight angle prediction branch network and multiple rotation prediction branch networks, including:
[0021] Perform illumination transformation on the target domain image to obtain a target domain input image set;
[0022] Input the illumination-transformed image in the target domain input image set into the pre-trained line-of-sight prediction network to output a first pre-training result;
[0023] Determine the target domain pseudo-label according to the first pre-training result, the rotation angles of the original line-of-sight angle prediction branch network and multiple rotation prediction branch networks.
[0024] In an embodiment of the present invention, input the original image corresponding to the illumination transformation of the target domain image into the pre-trained line-of-sight prediction network, and train the pre-trained line-of-sight prediction network in combination with the target domain pseudo-label to obtain a predicted target line-of-sight prediction network, including:
[0025] Input the original image corresponding to the illumination transformation of the target domain image into the pre-trained line-of-sight prediction network to output a second pre-training result;
[0026] Determine the variance minimization loss according to the second pre-training result;
[0027] Determine the pseudo-label supervision minimization loss according to the second pre-training result and the target domain pseudo-label;
[0028] Determine the target loss according to the variance minimization loss and the pseudo-label supervision minimization loss;
[0029] Update the parameters of the pre-trained line-of-sight prediction network according to the target loss to obtain a predicted target line-of-sight prediction network.
[0030] In one embodiment of the present invention, the calculation formula for minimizing the root mean square error between the first output result and the true line-of-sight label is as follows:
[0031]
[0032] where, when θ = 0, θ represents the rotation angle of the original line-of-sight angle prediction branch network, G({I s}, θ) s represents the first output result, {I s} represents the source domain input image set, ‖·‖ 2 represents calculating the root mean square error, and g s represents the true line-of-sight label;
[0033] The calculation formula for minimizing the root mean square error between the second output result and the true line-of-sight label is as follows:
[0034]
[0035] where, when θ ≠ 0, θ represents the rotation angles of multiple rotation prediction branch networks, G({I s}, θ) s represents the second output result, represents the rotated line-of-sight label,
[0036] In one embodiment of the present invention, the expression of the target domain pseudo-label is:
[0037]
[0038] where K represents the number of the original line-of-sight angle prediction branch network and multiple rotation prediction branch networks; represents the first pre-training result, represents the illumination transformation image in the target domain input image set.
[0039] In one embodiment of the present invention, the expression of the target loss is:
[0040]
[0041] where I t represents the original image in the target domain input image set, α represents the weight hyperparameter,
[0042]
[0043] represents the second pre-training result.
[0044] In a second aspect of the embodiments of the present invention, a cross-domain line-of-sight prediction device based on minimizing the rotation variance of a multi-branch network is provided, including:
[0045] A transformation module, configured to perform illumination transformation on the source domain image to obtain a source domain input image set;
[0046] A pre-training module, configured to input the source domain input image set into an initial line-of-sight prediction network to obtain an output result, and minimize the root mean square error between the output result and the true line-of-sight label to update the parameters of the initial line-of-sight prediction network, so as to obtain a pre-trained line-of-sight prediction network; wherein, the line-of-sight prediction network includes: a backbone network, an original line-of-sight angle prediction branch network and a plurality of rotation prediction branch networks arranged in parallel at the output end of the backbone network;
[0047] A target domain training module, configured to input the illumination transformation image obtained by performing illumination transformation on the target domain image into the pre-trained line-of-sight prediction network, output a first pre-training result, and determine a target domain pseudo-label according to the first pre-training result, the rotation angles of the original line-of-sight angle prediction branch network and the plurality of rotation prediction branch networks;
[0048] The target domain training module is further configured to input the original image corresponding to the illumination transformation of the target domain image into the pre-trained line-of-sight prediction network, and train the pre-trained line-of-sight prediction network in combination with the target domain pseudo-label to obtain a predicted target line-of-sight prediction network;
[0049] A prediction module, configured to input the target domain image into the predicted target line-of-sight prediction network, and determine the predicted target line of sight according to the output prediction result.
[0050] In a third aspect of the embodiments of the present invention, an electronic device is provided, including a memory, a processor, and a computer program stored on the memory and executable on the processor, and when the processor executes the program, it implements a cross-domain line-of-sight prediction method based on minimizing the rotation variance of a multi-branch network provided in the first aspect of the embodiments of the present invention.
[0051] In a fourth aspect of the embodiments of the present invention, a computer-readable storage medium is provided, on which a computer program is stored, and when the computer program is executed by a processor, it implements a cross-domain line-of-sight prediction method based on minimizing the rotation variance of a multi-branch network provided in the first aspect of the embodiments of the present invention.
[0052] Advantages of the present invention:
[0053] The present invention provides a data augmentation method for randomly changing the illumination of an input image by suppressing the obvious distribution differences between the source domain and target domain data caused by illumination, appearance, etc., enabling the model to ignore the domain differences in the data and thus extract more robust gaze-related features; by minimizing the variance of the rotation prediction branch on the same input, averaging the prediction results of multiple branches reduces the uncertainty of the model prediction and decreases the predicted gaze error; during the target domain adaptation training, pseudo-labels are obtained from the images with illumination changes to ensure the gaze estimation ability of the model and, to a certain extent, make the features more robust.
[0054] Other features and advantages of the present invention will be described in the following specification, and, in part, will be obvious from the specification, or will be understood by implementing the present invention. The objectives and other advantages of the present invention can be achieved and obtained by the structures specifically pointed out in the written specification, claims, and drawings.
[0055] The technical solutions of the present invention will be further described in detail below through the accompanying drawings and embodiments. Description of the Drawings
[0056] The accompanying drawings are used to provide a further understanding of the present invention, and constitute a part of the specification. They are used together with the embodiments of the present invention to explain the present invention, and do not constitute a limitation to the present invention. In the drawings:
[0057] Figure 1 is a schematic flowchart of a cross-domain gaze prediction method based on minimizing the rotation variance of a multi-branch network provided by an embodiment of the present invention;
[0058] Figure 2 is a schematic diagram of an algorithm framework of a cross-domain gaze prediction method based on minimizing the rotation variance of a multi-branch network provided by an embodiment of the present invention;
[0059] Figure 3 is a schematic diagram of a cross-domain gaze prediction device based on minimizing the rotation variance of a multi-branch network provided by an embodiment of the present invention. Detailed Embodiments
[0060] The present invention will be further described in detail below in combination with specific embodiments, but the embodiments of the present invention are not limited thereto.
[0061] As Figure 1 shown, the first aspect of the embodiment of the present invention provides a cross-domain gaze prediction method based on minimizing the rotation variance of a multi-branch network, including the following steps:
[0062] Step 11: Perform illumination transformation on the source domain images to obtain a source domain input image set.
[0063] Step 12: Input the source domain input image set into the initial gaze prediction network to obtain an output result, and minimize the root mean square error between the output result and the gaze ground truth label to update the parameters of the initial gaze prediction network, thereby obtaining a pre-trained gaze prediction network.
[0064] The gaze prediction network includes: a backbone network, a parallel original gaze angle prediction branch network located at the output end of the backbone network, and multiple rotation prediction branch networks.
[0065] Step 13: Input the illumination-transformed image obtained by performing illumination transformation on the target domain image into the pre-trained gaze prediction network, output a first pre-trained result, and determine the target domain pseudo-label according to the first pre-trained result, the rotation angles of the original gaze angle prediction branch network and the multiple rotation prediction branch networks.
[0066] Step 14: Input the original image corresponding to the illumination transformation of the target domain image into the pre-trained gaze prediction network, and train the pre-trained gaze prediction network in combination with the target domain pseudo-label to obtain a predicted target gaze prediction network.
[0067] Step 15: Input the target domain image into the predicted target gaze prediction network, and determine the predicted target gaze according to the output prediction result.
[0068] To solve the cross-domain gaze estimation problem, in this embodiment, by reducing the sample uncertainty and model uncertainty of the unlabeled target domain data, the gaze estimation model trained in the source domain can better adapt to the target domain. Sample uncertainty mainly refers to the inherent noise in the input image, such as sensor noise and motion blur. The uncertainty of the model is caused by the implementation method of the deep network itself. A simple model may not achieve the expected fitting effect, while a too complex model will lead to overfitting. In the method of this embodiment, the data augmentation method reduces the sample uncertainty and can extract more robust gaze features. The variance minimization method reduces the model uncertainty and can better adapt to a small amount of unlabeled target domain data, improving the accuracy of cross-domain gaze estimation.
[0069] The method involved in this embodiment mainly completes training in a labeled face image dataset and predicts the gaze direction of a person's eyes in an unlabeled face image. It can be used in the medical and health field, such as patients with amyotrophic lateral sclerosis completing daily activities with the help of an eye tracker; in the field of assisted driving, such as providing a human-computer interaction function to help free the driver's hands; and in the virtual reality (VR) field, such as realizing a human-machine synchronous immersive real-scene interaction.
[0070] As Figure 2 shown, the second aspect of the embodiment of the present invention provides a cross-domain gaze prediction method based on minimizing the multi-branch rotation variance, including the following steps:
[0071] Step 21: Perform illumination transformation on the source domain images to obtain the source domain input image set.
[0072] In this step, if I represents the input image and g represents the line of sight, then the source domain labeled data can be represented by where I s represents the source domain image and g s represents the true label of the line of sight corresponding to the image. To reduce the uncertainty of the input samples, that is, to suppress the domain difference and noise, perform random illumination transformation on the input image. The method of illumination transformation adopts the gamma transformation method, which is a process of non-linear brightness transformation of the image. The implementation method is as follows:
[0073] o = c·r γ , r ∈ [0, 1]
[0074] where c is a constant, usually taken as 1, r is the normalized pixel value of the input image I, o is the normalized pixel value of the output illumination transformation image, and γ is the Gamma value of the gamma transformation. Different degrees of illumination transformation are achieved by adjusting γ. Here, three degrees of gamma transformation are implemented, namely γ < 1, γ = 1, γ > 1. Therefore, the transformed images are 3, one original image, one image with randomly enhanced illumination, and one image with randomly weakened illumination. The transformed source domain input images can be represented by where {I s} represents the set of the original image I s , the image with randomly enhanced illumination and the image with randomly weakened illumination , that is, the source domain input image set.
[0075] Step 22: Input the source domain input image set into the initial line of sight prediction network to obtain the output result, and minimize the root mean square error between the output result and the true label of the line of sight to update the parameters of the initial line of sight prediction network, and obtain the pre-trained line of sight prediction network.
[0076] Among them, the line of sight prediction network includes: a backbone network, an original line of sight angle prediction branch network and multiple rotation prediction branch networks arranged in parallel at the output end of the backbone network.
[0077] In this step, the method disclosed in the literature "K. He, X. Zhang, S. Ren, and J. Sun, Deep Residual Learning for Image Recognition. in Proceedings of IEEE Conference on Computer Vision and Pattern Recognition, 2016, 770 - 778" is used to construct a backbone network as ResNet18. The backbone network is used to extract gaze features from the input image, and then the extracted gaze features are globally average pooled and fed into a fully connected layer to predict the gaze. At the same time, in order to simplify the rotation augmentation model, combined with the variance minimization method, a primary gaze angle prediction branch network is constructed to predict the primary gaze, and multiple rotation prediction branch networks directly predict the rotated gaze. The primary gaze angle prediction branch network and the rotation prediction branch networks are in parallel. Let θ represent the rotation angle, and the initial parameters of the backbone network and a prediction branch network with a rotation angle of θ can be represented by G(*, θ|θ = 0,..., K) 0 For the primary gaze angle prediction branch network, the rotation angle θ = 0, and for the rotation prediction branch networks, the rotation angle θ ≠ 0.
[0078] The specific steps of step 22 include steps 221 - 226:
[0079] Step 221, input the source domain input image set into the backbone network of the initial gaze prediction network to output gaze features. For any input sample, the three input images {I s} of each sample pass through the backbone network to obtain the corresponding three gaze features, and these three gaze features should be equivalent.
[0080] Step 222, input the gaze features into the primary gaze angle prediction branch network of the initial gaze prediction network to obtain the first output result.
[0081] Step 223, minimize the root mean square error between the first output result and the gaze ground truth label to update the parameters of the backbone network and the primary gaze angle prediction branch network of the initial gaze prediction network.
[0082] All these three gaze features are input into the primary gaze angle prediction branch network with θ = 0, that is, to predict the primary gaze without rotation, and then minimize the error between the predicted gaze (the first output result) and the true label. The calculation method is as follows:
[0083]
[0084] where G({I s}, θ) s represents the network parameters obtained from source domain training, ‖·‖2 denotes calculating the root mean square error between the two, g s denotes the true label of the line of sight. The labels of the three input images are the same. Update the parameters of the original line-of-sight angle prediction branch network and the backbone network using the original line-of-sight angle prediction error.
[0085] Here, the three line-of-sight features correspond to three output results and three minimized root mean square errors. After summing the three minimized root mean square errors, update the corresponding network parameters.
[0086] Step 224: Input the line-of-sight features into multiple rotation prediction branch networks of the initial line-of-sight prediction network to obtain a second output result.
[0087] Step 225: Minimize the root mean square error between the second output result and the true line-of-sight label to update the parameters of the multiple rotation prediction branch networks of the initial line-of-sight prediction network.
[0088] Meanwhile, all three line-of-sight features are input into multiple parallel rotation prediction branch networks to directly predict the rotated line of sight. In the multiple rotation prediction branch networks, θ represents the rotation angle, and the rotated line-of-sight label is Similarly, minimize the error between the predicted line of sight and the true label. The calculation method is as follows:
[0089]
[0090] Among them, the calculation method of the error remains the same, which is still the root mean square error. However, when updating the network parameters, since the features of the backbone network are consistent, here a fully connected layer is directly used to predict the linear rotation angle. In order to prevent the feature extraction ability of the backbone network from being damaged, only the backbone network will be updated when predicting the original line-of-sight angle, and the parameters of the backbone network will be frozen when predicting the rotated line-of-sight angle and will not be updated.
[0091] Here, the three line-of-sight features correspond to three output results and three minimized root mean square errors. After summing the three minimized root mean square errors, update the corresponding network parameters.
[0092] Among them, Step 222 and Step 224 are executed in parallel.
[0093] Step 226: After the update is completed, obtain the pre-trained line-of-sight prediction network. G(*, θ) s denotes the network parameters of the pre-trained line-of-sight prediction network.
[0094] The training in the target domain adaptation stage is carried out on the basis of the source domain training. Using the model trained in the source domain as the pre-trained base model for target domain training includes the following Steps 23 - Step 24:
[0095] Step 23: Input the illumination-transformed image of the target-domain image into the pre-trained gaze prediction network, output the first pre-trained result, and determine the target-domain pseudo-label based on the first pre-trained result, the original gaze angle prediction branch network, and the rotation angles of multiple rotation prediction branch networks.
[0096] The specific steps of Step 23 include Step 231 - Step 233:
[0097] Step 231: Perform illumination transformation on the target-domain image to obtain the target-domain input image set.
[0098] Since the target-domain data lacks labels, it is possible to consider using the enhanced images to calculate the pseudo-labels. Let represent the target-domain data, where {I t} represents the set of the original image I t , the illumination randomly enhanced image and the illumination randomly weakened image , which is also the target-domain input image set.
[0099] Step 232: Input the illumination-transformed image in the target-domain input image set into the pre-trained gaze prediction network, and output the first pre-trained result.
[0100] Step 233: Determine the target-domain pseudo-label based on the first pre-trained result, the original gaze angle prediction branch network, and the rotation angles of multiple rotation prediction branch networks.
[0101] Two images with random illumination transformation are fed into the network to predict multiple gazes, and then the original gaze prediction angle and the multiple rotation prediction gaze angles are averaged to obtain the pseudo-label p t , as follows:
[0102]
[0103] where K represents the number of the original gaze angle prediction branch network and multiple rotation prediction branch networks; represents the first pre-trained result, represents the illumination-transformed image in the target-domain input image set, which is also the set of the illumination randomly enhanced image and the illumination randomly weakened image . The obtained pseudo-label is used to ensure that the gaze estimation ability of the model will not fluctuate greatly during subsequent training in the target domain.
[0104] Step 24: Input the original image corresponding to the illumination transformation of the target-domain image into the pre-trained gaze prediction network, and train the pre-trained gaze prediction network in combination with the target-domain pseudo-label to obtain the predicted target gaze prediction network.
[0105] The specific steps of step 24 include steps 241 - 245:
[0106] In step 241, the original image corresponding to the illumination transformation of the target domain image is input into the pre-trained gaze prediction network, and a second pre-trained result is output.
[0107] In step 242, the variance minimization loss is determined according to the second pre-trained result.
[0108] The purpose of multi-branch rotation prediction is to reduce the uncertainty of the model, which can be understood here as minimizing the variance between predictions in multi-branch rotation prediction under the same input. t Denote the input original image of the target domain, the variance minimization loss is as follows:
[0109]
[0110] Among them, Denote the second pre-trained result, which is the predicted gaze of the target domain of the input original image. According to different rotation angles, there are K predicted gazes, and the variance of these K predicted gazes is minimized.
[0111] In step 243, the pseudo-label supervised minimization loss is determined according to the second pre-trained result and the target domain pseudo-label.
[0112] Meanwhile, using the obtained pseudo-label, a pseudo-label supervised minimization loss is calculated for the predicted gaze, as follows:
[0113]
[0114] Among them, the loss between the predicted gaze of the target domain and the pseudo-label is calculated using the same root mean square error, and the root mean square errors of multiple predicted angles are directly averaged.
[0115] In step 244, the target loss is determined according to the variance minimization loss and the pseudo-label supervised minimization loss.
[0116] To balance the above two losses, a weight parameter is added. Therefore, the final target loss in the target domain adaptation training phase is as follows:
[0117]
[0118] Among them, α is the weight hyperparameter.
[0119] In step 245, the parameters of the pre-trained gaze prediction network are updated according to the target loss to obtain the predicted target gaze prediction network.
[0120] Here, steps 23 and 24 are the processes of target domain adaptation training, and the initial parameters of the network during training are G(*,θ)s ,G(*,θ) t represents the network parameters obtained from training in the target domain.
[0121] Step 25: Input the target domain image into the prediction target line-of-sight prediction network, and determine the predicted target line-of-sight according to the output prediction result.
[0122] In the inference stage, input the target domain image I t , using the results of multi-branch rotation prediction, after subtracting the corresponding rotation angle from each rotation line-of-sight, take the average to obtain the final predicted line-of-sight g. The evaluation index uses the common angular error, and the calculation formula is as follows:
[0123]
[0124] where the angular error L angular is usually used to evaluate the accuracy of the three-dimensional line-of-sight estimation method. Among them, the estimated gaze direction is g ∈ R 3 , and the actual gaze direction is
[0125] The effects of the present invention can be further illustrated by the following simulation experiments.
[0126] 1. Simulation conditions
[0127] This invention is carried out on a central processing unit of Intel(R) Core(TM) i7-7820X CPU@3.60GHz, NVIDIA GeForce RTX 2080Ti, and Ubuntu 22.04 operating system. The source domain dataset uses the Gaze360 dataset disclosed in the literature "P. Kellnhofer, A. Recasens, and S. Stent. Gaze360: Physically unconstrained gaze estimation in the wild. in Proceedings of IEEE International Conference on Computer Vision, 2019, 6912-6921", abbreviated as DG. The target domain dataset uses the MPIIFaceGaze dataset disclosed in the literature "X. Zhang, Y. Sugano, and M. Fritz. It's written all over your face: Full-face appearance-based gaze estimation. in Proceedings of IEEE Conference on Computer Vision and Pattern Recognition, 2017, 51-60", abbreviated as DM.
[0128] The methods compared in the experiment are as follows:
[0129] One is the method based on outlier collaborative adaptation, denoted as PnP-GA, and the reference is "Y. Liu, R. liu, H. wang, and F. Lu. Generalizing Gaze Estimation with outlier-guided Collaborative Adaptation. in Proceedings of IEEE Conference on Computer Vision and Pattern Recognition, 2021, 3835-3844".
[0130] One is the method based on rotation consistency, denoted as RUDA, and the reference is "Y. Bao, Y. Liu, and H. Wang. Generalizing gaze estimation with rotation consistency. in Proceedings of IEEE Conference on Computer Vision and Pattern Recognition, 2022, 4207 - 4216".
[0131] The other is the method based on contrastive learning, denoted as CRGA, and the reference is "Y. Wang, Y. Jiang, and J. Li. Contrastive regression for domain adaptation on gaze estimation. in Proceedings of IEEE Conference on Computer Vision and Pattern Recognition, 2022, 19376 - 19385".
[0132] 2. Simulation content
[0133] According to the specific embodiments of the present invention, the angular error of cross - domain gaze estimation is calculated. The source domain is DG and the target domain is DM, and it is compared with the angular errors of the RUDA method and the CRGA method. The results are shown in Table 1.
[0134] Table 1 Cross - domain gaze estimation angular error
[0135] Method PnP-GA RUDA CRGA The present invention Error 6.18 6.20 5.89 5.78
[0136] As can be seen from Table 1, due to the multi - branch rotation - minimized variance method adopted by the present invention, it can reduce the sample uncertainty and model uncertainty of the unlabeled target domain data, and thus can achieve a lower predicted angular error in cross - domain gaze estimation, verifying the effectiveness of the present invention.
[0137] As Figure 3 shown, the third aspect of the embodiment of the present invention provides a cross - domain gaze prediction device based on minimizing the rotation variance of a multi - branch network, including:
[0138] A transformation module 31, configured to perform illumination transformation on the source domain image to obtain a source domain input image set;
[0139] The pre-training module 32 is configured to input the source domain input image set into the initial line-of-sight prediction network to obtain an output result, and minimize the root mean square error between the output result and the line-of-sight ground truth label to update the parameters of the initial line-of-sight prediction network, thereby obtaining a pre-trained line-of-sight prediction network. The line-of-sight prediction network includes: a backbone network, an original line-of-sight angle prediction branch network and a plurality of rotation prediction branch networks arranged in parallel at the output end of the backbone network.
[0140] The target domain training module 33 is configured to input the illumination-transformed image obtained by performing illumination transformation on the target domain image into the pre-trained line-of-sight prediction network, output a first pre-trained result, and determine a target domain pseudo-label according to the first pre-trained result, the rotation angles of the original line-of-sight angle prediction branch network and the plurality of rotation prediction branch networks.
[0141] The target domain training module 33 is further configured to input the original image corresponding to the illumination transformation of the target domain image into the pre-trained line-of-sight prediction network, and train the pre-trained line-of-sight prediction network in combination with the target domain pseudo-label to obtain a predicted target line-of-sight prediction network.
[0142] The prediction module 34 is configured to input the target domain image into the predicted target line-of-sight prediction network, and determine the predicted target line of sight according to the output prediction result.
[0143] In an embodiment of the present invention, inputting the source domain input image set into the initial line-of-sight prediction network to obtain an output result, and minimizing the root mean square error between the output result and the line-of-sight ground truth label to update the parameters of the initial line-of-sight prediction network, thereby obtaining a pre-trained line-of-sight prediction network, includes:
[0144] Input the source domain input image set into the backbone network of the initial line-of-sight prediction network to output line-of-sight features.
[0145] Input the line-of-sight features into the original line-of-sight angle prediction branch network of the initial line-of-sight prediction network to obtain a first output result.
[0146] Minimize the root mean square error between the first output result and the line-of-sight ground truth label to update the parameters of the backbone network and the original line-of-sight angle prediction branch network of the initial line-of-sight prediction network.
[0147] Input the line-of-sight features into the plurality of rotation prediction branch networks of the initial line-of-sight prediction network to obtain a second output result.
[0148] Minimize the root mean square error between the second output result and the line-of-sight ground truth label to update the parameters of the plurality of rotation prediction branch networks of the initial line-of-sight prediction network.
[0149] After the update is completed, a pre-trained line-of-sight prediction network is obtained.
[0150] In one embodiment of the present invention, the illumination-transformed image obtained by performing illumination transformation on the target-domain image is input into a pre-trained gaze prediction network to output a first pre-trained result, and a target-domain pseudo-label is determined according to the first pre-trained result, the rotation angles of the original gaze angle prediction branch network and multiple rotation prediction branch networks, including:
[0151] Perform illumination transformation on the target-domain image to obtain a target-domain input image set;
[0152] Input the illumination-transformed image in the target-domain input image set into a pre-trained gaze prediction network to output a first pre-trained result;
[0153] Determine the target-domain pseudo-label according to the first pre-trained result, the rotation angles of the original gaze angle prediction branch network and multiple rotation prediction branch networks.
[0154] In one embodiment of the present invention, the original image corresponding to the illumination transformation of the target-domain image is input into a pre-trained gaze prediction network, and the pre-trained gaze prediction network is trained in combination with the target-domain pseudo-label to obtain a prediction target gaze prediction network, including:
[0155] Input the original image corresponding to the illumination transformation of the target-domain image into a pre-trained gaze prediction network to output a second pre-trained result;
[0156] Determine the variance minimization loss according to the second pre-trained result;
[0157] Determine the pseudo-label supervision minimization loss according to the second pre-trained result and the target-domain pseudo-label;
[0158] Determine the target loss according to the variance minimization loss and the pseudo-label supervision minimization loss;
[0159] Update the parameters of the pre-trained gaze prediction network according to the target loss to obtain a prediction target gaze prediction network.
[0160] In one embodiment of the present invention, the calculation formula for minimizing the root mean square error between the first output result and the gaze true label is:
[0161]
[0162] Where, when θ = 0, θ represents the rotation angle of the original gaze angle prediction branch network, G({I s}, θ) s represents the first output result, {I s} represents the source-domain input image set, ‖·‖ 2 represents the calculation of the root mean square error, and g s represents the gaze true label;
[0163] The calculation formula for minimizing the root mean square error between the second output result and the true label of the line of sight is as follows:
[0164]
[0165] Among them, when θ≠0, θ represents the rotation angle of multiple rotation prediction branch networks, and G({I s},θ) s represents the second output result, represents the rotated line-of-sight label,
[0166] In an embodiment of the present invention, the expression of the target domain pseudo-label is:
[0167]
[0168] Among them, K represents the number of the original line-of-sight angle prediction branch network and multiple rotation prediction branch networks; represents the first pre-training result, represents the illumination transformation image in the target domain input image set.
[0169] In an embodiment of the present invention, the expression of the target loss is:
[0170]
[0171] Among them, I t represents the original image in the target domain input image set, α represents the weight hyperparameter,
[0172]
[0173] represents the second pre-training result.
[0174] The fourth aspect of the embodiments of the present invention provides an electronic device, including a memory, a processor, and a computer program stored on the memory and executable on the processor. When the processor executes the program, it implements the cross-domain line-of-sight prediction method provided by the embodiments of the present invention based on minimizing the rotation variance of a multi-branch network.
[0175] The fifth aspect of the embodiments of the present invention further provides a computer-readable storage medium, on which a computer program is stored. When the computer program is executed by a processor, it implements the steps of the cross-domain line-of-sight prediction method provided by the embodiments of the present invention based on minimizing the rotation variance of a multi-branch network.
[0176] Among them, the memory may include a Random Access Memory (RAM), or may also include a Non-Volatile Memory (NVM), such as at least one disk memory. Optionally, the memory may also be at least one storage device located far from the aforementioned processor.
[0177] The aforementioned processor may be a general-purpose processor, including a Central Processing Unit (CPU), a Network Processor (NP), etc.; it may also be a Digital Signal Processor (DSP), an Application Specific Integrated Circuit (ASIC), a Field-Programmable Gate Array (FPGA), or other programmable logic devices, discrete gate or transistor logic devices, discrete hardware devices.
[0178] The method provided by the embodiments of the present invention can be applied to an electronic device. Specifically, the electronic device may be: a desktop computer, a portable computer, a smart mobile terminal, a server, etc. This is not limited herein, and any electronic device that can implement the present invention belongs to the protection scope of the present invention.
[0179] For the device / electronic device embodiments, since they are basically similar to the method embodiments, the description is relatively simple, and for the relevant parts, refer to the partial description of the method embodiments.
[0180] The present invention is described with reference to the flowcharts and / or block diagrams of methods, devices (apparatus), and computer program products according to the embodiments of the present invention. It should be understood that each process and / or block in the flowchart and / or block diagram, and the combination of processes and / or blocks in the flowchart and / or block diagram, can be implemented by computer program instructions. These computer program instructions can be provided to the processor of a general-purpose computer, a special-purpose computer, an embedded processor, or other programmable data processing devices to generate a machine, so that the instructions executed by the processor of the computer or other programmable data processing devices generate a device for implementing the functions specified in Figure 1 one process or multiple processes and / or blocks Figure 1 one block or multiple blocks.
[0181] These computer program instructions can also be stored in a computer-readable memory that can direct a computer or other programmable data processing apparatus to operate in a particular manner, such that the instructions stored in the computer-readable memory produce a manufacture including instruction means embodying the functionality specified in the flowchart Figure 1 a flowchart or multiple flowcharts and / or blocks Figure 1 specified in a block or multiple blocks.
[0182] These computer program instructions can also be loaded onto a computer or other programmable data processing apparatus to cause a series of operational steps to be performed on the computer or other programmable apparatus to produce a computer-implemented process, whereby the instructions executed on the computer or other programmable apparatus provide steps for implementing the functionality specified in the flowchart Figure 1 a flowchart or multiple flowcharts and / or blocks Figure 1 specified in a block or multiple blocks.
[0183] Obviously, those skilled in the art can make various changes and modifications to the present invention without departing from the spirit and scope of the present invention. Thus, if these modifications and variations of the present invention fall within the scope of the claims of the present invention and their equivalent technologies, the present invention is also intended to include these modifications and variations.
Claims
1. A cross-domain sight line prediction method based on minimizing the rotation variance of a multi-branch network, characterized in that: The following steps are involved: Perform illumination transformation on the source domain image to obtain a source domain input image set; Inputting the source domain input image set into the initial sight line prediction network to obtain an output result, and minimizing the root mean square error between the output result and the true sight line label to update the parameters of the initial sight line prediction network to obtain a pre-trained sight line prediction network; wherein the sight line prediction network includes: a backbone network, an original sight line angle prediction branch network located at the output end of the backbone network and parallel to each other, and a plurality of rotation prediction branch networks; Inputting the illumination transformed image of the target domain image after the illumination transformation into the pre-trained sight line prediction network, outputting a first pre-training result, and determining the target domain pseudo label according to the first pre-training result, the original sight line angle prediction branch network and the rotation angles of the plurality of rotation prediction branch networks; Inputting the original image corresponding to the illumination transformation of the target domain image into the pre-trained sight line prediction network, and training the pre-trained sight line prediction network in combination with the target domain pseudo label to obtain a predicted target sight line prediction network; The target domain image is input into the predicted target sight line prediction network, and the predicted target sight line is determined according to the output prediction result.
2. The method according to claim 1, characterized in that The source domain input image set is input into the initial sight line prediction network to obtain an output result, and the root mean square error between the output result and the true sight line label is minimized to update the parameters of the initial sight line prediction network to obtain a pre-trained sight line prediction network, including: The source domain input image set is input into the backbone network of the initial sight line prediction network, and the sight line features are output; Inputting the sight line feature into an original sight line angle prediction branch network of an initial sight line prediction network to obtain a first output result; Minimize the root mean square error between the first output result and the true line of sight label to update the parameters of the backbone network of the initial line of sight prediction network and the original line of sight angle prediction branch network; Inputting the sight line feature into a plurality of rotation prediction branch networks of an initial sight line prediction network to obtain a second output result; Minimizing the root mean square error between the second output result and the true line of sight label to update the parameters of multiple rotation prediction branch networks of the initial line of sight prediction network; After the update is completed, the pre-trained sight prediction network is obtained.
3. The method according to claim 1, characterized in that The step of inputting the illumination transformed image after the illumination transformation of the target domain image into the pre-trained sight line prediction network, outputting a first pre-training result, and determining the target domain pseudo label according to the first pre-training result, the original sight line angle prediction branch network, and the rotation angles of the plurality of rotation prediction branch networks, comprises: Perform illumination transformation on the target domain image to obtain a target domain input image set; Inputting the illumination transformation image in the target domain input image set into the pre-trained sight line prediction network, and outputting a first pre-training result; A target domain pseudo label is determined according to the first pre-training result, the original sight angle prediction branch network, and the rotation angles of multiple rotation prediction branch networks.
4. The method according to claim 1, characterized in that The original image corresponding to the target domain image subjected to illumination transformation is input into the pre-trained sight line prediction network, and the pre-trained sight line prediction network is trained in combination with the target domain pseudo label to obtain a predicted target sight line prediction network, including: Inputting the original image corresponding to the illumination transformation of the target domain image into the pre-trained sight line prediction network, and outputting a second pre-training result; Determine a variance minimization loss according to the second pre-training result; Determine a pseudo-label supervision minimization loss according to the second pre-training result and the target domain pseudo-label; Determining a target loss according to the variance minimization loss and the pseudo-label supervision minimization loss; The parameters of the pre-trained sight line prediction network are updated according to the target loss to obtain a predicted target sight line prediction network.
5. The method according to claim 2, characterized in that The calculation formula for minimizing the root mean square error between the first output result and the true line of sight label is: When θ = 0, θ represents the rotation angle of the original line of sight angle prediction branch network, G({I s },θ) s Represents the first output result, {I s } represents the source domain input image set, ‖·‖2 represents the calculation of the root mean square error, g s Indicates the true label of the sight line; The calculation formula for minimizing the root mean square error between the second output result and the true line of sight label is: When θ≠0, θ represents the rotation angle of multiple rotation prediction branch networks, G({I s },θ) s Represents the second output result, Represents the rotated sight line label, 6. The method according to claim 3, characterized in that The expression of the target domain pseudo label is: Wherein, K represents the number of the original sight angle prediction branch network and multiple rotation prediction branch networks; represents the first pre-training result, Represents the illumination transformed image in the target domain input image set.
7. The method according to claim 4, characterized in that The expression of the target loss is: Among them, I t represents the original image in the target domain input image set, α represents the weight hyperparameter, Represents the second pre-training result.
8. A cross-domain sight line prediction device based on minimizing the rotation variance of a multi-branch network, characterized in that: include: A transformation module, used to perform illumination transformation on the source domain image to obtain a source domain input image set; A pre-training module is used to input a source domain input image set into an initial sight prediction network to obtain an output result, and minimize the root mean square error between the output result and the true sight label to update the parameters of the initial sight prediction network to obtain a pre-trained sight prediction network; wherein the sight prediction network includes: a backbone network, an original sight angle prediction branch network located at an output end of the backbone network and parallel to each other, and a plurality of rotation prediction branch networks; a target domain training module, configured to input an illumination transformed image of a target domain image after illumination transformation into the pre-trained sight line prediction network, output a first pre-training result, and determine a target domain pseudo label according to the first pre-training result, an original sight line angle prediction branch network, and a rotation angle of a plurality of rotation prediction branch networks; The target domain training module is further used to input the original image corresponding to the illumination transformation of the target domain image into the pre-trained sight line prediction network, and train the pre-trained sight line prediction network in combination with the target domain pseudo label to obtain a predicted target sight line prediction network; The prediction module is used to input the target domain image into the predicted target sight line prediction network, and determine the predicted target sight line according to the output prediction result.
9. An electronic device comprising a memory, a processor, and a computer program stored in the memory and executable on the processor, characterized in that: When the processor executes the program, the cross-domain sight line prediction method based on minimizing the multi-branch network rotation variance as described in any one of claims 1 to 7 is implemented.
10. A computer-readable storage medium having a computer program stored thereon, characterized in that: When the computer program is executed by a processor, the cross-domain sight line prediction method based on minimizing the rotation variance of a multi-branch network according to any one of claims 1 to 7 is implemented.
Citation Information
Patent Citations
Sight line prediction method, device and system and readable storage medium
CN110008835A
Sight line estimation method and device and storage medium
CN114863200A
Prediction model training method and device, equipment, medium and program product
CN117217368A
Systems and methods for training machine learning model based on cross-domain data
US20220198339A1
Gaze estimation cross-scene adaptation method and device based on outlier guidance
US20220405953A1
Cited By
Face feature line-of-sight estimation method based on Gaussian probability embedding
CN121438376A