A Method and Device for Counterfactual Prediction of Entity Trajectories
Through the mixed factor estimation and transmission model, combined with causal correlation and object graph feature extraction, the problem of insufficient generalization ability of the entity trajectory prediction model in the prior art in different game scenarios is solved, and the accurate prediction of entity motion trajectory and the saving of computing resources are achieved.
Patent Information
- Application Number
- CN202111478788.3
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2021-12-06
- Publication Date
- 2025-07-11
- Estimated Expiration
- 2041-12-06
AI Technical Summary
The counterfactual prediction model used for entity trajectory prediction in the prior art has insufficient generalization capabilities in different game scenarios, resulting in increased computing resource consumption.
The confounding factor estimation model and confounding factor transmission model are used to obtain the historical video sequence and video frames to be tested during the game process, and the causal relationship between entities is modeled using the absolute position encoding layer, the global causal correlation attention layer and the scaling dot product self-attention mechanism, and feature extraction and prediction are combined with the causal graph and the object graph to achieve the prediction of the entity's motion trajectory after the disturbance.
It realizes accurate prediction of entity motion trajectories in different game scenarios, reduces the consumption of computing resources, and improves the generalization ability of the model.
Smart Images

Figure CN114377398B_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of image processing, and in particular, to a counterfactual prediction method and device for entity trajectories. Background Art
[0002] In the physical world, discovering potential causal associations is an important ability for inferring the surrounding environment and predicting future states. Counterfactual prediction from visual input requires simulating future states based on scenarios that did not occur in the past, which is an important part of causal association task research and has received increasing attention. At the same time, this prediction technology can also be widely applied to kinetic games, such as block stacking and trajectory prediction games.
[0003] In kinetic game scenarios, external perturbations are usually encountered, such as the actions performed by players in single-player games, and these external perturbations will cause changes in the movement trajectories of entities in the game. How to predict the movement trajectories of entities after being affected by external perturbations is a key factor in enhancing the game experience. In order to accurately predict the movement trajectories of entities after being affected by external perturbations, it is very important to discover the hidden causal associations in the game scenario and model the relationships between entities. Although existing methods can learn limited intuitive physical information in some specific game scenarios, they rely on direct supervision information of potential physical attributes in the game, which will limit the model to specific game scenarios and lack generalization ability, that is, specific models can only be designed for specific game scenarios, which will greatly increase the consumption of computing resources. Summary of the Invention
[0004] The present invention provides a counterfactual prediction method and device for entity trajectories, which are used to solve the defect that the counterfactual prediction model for entity trajectory prediction in the prior art has poor generalization ability, and realize the applicability of the counterfactual prediction model for entity trajectory prediction to different game scenarios and strong generalization ability.
[0005] The present invention provides a counterfactual prediction method for entity trajectories, including:
[0006] Obtain a historical video sequence and a to-be-tested video frame during the game process; wherein, the to-be-tested video frame is image data corresponding to the moment of adding perturbation during the game process;
[0007] Input the historical video sequence and the to-be-tested video frame into a perception model to obtain 3D position information of each entity in the historical video sequence and the to-be-tested video frame;
[0008] Input the 3D position information of each entity in the historical video sequence and the to-be-tested video frame into a counterfactual prediction model to obtain a prediction result of the movement trajectories of each entity in the game after adding perturbation;
[0009] Among them, the counterfactual prediction model includes a confounding factor estimation model and a confounding factor transmission model. The confounding factor estimation model is used to obtain the confounding factors in the game according to the 3D position information of each entity in the historical video sequence; the confounding factor transmission model is used to obtain the prediction results of the motion trajectories of each entity in the game after perturbation according to the 3D position information of each entity in the video frame to be measured and the confounding factors.
[0010] According to a method for counterfactual prediction of entity trajectories provided by the present invention, the structure of the confounding factor estimation model includes:
[0011] An absolute position encoding layer for calculating the absolute position information of each entity in the historical video sequence;
[0012] A global causal association attention layer for modeling the causal relationship between each entity in the game by using a scaled dot-product self-attention mechanism according to the 3D position information and the absolute position information of each entity in the historical video sequence, and obtaining the confounding factors in the game based on the causal relationship.
[0013] According to a method for counterfactual prediction of entity trajectories provided by the present invention, calculating the absolute position information of each entity in the historical video sequence includes:
[0014] Obtaining the sequence information of each entity in the historical video sequence;
[0015] According to the sequence information of each entity in the historical video sequence, using a sine function to calculate and obtain the absolute position information of each entity.
[0016] According to a method for counterfactual prediction of entity trajectories provided by the present invention, modeling the causal relationship between each entity in the game by using a scaled dot-product self-attention mechanism includes:
[0017] Using a scaled dot-product self-attention mechanism to calculate the correlation degree between each pair of entities to obtain the causal relationship between all entities in the historical video sequence; the calculation of the correlation degree is shown in Equation 1:
[0018]
[0019] In the formula, are the query vector, the key vector, and the value vector respectively. The query vector, the key vector, and the value vector are respectively obtained by multiplying the 3D position matrix and / or the absolute position matrix by the corresponding weight matrices W qsi 、W krj 、W vrj ; the 3D position matrix and the absolute position matrix are respectively used to store the 3D position information and the absolute position information of each entity in the historical video sequence. Represents the association degree between entity i in video frame s of the historical video sequence and entity j in video frame r; d k Represents the dimension of the key vector, and softmax() represents a probability-based multi-classification function.
[0020] According to a counterfactual prediction method for entity trajectories provided by the present invention, the structure of the confounding factor transmission model includes:
[0021] A splicing layer for superimposing the causal graph and the object graph to obtain a superimposed graph; wherein, the causal graph is constructed based on the confounding factor, and the object graph is constructed based on the 3D position information of each entity in the to-be-detected video frame and the position prediction results of each entity at future moments;
[0022] An empty sequence information enhancement layer for extracting features from the superimposed graph in the empty sequence dimension;
[0023] A temporal sequence information aggregation layer for extracting features from the superimposed graph in the temporal sequence dimension according to the feature extraction result of the superimposed graph in the empty sequence dimension;
[0024] A spatio-temporal information transmission layer for predicting the 3D position information of each entity at the next moment according to the feature extraction result of the superimposed graph in the temporal sequence dimension.
[0025] According to a counterfactual prediction method for entity trajectories provided by the present invention, the expressions of the empty sequence information enhancement layer, the temporal sequence information aggregation layer, and the spatio-temporal information transmission layer are shown in Equations 2-4 respectively:
[0026]
[0027]
[0028]
[0029] In the formula, f(), respectively represent an empty sequence feature extraction function, a temporal sequence feature extraction function, and a spatio-temporal information transmission function; respectively represent the node corresponding to entity i and the edge between the nodes corresponding to entity i and entity j after the superimposed graph at time t is extracted by the empty sequence feature; respectively represent the node corresponding to entity i and the edge between the nodes corresponding to entity i and entity j after the superimposed graph is extracted by the temporal sequence feature; represents the predicted result of the 3D position information of entity i at time t + 1; represents the object graph at time t, represents the causal graph at time t.
[0030] The present invention also provides a counterfactual prediction device for entity trajectories, including:
[0031] A data acquisition module, configured to acquire a historical video sequence and a to-be-tested video frame during a game process; wherein, the to-be-tested video frame is image data corresponding to the moment of adding perturbation during the game process;
[0032] A position information extraction module, configured to input the historical video sequence and the to-be-tested video frame into a perception model to obtain 3D position information of each entity in the historical video sequence and the to-be-tested video frame;
[0033] A trajectory prediction module, configured to input the 3D position information of each entity in the historical video sequence and the to-be-tested video frame into a counterfactual prediction model to obtain a prediction result of the movement trajectories of each entity in the game after adding perturbation;
[0034] Wherein, the counterfactual prediction model includes a confounding factor estimation model and a confounding factor transmission model, the confounding factor estimation model is configured to obtain the confounding factors in the game according to the 3D position information of each entity in the historical video sequence; the confounding factor transmission model is configured to obtain the prediction result of the movement trajectories of each entity in the game after adding perturbation according to the 3D position information of each entity in the to-be-tested video frame and the confounding factors.
[0035] The present invention also provides an electronic device, 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, the steps of any one of the above-mentioned counterfactual prediction methods for entity trajectories are implemented.
[0036] The present invention also provides a non-transitory computer-readable storage medium, on which a computer program is stored, and when the computer program is executed by a processor, the steps of any one of the above-mentioned counterfactual prediction methods for entity trajectories are implemented.
[0037] The present invention also provides a computer program product, including a computer program, and when the computer program is executed by a processor, the steps of any one of the above-mentioned counterfactual prediction methods for entity trajectories are implemented.
[0038] The counterfactual prediction method and device for entity trajectories provided by the present invention obtain the confounding factors in the game scene according to the historical video sequence during the game process, and predict the movement trajectories of each entity after adding perturbation according to the confounding factors and the 3D position information of each entity in the game image corresponding to the moment of adding perturbation. The prediction process does not depend on the physical information in the game scene, can be applied to various different game scenes, does not need to design a specific prediction model for a specific game scene, has strong generalization ability, and effectively reduces the consumption of computing resources. Brief Description of the Drawings
[0039] In order to more clearly illustrate the technical solutions in the present invention or the prior art, the following will briefly introduce the drawings required in the description of the embodiments or the prior art. Obviously, the drawings in the following description are some embodiments of the present invention. For those of ordinary skill in the art, without creative efforts, other drawings can also be obtained based on these drawings.
[0040] Figure 1 is a schematic flowchart of the counterfactual prediction method for the entity trajectory provided by the present invention;
[0041] Figure 2 is a schematic structural diagram of the counterfactual prediction device for the entity trajectory provided by the present invention;
[0042] Figure 3 is a schematic structural diagram of the electronic device provided by the present invention. Detailed Embodiments
[0043] To make the objectives, technical solutions, and advantages of the present invention clearer, the following will clearly and completely describe the technical solutions in the present invention with reference to the drawings in the present invention. Obviously, the described embodiments are some, but not all, of the embodiments of the present invention. Based on the embodiments in the present invention, all other embodiments obtained by those of ordinary skill in the art without creative efforts fall within the protection scope of the present invention.
[0044] The following will describe Figure 1 the counterfactual prediction method for the entity trajectory of the present invention. As Figure 1 shown, the method includes:
[0045] S100. Obtain the historical video sequence and the to-be-tested video frame during the game process; wherein, the to-be-tested video frame is the image data corresponding to the moment of adding perturbation during the game process.
[0046] Specifically, the historical video sequence refers to a video sequence before adding perturbation in the game scene. The historical video sequence includes multiple entities, and the multiple entities may be in the same frame image or in different frame images. During the acquisition of the historical video sequence, all entities in the game should be included as much as possible. The to-be-tested video frame is the image data corresponding to the moment of adding perturbation during the game process. The perturbation is the change of relevant factors during the game process or some actions artificially imposed that can affect the movement trajectory of the entity. Among them, the historical video sequence and the to-be-tested video frame belong to the same game scene.
[0047] S200. Input the historical video sequence and the video frame to be measured into the perception model to obtain the 3D position information of each entity in the historical video sequence and the video frame to be measured.
[0048] Specifically, input the historical video sequence and the video frame to be measured into the trained perception model. The perception model processes the images in the historical video sequence frame by frame to obtain the 3D position information of all entities in each frame of the image. At the same time, the perception model extracts the 3D position information of each entity in the video frame to be measured. Here, there is no specific requirement for the specific structure of the perception model, as long as it can identify the 3D position information of the entity. For example, ResNet18 can be used as the backbone network.
[0049] S300. Input the 3D position information of each entity in the historical video sequence and the video frame to be measured into the counterfactual prediction model to obtain the prediction results of the movement trajectories of each entity in the game after perturbation.
[0050] Among them, the counterfactual prediction model includes a confounding factor estimation model and a confounding factor transmission model. The confounding factor estimation model is used to obtain the confounding factors in the game according to the 3D position information of each entity in the historical video sequence. The confounding factor transmission model is used to obtain the prediction results of the movement trajectories of each entity in the game after perturbation according to the 3D position information of each entity in the video frame to be measured and the confounding factors.
[0051] Specifically, input the 3D position information of each entity in each frame of the historical video sequence into the confounding factor estimation model to obtain the confounding factors in the game scene. Among them, the confounding factor is the game scene information that affects the movement trajectories of each entity in the historical video sequence and is difficult to directly observe. It has a specific meaning according to a specific game scene. For example, the friction coefficient of a wooden block and the deformation coefficient of a small ball. Input the obtained confounding factors and the 3D position information of each entity in the video frame to be measured into the confounding factor transmission model to obtain the prediction results of the 3D position information of the entity at each future moment. According to the prediction results of the 3D position information of the entity at each moment, obtain the prediction results of the movement trajectories of each entity after perturbation.
[0052] It can be seen that in the embodiment of the present invention, the confounding factors in the game scene are obtained according to the historical video sequence in the game process, and the movement trajectories of each entity after perturbation are predicted according to the confounding factors and the 3D position information of each entity in the game image corresponding to the perturbation moment. The prediction process does not depend on the physical information in the game scene, can be applied to various different game scenes, does not need to design a specific prediction model for a specific game scene, has strong generalization ability, and effectively reduces the consumption of computing resources.
[0053] Based on the above embodiments, the structure of the confounding factor estimation model includes:
[0054] An absolute position encoding layer, configured to calculate the absolute position information of each entity in the historical video sequence;
[0055] A global causal association attention layer, configured to model the causal relationship between each entity in the game by using a scaled dot-product self-attention mechanism according to the 3D position information and the absolute position information of each entity in the historical video sequence, and obtain the confounding factor in the game based on the causal relationship.
[0056] Specifically, the absolute position information of each entity in the historical video sequence represents the order information of each entity in each frame of the historical video sequence, such as red No. 1, blue No. 2. The causal relationship between each entity in the game is the mutual association between each entity in the game; in the self-attention mechanism, the order information of the entities has a great influence on the causal relationship between the entities. Therefore, through the 3D position information and the absolute position information of the entities, the causal relationship between each entity in the game can be accurately modeled.
[0057] It can be seen that in the embodiment of the present invention, the order information of each entity is characterized by calculating the absolute position information of each entity in the historical video sequence. Based on the 3D position information and the absolute position information of each entity, the causal relationship between each entity in the game is modeled by using a scaled dot-product self-attention mechanism, which can fully exploit and utilize the indirect causal chain between the entities to estimate the confounding factor, effectively improving the estimation ability of the counterfactual prediction model for the confounding factor in complex game scenarios; at the same time, only the 3D position information and the absolute position information of each entity in each frame of the historical video sequence are required in the process of estimating the confounding factor, without relying on the physical information in the game scene, and it can be applied to various different game scenes, without designing a specific prediction model for a specific game scene, with strong generalization ability and effectively reducing the consumption of computing resources.
[0058] Based on any of the above embodiments, calculating the absolute position information of each entity in the historical video sequence includes:
[0059] Obtaining the order information of each entity in the historical video sequence;
[0060] According to the order information of each entity in the historical video sequence, using a sine function to calculate the absolute position information of each entity.
[0061] Specifically, the order information of each entity in the historical video sequence, such as Red No. 1 and Blue No. 2, has different formats in different game scenarios and there is no fixed value range. Therefore, in the embodiments of the present invention, the sine function is used to perform absolute position encoding on the order information of each entity in the historical video sequence to obtain absolute position information, so that the absolute position information of each entity can all fall within the interval [-1, 1], ensuring that the absolute position information of each entity has the same format and value range, and can effectively reflect the order sequence between different entities.
[0062] Based on any of the above embodiments, a scaled dot product self-attention mechanism is used to model the causal relationship between each entity in the game, including:
[0063] The scaled dot product self-attention mechanism is used to calculate the correlation degree between each pair of entities to obtain the causal relationship between all entities in the historical video sequence; the calculation of the correlation degree is shown in Equation (1):
[0064]
[0065] In the formula, are the query vector, key vector, and value vector respectively. The query vector, key vector, and value vector are obtained by multiplying the 3D position matrix and / or the absolute position matrix with the corresponding weight matrices W qsi 、W krj 、W vrj respectively; the 3D position matrix and the absolute position matrix are used to store the 3D position information and absolute position information of each entity in the historical video sequence respectively; represents the correlation degree between entity i in video frame s and entity j in video frame r of the historical video sequence; d k represents the dimension of the key vector, and softmax() represents a probability-based multi-classification function; T represents the transpose of the matrix.
[0066] Specifically, existing confounding factor estimation methods ignore the causal associations between different entities in different frames and cannot effectively model the associations between entities, especially the associations between entities in long time series. In contrast, the embodiments of the present invention encode the association information between different entities in different frames of a long-distance video sequence based on the 3D position information and absolute position information of each entity in each frame image of the historical video sequence, mine and utilize indirect causal chains to model the associations between entities; according to Equation (1), in the process of calculating the association degree between each pair of entities using the scaled dot-product self-attention mechanism in the embodiments of the present invention, the inter-frame and intra-frame attention mechanisms are introduced, which can prompt the confounding factor estimation model to model the causal associations between entities in a long-distance video sequence, thereby effectively modeling the associations between entities in a long time series, further enhancing the ability to estimate confounding factors in complex game environments, and providing a data basis for accurately predicting the movement trajectories of each entity after adding perturbations in the game.
[0067] In addition, the embodiments of the present invention extend the construction of a confounding factor estimation model based on a transformer model, and use the scaled dot-product self-attention mechanism to calculate the associations between different objects. Through the self-attention mechanism, the 3D position information and absolute position information of each entity can be associated to calculate the association mechanism of each entity in different frames. Based on the self-attention mechanism of global causal association, it can help the confounding factor estimation model to more fully model the correlations between entities. The scaled dot-product self-attention mechanism is implemented using highly optimized matrix multiplication, with fast calculation speed and little occupied space, which can effectively improve the efficiency of confounding factor estimation.
[0068] Among them, the extended transformer-based model in the embodiments of the present invention includes an absolute position encoding block, an entity information encoding block, and an entity information decoding block; the absolute position encoding module, i.e., the absolute position encoding layer, is used to calculate the absolute position information of each entity using a sine function according to the sequence information of each entity in each frame image of the historical video sequence; the entity information encoding block fully mines the causal relationships between entities and outputs the potential causal relationships between entities; the entity information decoding block decodes the causal relationships and outputs the estimation result of the confounding factor in the game. Among them, both the entity information encoding block and the entity information decoding block include a multi-head self-attention layer, a feed-forward neural network layer, and a normalization layer connected in sequence, which are used to encode and decode the causal chain between objects, thereby improving the prediction accuracy of the confounding factor.
[0069] The existing methods for estimating confounding factors mainly rely on recurrent neural networks, which only consider updating the null-order information at the first moment and cannot utilize and update the null-order information in a timely manner during the subsequent process. In the embodiments of the present invention, the 3D position information and absolute position information of different entities in different frames of the historical video sequence are input into the transformer model for estimating the confounding factors, so that the transformer model can update the temporal and null-order information multiple times, and the temporal and null-order information are updated simultaneously, avoiding the defect that the existing methods for estimating confounding factors only update the null-order information at the first moment, and effectively improving the ability of the confounding factor estimation model to estimate the confounding factors in complex game scenarios.
[0070] Based on any of the above embodiments, the structure of the confounding factor transmission model includes:
[0071] A splicing layer for superimposing the causal graph and the object graph to obtain a superimposed graph; wherein, the causal graph is constructed based on the confounding factors, and the object graph is constructed based on the 3D position information of each entity in the video frame to be measured and the position prediction results of each entity at future moments.
[0072] A null-order information enhancement layer for extracting features of the superimposed graph in the null-order dimension.
[0073] A temporal information aggregation layer for extracting features of the superimposed graph in the temporal dimension according to the feature extraction result of the superimposed graph in the null-order dimension.
[0074] A spatio-temporal information transmission layer for predicting the 3D position information of each entity at the next moment according to the feature extraction result of the superimposed graph in the temporal dimension.
[0075] Specifically, in the existing counterfactual prediction model designed in the game system, the forward propagation sub-module often cannot make full and effective use of the confounding factors in the complex game environment that have been estimated, and the number of updates of the null-order information is limited, and it is easy to ignore the potential correlation information between entities. In addition, the existing counterfactual prediction model also lacks the exploration and understanding of the potential causal graph, which also leads to inaccurate simulation prediction results of the object trajectories in the final game system.
[0076] In the embodiment of the present invention, the causal graph and the object graph are superimposed through a splicing layer, wherein the causal graph is constructed based on confounding factors; the object graph is obtained by splicing the 3D position information of each entity in the video frame to be measured and the position prediction results of each entity at each future moment, that is, the object graph is composed of the 3D position information of each entity. For example, after adding perturbations, when predicting the 3D position information of the entity in the second frame image, the object graph is constructed according to the 3D position information of each entity in the video frame to be measured (the first frame, that is, the video frame with perturbations). When predicting the 3D position information of the entity in the third frame image, the 3D position information of each entity in the video frame to be measured and the 3D position information of each entity in the predicted second frame image are spliced. By superimposing the causal graph and the object graph, a superimposed graph is obtained, and the object graph and the causal graph information are efficiently encoded and transmitted through the null sequence information enhancement layer, the temporal sequence information aggregation layer, and the spatio-temporal information transmission layer, realizing further extraction and enhancement of the correlation between entities, effectively improving the ability of the confounding factor transmission model to understand and utilize confounding factors, and further improving the prediction accuracy of the counterfactual prediction model for the movement trajectories of each entity in the game.
[0077] Based on any of the above embodiments, the expressions of the null sequence information enhancement layer, the temporal sequence information aggregation layer, and the spatio-temporal information transmission layer are respectively shown in formulas (2)-(4):
[0078]
[0079]
[0080]
[0081] In the formulas, f(), respectively represent the null sequence feature extraction function, the temporal sequence feature extraction function, and the spatio-temporal information transmission function; respectively represent the node corresponding to entity i and the edge between the nodes corresponding to entity i and entity j after the superimposed graph at time t is extracted by the null sequence feature; respectively represent the node corresponding to entity i and the edge between the nodes corresponding to entity i and entity j after the superimposed graph is extracted by the temporal sequence feature; represents the predicted result of the 3D position information of entity i at time t + 1; represents the object graph at time t, represents the causal graph at time t.
[0082] Specifically, the object graph is composed of the 3D position information (position coordinates) of each entity. The nodes on the object graph are composed of the embedding encodings of the 3D position information of each entity. The edges on the object graph are obtained by stacking adjacent nodes, and the stacking form is shown in Equation (5); the nodes in the causal graph represent confounding factor information, and the edges represent the contact information between entities, which are learnable vectors and are randomly initialized. The empty sequence feature extraction function Spatio-temporal information transmission function Both adopt the traditional graph neural network structure, and the temporal feature extraction function f() adopts the GRU (Gate Recurrent Unit) structure.
[0083]
[0084] Among them, is the edge between the nodes corresponding to entity i and entity j, are the nodes corresponding to entity i and entity j respectively.
[0085] During the prediction process of the 3D position information of entities at different times, Equations (2) - (4) are continuously iteratively used, so that the object graph information and the causal graph information are fully utilized and updated in both the empty sequence and temporal dimensions, thereby ensuring that the confounding factor transmission model can effectively utilize the potential causal chains between entities and the already estimated confounding factors, thus effectively improving the accuracy of the entity trajectory prediction results in the game.
[0086] In addition, during the training process of the counterfactual prediction model, based on each sample in the obtained training sample set in turn, the counterfactual prediction model is trained through the constructed loss function until the model converges or reaches the preset number of training times, thereby obtaining the trained counterfactual prediction model. Among them, the loss function L e2e is shown in Equation (6):
[0087]
[0088] In the formula, are the predicted 3D position information and the true 3D position information of entity m at time t respectively; T represents the total duration of the training samples (that is, the training samples include T frames of images); M is the total number of entities in the training samples; L mse () represents the mean squared error loss.
[0089] It can be seen that the embodiment of the present invention is based on a self-attention mechanism of global causal association, which helps the counterfactual prediction model to more fully model the correlation between objects. Next, in order to strengthen the encoding and utilization ability of the counterfactual prediction model for confounding factors, the embodiment of the present invention proposes a confounding factor transmission architecture, which significantly improves the model's ability to utilize confounding factors and enhances the robustness of the model, enabling the counterfactual prediction model to be better generalized and deployed in different game systems and perform dynamic simulations on the game systems, ultimately contributing to the improvement of the prediction accuracy of the entity movement trajectories in the game system.
[0090] The counterfactual prediction device for entity trajectories provided by the present invention will be described below. The counterfactual prediction device for entity trajectories described below can be mutually referred to the counterfactual prediction method for entity trajectories described above. As Figure 2 shown, the device includes:
[0091] A data acquisition module 210, configured to acquire a historical video sequence and a to-be-tested video frame during a game process; wherein, the to-be-tested video frame is image data corresponding to the moment of adding perturbations during the game process;
[0092] A position information extraction module 220, configured to input the historical video sequence and the to-be-tested video frame into a perception model to obtain 3D position information of each entity in the historical video sequence and the to-be-tested video frame;
[0093] A trajectory prediction module 230, configured to input the 3D position information of each entity in the historical video sequence and the to-be-tested video frame into a counterfactual prediction model to obtain a prediction result of the movement trajectories of each entity in the game after adding perturbations;
[0094] Wherein, the counterfactual prediction model includes a confounding factor estimation model and a confounding factor transmission model. The confounding factor estimation model is configured to obtain the confounding factors in the game according to the 3D position information of each entity in the historical video sequence; the confounding factor transmission model is configured to obtain the prediction result of the movement trajectories of each entity in the game after adding perturbations according to the 3D position information of each entity in the to-be-tested video frame and the confounding factors.
[0095] Based on the above embodiment, the structure of the confounding factor estimation model includes:
[0096] An absolute position encoding layer, configured to calculate the absolute position information of each entity in the historical video sequence;
[0097] A global causal association attention layer, configured to model the causal relationship between each entity in the game by using a scaled dot-product self-attention mechanism according to the 3D position information and the absolute position information of each entity in the historical video sequence, and obtain the confounding factors in the game based on the causal relationship.
[0098] Based on any of the above embodiments, calculating the absolute position information of each of the entities in the historical video sequence includes:
[0099] Obtaining the sequence information of each of the entities in the historical video sequence;
[0100] According to the sequence information of each of the entities in the historical video sequence, using the sine function to calculate the absolute position information of each of the entities.
[0101] Based on any of the above embodiments, using the scaled dot-product self-attention mechanism to model the causal relationship between each entity in the game includes:
[0102] Using the scaled dot-product self-attention mechanism to calculate the correlation degree between each pair of the entities, and obtaining the causal relationship between all the entities in the historical video sequence; the calculation of the correlation degree is shown in Equation (1):
[0103]
[0104] In the formula, are the query vector, the key vector, and the value vector respectively, and the query vector, the key vector, and the value vector are respectively obtained by multiplying the 3D position matrix and / or the absolute position matrix with the corresponding weight matrices W qsi 、W krj 、W vrj ; the 3D position matrix and the absolute position matrix are respectively used to store the 3D position information and the absolute position information of each entity in the historical video sequence; represents the correlation degree between entity i in video frame s and entity j in video frame r of the historical video sequence; d k represents the dimension of the key vector, and softmax() represents the probability-based multi-classification function.
[0105] Based on any of the above embodiments, the structure of the confounding factor transmission model includes:
[0106] A splicing layer for superimposing the causal graph and the object graph to obtain a superimposed graph; wherein, the causal graph is constructed based on the confounding factor, and the object graph is constructed based on the 3D position information of each entity in the to-be-tested video frame and the position prediction results of each entity at future moments;
[0107] An empty sequence information enhancement layer for extracting features of the superimposed graph in the empty sequence dimension;
[0108] A temporal sequence information aggregation layer for extracting features of the superimposed graph in the temporal sequence dimension according to the feature extraction result of the superimposed graph in the empty sequence dimension;
[0109] A spatio-temporal information transmission layer, configured to predict the 3D position information of each entity at the next moment according to the feature extraction result of the superimposed graph in the temporal dimension.
[0110] Based on any of the above embodiments, the expressions of the spatial-order information enhancement layer, the temporal information aggregation layer, and the spatio-temporal information transmission layer are respectively shown in Formulas (2)-(4):
[0111]
[0112]
[0113]
[0114] In the formulas, f(), respectively represent a spatial-order feature extraction function, a temporal feature extraction function, and a spatio-temporal information transmission function; respectively represent the node corresponding to entity i and the edge between the nodes corresponding to entity i and entity j after the superimposed graph at time t is subjected to spatial-order feature extraction; respectively represent the node corresponding to entity i and the edge between the nodes corresponding to entity i and entity j after the superimposed graph is subjected to temporal feature extraction; represents the prediction result of the 3D position information of entity i at time t + 1; represents the object graph at time t, represents the causal graph at time t.
[0115] Figure 3 Illustrates a schematic diagram of the entity structure of an electronic device, as Figure 3 shown. The electronic device may include: a processor 310, a communication interface 320, a memory 330, and a communication bus 340. Among them, the processor 310, the communication interface 320, and the memory 330 complete communication with each other through the communication bus 340. The processor 310 may call the logical instructions in the memory 330 to execute the counterfactual prediction method for the entity trajectory, and the method includes: obtaining a historical video sequence and a to-be-tested video frame during the game process; wherein, the to-be-tested video frame is the image data corresponding to the moment of adding perturbation during the game process;
[0116] Inputting the historical video sequence and the to-be-tested video frame into a perception model to obtain the 3D position information of each entity in the historical video sequence and the to-be-tested video frame;
[0117] Input the 3D position information of each of the entities in the historical video sequence and the video frame to be measured into a counterfactual prediction model to obtain the prediction results of the movement trajectories of each of the entities in the game after perturbation;
[0118] Among them, the counterfactual prediction model includes a confounding factor estimation model and a confounding factor transmission model. The confounding factor estimation model is used to obtain the confounding factors in the game according to the 3D position information of each of the entities in the historical video sequence; the confounding factor transmission model is used to obtain the prediction results of the movement trajectories of each of the entities in the game after perturbation according to the 3D position information of each of the entities in the video frame to be measured and the confounding factors.
[0119] In addition, when the logical instructions in the above-mentioned memory 330 can be implemented in the form of software functional units and sold or used as an independent product, they can be stored in a computer-readable storage medium. Based on such an understanding, the technical solution of the present invention, in essence, or the part that contributes to the prior art, or a part of this technical solution, can be embodied in the form of a software product. This computer software product is stored in a storage medium and includes several instructions for causing a computer device (which can be a personal computer, a server, or a network device, etc.) to execute all or part of the steps of the methods described in various embodiments of the present invention. The foregoing storage medium includes: various media such as USB flash drives, mobile hard disks, read-only memories (ROM, Read-Only Memory), random access memories (RAM, Random Access Memory), magnetic disks, or optical discs that can store program codes.
[0120] On the other hand, the present invention also provides a computer program product. The computer program product includes a computer program that can be stored on a non-transitory computer-readable storage medium. When the computer program is executed by a processor, the computer can execute the counterfactual prediction method for the entity trajectory provided by the above-mentioned various methods. The method includes: obtaining a historical video sequence and a video frame to be measured during the game process; among them, the video frame to be measured is the image data corresponding to the perturbation moment during the game process;
[0121] Input the historical video sequence and the video frame to be measured into a perception model to obtain the 3D position information of each entity in the historical video sequence and the video frame to be measured;
[0122] Input the 3D position information of each of the entities in the historical video sequence and the video frame to be measured into a counterfactual prediction model to obtain the prediction results of the movement trajectories of each of the entities in the game after perturbation;
[0123] Among them, the counterfactual prediction model includes a confounding factor estimation model and a confounding factor transmission model. The confounding factor estimation model is used to obtain the confounding factors in the game according to the 3D position information of each entity in the historical video sequence; the confounding factor transmission model is used to obtain the prediction result of the motion trajectory of each entity in the game after perturbation according to the 3D position information of each entity in the to-be-tested video frame and the confounding factors.
[0124] In another aspect, the present invention also provides a non-transitory computer-readable storage medium, on which a computer program is stored. When the computer program is executed by a processor, it is configured to execute the counterfactual prediction method for entity trajectories provided by the above-mentioned various methods. The method includes: obtaining a historical video sequence and a to-be-tested video frame during the game process; among them, the to-be-tested video frame is the image data corresponding to the perturbation moment during the game process.
[0125] Input the historical video sequence and the to-be-tested video frame into the perception model to obtain the 3D position information of each entity in the historical video sequence and the to-be-tested video frame.
[0126] Input the 3D position information of each entity in the historical video sequence and the to-be-tested video frame into the counterfactual prediction model to obtain the prediction result of the motion trajectory of each entity in the game after perturbation.
[0127] Among them, the counterfactual prediction model includes a confounding factor estimation model and a confounding factor transmission model. The confounding factor estimation model is used to obtain the confounding factors in the game according to the 3D position information of each entity in the historical video sequence; the confounding factor transmission model is used to obtain the prediction result of the motion trajectory of each entity in the game after perturbation according to the 3D position information of each entity in the to-be-tested video frame and the confounding factors.
[0128] The device embodiments described above are merely illustrative. The units described as separate components may or may not be physically separated. The components shown as units may or may not be physical units, that is, they may be located in one place or distributed to multiple network units. Some or all of the modules can be selected according to actual needs to achieve the purpose of the solution of this embodiment. Those of ordinary skill in the art can understand and implement it without creative efforts.
[0129] Through the description of the above embodiments, those skilled in the art can clearly understand that each embodiment can be implemented by means of software plus a necessary general hardware platform, and of course, it can also be implemented by hardware. Based on such an understanding, the above technical solution, in essence, or the part that contributes to the prior art can be embodied in the form of a software product. This computer software product can be stored in a computer-readable storage medium, such as ROM / RAM, magnetic disk, optical disk, etc., and includes several instructions for causing a computer device (which can be a personal computer, a server, or a network device, etc.) to execute the methods described in each embodiment or some parts of the embodiments.
[0130] Finally, it should be noted that the above embodiments are only used to illustrate the technical solutions of the present invention and are not intended to limit them. Although the present invention has been described in detail with reference to the foregoing embodiments, those of ordinary skill in the art should understand that they can still modify the technical solutions described in the foregoing embodiments, or perform equivalent replacements for some of the technical features. And these modifications or replacements do not cause the essence of the corresponding technical solutions to deviate from the spirit and scope of the technical solutions of the embodiments of the present invention.
Claims
1. A counterfactual prediction method for entity trajectories, characterized in that Including: Obtain a historical video sequence and a to-be-tested video frame during the game process; wherein, the to-be-tested video frame is image data corresponding to the moment of adding perturbation during the game process; Input the historical video sequence and the to-be-tested video frame into a perception model to obtain 3D position information of each entity in the historical video sequence and the to-be-tested video frame; Input the 3D position information of each entity in the historical video sequence and the to-be-tested video frame into a counterfactual prediction model to obtain a prediction result of the movement trajectory of each entity in the game after adding perturbation; Wherein, the counterfactual prediction model includes a confounding factor estimation model and a confounding factor transmission model. The confounding factor estimation model is used to obtain the confounding factor in the game according to the 3D position information of each entity in the historical video sequence; the confounding factor transmission model is used to obtain the prediction result of the movement trajectory of each entity in the game after adding perturbation according to the 3D position information of each entity in the to-be-tested video frame and the confounding factor; The structure of the confounding factor estimation model includes: An absolute position encoding layer for calculating the absolute position information of each entity in the historical video sequence; A global causal correlation attention layer for modeling the causal relationship between each entity in the game by using a scaled dot-product self-attention mechanism according to the 3D position information and absolute position information of each entity in the historical video sequence, and obtaining the confounding factor in the game based on the causal relationship; The structure of the confounding factor transmission model includes: A splicing layer for superimposing a causal graph and an object graph to obtain a superimposed graph; wherein, the causal graph is constructed based on the confounding factor, and the object graph is constructed based on the 3D position information of each entity in the to-be-tested video frame and the position prediction result of each entity at each future moment; A null sequence information enhancement layer for extracting features of the superimposed graph in the null sequence dimension; A temporal sequence information aggregation layer for extracting features of the superimposed graph in the temporal sequence dimension according to the feature extraction result of the superimposed graph in the null sequence dimension; A spatio-temporal information transmission layer for predicting the 3D position information of each entity at the next moment according to the feature extraction result of the superimposed graph in the temporal sequence dimension.
2. The counterfactual prediction method for an entity trajectory according to claim 1, wherein The calculation of the absolute position information of each entity in the historical video sequence includes: Obtain the sequence information of each entity in the historical video sequence; According to the sequence information of each entity in the historical video sequence, use the sine function to calculate the absolute position information of each entity.
3. The counterfactual prediction method for an entity trajectory according to claim 1, wherein The modeling of the causal relationship between each entity in the game by using a scaled dot-product self-attention mechanism includes: Use a scaled dot-product self-attention mechanism to calculate the correlation degree between each pair of entities to obtain the causal relationship between all entities in the historical video sequence; the calculation of the correlation degree is shown in Equation 1: In the formula, are the query vector, key vector, and value vector respectively. The query vector, key vector, and value vector are obtained by multiplying with the corresponding weight matrices W qsi , W krj , W vrj respectively; the 3D position matrix and the absolute position matrix are used to store the 3D position information and absolute position information of each entity in the historical video sequence; represents the correlation degree between entity i in video frame s and entity j in video frame r of the historical video sequence; d k represents the dimension of the key vector, and softmax() represents a probability-based multi-classification function.
4. A counterfactual prediction method for an entity trajectory according to claim 1, wherein The expressions of the null sequence information enhancement layer, the temporal sequence information aggregation layer, and the spatio-temporal information transmission layer are shown in Equations 2-4 respectively: In the formula, f(), respectively represent the empty sequence feature extraction function, the time sequence feature extraction function, and the spatio-temporal information transmission function; respectively represent the nodes corresponding to entity i and the edges between the nodes corresponding to entity i and entity j after the superimposed graph at time t is subjected to empty sequence feature extraction; respectively represent the nodes corresponding to entity i and the edges between the nodes corresponding to entity i and entity j after the superimposed graph is subjected to time sequence feature extraction; represents the prediction result of the 3D position information of entity i at time t+1; represents the object graph at time t, represents the causal graph at time t.
5. A counterfactual prediction device for an entity trajectory, characterized in that, Including: A data acquisition module, configured to acquire a historical video sequence and a to-be-tested video frame during a game process; wherein, the to-be-tested video frame is image data corresponding to the moment of adding perturbation during the game process; A position information extraction module, configured to input the historical video sequence and the to-be-tested video frame into a perception model to obtain 3D position information of each entity in the historical video sequence and the to-be-tested video frame; A trajectory prediction module, configured to input the 3D position information of each entity in the historical video sequence and the to-be-tested video frame into a counterfactual prediction model to obtain a prediction result of the motion trajectory of each entity in the game after adding perturbation; Wherein, the counterfactual prediction model includes a confounding factor estimation model and a confounding factor transmission model, and the confounding factor estimation model is configured to obtain the confounding factor in the game according to the 3D position information of each entity in the historical video sequence; the confounding factor transmission model is configured to obtain the prediction result of the motion trajectory of each entity in the game after adding perturbation according to the 3D position information of each entity in the to-be-tested video frame and the confounding factor; The structure of the confounding factor estimation model includes: An absolute position encoding layer, configured to calculate the absolute position information of each entity in the historical video sequence; A global causal correlation attention layer, configured to model the causal relationship between each entity in the game by using a scaled dot-product self-attention mechanism according to the 3D position information and the absolute position information of each entity in the historical video sequence, and obtain the confounding factor in the game based on the causal relationship; The structure of the confounding factor transmission model includes: A splicing layer, configured to superimpose a causal graph and an object graph to obtain a superimposed graph; wherein, the causal graph is constructed based on the confounding factor, and the object graph is constructed based on the 3D position information of each entity in the to-be-tested video frame and the position prediction result of each entity at each future moment; An empty sequence information enhancement layer, configured to extract features from the superimposed graph in the empty sequence dimension; A temporal sequence information aggregation layer, configured to extract features from the superimposed graph in the temporal sequence dimension according to the feature extraction result of the superimposed graph in the empty sequence dimension; A spatio-temporal information transmission layer, configured to predict the 3D position information of each entity at the next moment according to the feature extraction result of the superimposed graph in the temporal sequence dimension.
6. An electronic device, comprising a memory, a processor, and a computer program stored on the memory and executable on the processor, characterized in that, When the processor executes the program, it implements the steps of the counterfactual prediction method for the entity trajectory according to any one of claims 1 to 4.
7. A non-transitory computer-readable storage medium having a computer program stored thereon, characterized in that, When the computer program is executed by the processor, it implements the steps of the counterfactual prediction method for the entity trajectory according to any one of claims 1 to 4.
8. A computer program product, comprising a computer program, characterized in that, When the computer program is executed by the processor, it implements the steps of the counterfactual prediction method for the entity trajectory according to any one of claims 1 to 4.