Efficient real-world experimentation using causal inference models
An experiment design is determined for a plurality of physical experiment entities, based on a training loss that is dependent on a critic function and an action parameter individually associated with each physical experiment entity, with the aim of increasing (e.g., optimizing) information gain with respect to a plurality of test entities. The training loss encodes a predicted information gain between a predicted experiment outcome and a predicted test quantity. The predicted experiment outcome associated therewith is sampled from a joint probability distribution based on an entity context. A numerical output is computed using the critic function applied to the predicted experiment outcome and the predicted test quantity.
1 . A computer-implemented method comprising:
receiving, for a first physical experiment entity of a plurality of physical experiment entities, a first context that individually characterizes the first physical experiment entity;
receiving, for a first physical test entity of a plurality of physical test entities, a second context that individually characterizes the first physical test entity;
computing, for the first physical experiment entity, an initial value of an action parameter that is individually associated with the first physical experiment entity;
computing, in a sequence of training iterations, an updated value of the action parameter, the sequence comprising:
an initial training iteration based on the initial value of the action parameter, and
a subsequent training iteration, for updating the action parameter, based on a value of the action parameter computed in a preceding training iteration of the sequence;
wherein a training iteration of the sequence comprises:
sampling, for the first physical experiment entity, a first predicted experiment outcome from a joint probability distribution, the first predicted experiment outcome being associated with the first physical experiment entity;
sampling, for the first physical test entity, a first predicted test quantity from the joint probability distribution, the first predicted test quantity being associated with the first physical test entity, wherein the joint probability distribution is based on the first context, the second context, and the action parameter,
computing a numerical output using a critic function applied to the first predicted experiment outcome and the first predicted test quantity;
updating, using the numerical output, the action parameter based on a training loss that:
is dependent on the critic function and the action parameter, and
encodes a predicted information gain between the first predicted experiment outcome and the first predicted test quantity;
outputting, for the first physical experiment entity, an indication of a first real-world experiment action, as defined by the updated value of the action parameter associated with the first physical experiment entity as computed in a final training iteration in the sequence of training iterations;
determining a real-world test action individually associated with the first physical test entity based on an outcome of the first real-world experiment action performed on the first physical experiment entity; and
performing the real-world test action on the first physical test entity.
2 . The computer-implemented method of claim 1
wherein the joint probability distribution comprises a Bayesian model of a joint probability distribution over predicted experiment outcomes and predicted test quantities, the Bayesian model being parameterised by a world parameter vector, such that different world models are obtained by sampling different values of the world parameter vector, the joint probability distribution being conditional on the world parameter vector, the action parameter of the first physical experiment entity, the first context, and the second context.
3 . The computer-implemented method of claim 1 , comprising evaluating a real-world test quantity associated with the first physical test entity based on an outcome of the performing of the first real-world experiment action on the first physical experiment entity.
4 . The computer-implemented method of claim 1 , wherein:
the critic function is parameterised by a critic parameter and the computer-implemented method includes computing an initial value of the critic parameter;
the initial training iteration is based on the initial value of the critic parameter; and
the training iterations comprises computing, based on the training loss, an updated value of the action parameter, such that the critic parameter and the action parameter are jointly optimized across the sequence of training iterations.
5 . The computer-implemented method of claim 4 , wherein the critic parameter and the action parameter are jointly optimized across the sequence of training iterations via gradient-based optimization of the training loss.
6 . The computer-implemented method of claim 5 , wherein the sequence of training iterations is extended until a training limit, convergence threshold, or maximum number of iterations has been reached.
7 . The computer-implemented method of claim 5 , wherein at least one of the action parameter and the critic parameter is non-scalar.
8 . The computer-implemented method of claim 5 , wherein the critic parameter comprises a neural network weight.
9 . The computer-implemented method of claim 1 , wherein:
in the training iteration, a world state is sampled from a world-state distribution; and
the first predicted experiment outcome associated with the first physical experiment entity and the first predicted test quantity are sampled based on the sampled world state.
10 . The computer-implemented method of claim 9 , wherein:
in the training iteration, multiple world states are sampled from the world-state distribution; and
for the first physical experiment entity and for the first physical test entity, predicted experiment outcomes and predicted test quantities are sampled for the multiple world states.
11 . The computer-implemented method of claim 1 , further comprising:
receiving, for a second physical experiment entity of the plurality of physical experiment entities, a third context that individually characterizes the second physical experiment entity;
receiving, for a second physical test entity of the plurality of physical test entities, a fourth context that individually characterizes the second physical test entity;
computing, for the second physical experiment entity, an initial value of a second action parameter that is individually associated with the second physical experiment entity;
computing, in a second sequence of training iterations, an updated value of the second action parameter, the second sequence comprising:
an initial training iteration based on the initial value of the second action parameter, and
a subsequent training iteration, for updating the second action parameter, based on a value of the second action parameter computed in a preceding training iteration of the second sequence;
wherein a training iteration of the second sequence comprises:
sampling, for the second physical experiment entity, a second predicted experiment outcome from the joint probability distribution, the second predicted experiment outcome being associated with the second physical experiment entity;
sampling, for the second physical test entity, a second predicted test quantity from the joint probability distribution, the second predicted test quantity being associated with the second physical test entity, wherein the joint probability distribution is based on the third context, the fourth context, and the second action parameter;
computing a second numerical output using the critic function applied to the second predicted experiment outcome and the second predicted test quantity; and
updating, using the second numerical output, the second action parameter based on a second training loss that is dependent on the critic function and the second action parameter and that encodes a predicted information gain between the second predicted experiment outcome and the second predicted test quantity;
outputting, for the second physical experiment entity, an indication of a second real-world experiment action, as defined by the updated value of the second action parameter associated with the second physical experiment entity as computed in a final training iteration of the sequence of training iterations; and
performing a second real-world test action individually associated with the second physical test entity on the second physical test entity, wherein the second real-world test action is based on the first and second real-world experiment actions performed on the first and second physical experiment entities, respectively.
12 . The computer-implemented method of claim 1 , wherein the method is used to design a clinical trial, and the first physical experiment entity and the first physical test entity are living beings.
13 . The computer-implemented method of claim 1 , wherein the first physical experiment entity and the first physical test entity are particular physical configurations of a machine.
14 . The computer-implemented method of claim 1 , wherein the first predicted experiment outcome comprises a predicted reward associated with the first physical experiment entity, and the first predicted test quantity comprises a predicted maximum reward associated with the first physical test entity.
15 . The computer-implemented method of claim 1 , further comprising: causing a control signal to be transmitted to an actuator or device configured to implement the first real-world experiment action.
16 . A computer system comprising:
a memory embodying computer-readable instructions; and
a processor coupled to the memory and configured to execute the computer-readable instructions, the computer-readable instructions configured to cause the processor to:
receive, for a physical test entity of a plurality of physical test entities, a first context that individually characterizes the physical test entity;
receive, for a physical experiment entity of a plurality of physical experiment entities, a second context that individually characterizes the physical experiment entity;
determine, for the physical experiment entity, a first action parameter that is individually associated with the physical experiment entity;
sample a first predicted experiment outcome associated with the physical experiment entity and a predicted test quantity associated with the physical test entity from a first joint probability distribution based on: the first context, the second context, and the first action parameter of the physical experiment entity; sample a first predicted experiment outcome associated with the physical experiment entity and a first predicted test quantity associated with the physical test entity from a first joint probability distribution based on: the first context, the second context, and the first action parameter of the physical experiment entity;
compute a first numerical output using a critic function applied to the first predicted experiment outcome associated with the physical experiment entity and the first predicted test quantity associated with the physical test entity;
determine, for the physical experiment entity, a second action parameter individually associated with the physical experiment entity based on a training loss applied to the first numerical output computed using the critic function and the first action parameter of the physical experiment entity, the training loss applied to the first numerical output and the first action parameter encoding a predicted information gain between the first predicted experiment outcome associated with the physical experiment entity and the first predicted test quantity associated with the physical test entity;
sample a second predicted experiment outcome associated with the physical experiment entity and a second predicted test quantity associated with the physical test entity from a second joint probability distribution based on: the first context, the second context, and the second action parameter of the physical experiment entity;
compute a second numerical output using the critic function applied to the second predicted experiment outcome associated with the physical experiment entity and the second predicted test quantity associated with the physical test entity;
determine, for the physical experiment entity, a third action parameter individually associated with the physical experiment entity based on the training loss applied to the second numerical output computed using the critic function and the second action parameter of the physical experiment entity, the training loss applied to the second numerical output and the second action parameter encoding a predicted information gain between the second predicted experiment outcome associated with the physical experiment entity and the second predicted test quantity associated with the physical test entity;
output an experiment design based on the third action parameter, wherein a real-world experiment action according to the experiment design is performed on the physical experiment entity to obtain a real-world experiment result;
determine a real-world test action individually associated with the physical test entity by updating a probability distribution of outcomes using the real-world experiment result and applying the updated probability distribution to the first context; and
cause the real-world test action to be performed on the physical test entity.
17 . The computer system of claim 16 , wherein the computer-readable instructions are configured to cause the processor to:
sample a third predicted experiment outcome associated the physical experiment entity and a third predicted test quantity associated with the physical test entity from a third joint probability distribution based on: the first context, the second context, and the third action parameter of the physical experiment entity;
compute a third numerical output using the critic function applied to the third predicted experiment outcome associated with the physical experiment entity and the third predicted test quantity associated with the physical test entity; and
determine, for the physical experiment entity, a fourth action parameter individually associated with the physical experiment entity based on the training loss applied to the third numerical output computed using the critic function and the third action parameter of the physical experiment entity, the training loss applied to the third numerical output and the third action parameter encoding a predicted information gain between the third predicted experiment outcome associated with the physical experiment entity and the third predicted test quantity associated with the physical test entity; and
wherein the experiment design is additionally based on the fourth action parameter.
18 . The computer system of claim 16 , wherein the computer-readable instructions are configured to cause the processor to output the experiment design via a graphical user interface associated with the computer system.
19 . The computer system of claim 16 , wherein:
the critic function is parameterised by a first critic parameter, wherein the first numerical output is computed using the first critic parameter; and
a second critic parameter is computed based on the training loss applied to the first numerical output and the second action parameter of the physical experiment entity, wherein the second numerical output is computed using the critic function applied to the second predicted experiment outcome, the second predicted test quantity, and the second critic parameter.
20 . Computer-readable storage media embodying computer-readable instructions configured, when executed on a computer processor, to cause the computer processor to carry out operations comprising:
computing an initial value of a critic parameter;
computing, for a physical experiment entity of a plurality of experimental entities, an initial value of an action parameter individually associated with the physical experiment entity;
computing, in a sequence of training iterations, an updated value of the critic parameter and an updated value of the action parameter, the sequence including:
an initial training iteration based on the initial value of the critic parameter and the initial value of the action parameter, and
a subsequent training iteration for updating the critic parameter and the action parameter, based on a value of the action parameter computed in a preceding one of the training iterations,
wherein a training iteration of the sequence comprises:
sampling, for the physical experiment entity, a predicted experiment outcome from a joint probability distribution, the predicted experiment outcome being associated with the physical experiment entity;
sampling, for a physical test entity, a predicted test quantity associated with physical test entity, wherein the joint probability distribution is based on: a context of the physical experiment entity, a context of the physical test entity, and the action parameter,
computing a numerical output using a critic function parameterised by the critic parameter and applied to the predicted experiment outcome and the predicted test quantity, and
updating, using the numerical output, the critic parameter and the action parameter based on a training loss that is dependent on the critic function and the action parameter and that encodes a predicted information gain between the predicted experiment outcome and the predicted test quantity; updating, using the numerical output, the critic parameter and the action parameter based on a training loss that is dependent on the critic function and the action parameter and that encodes a predicted information gain between the predicted experiment outcomes and the predicted test quantity;
outputting, for the physical experiment entity, an indication of a real-world experiment action defined by the updated value of the action parameter;
determining a real-world test action individually associated with the physical test entity based on an outcome of the real-world experiment action performed on the physical experiment entity; and
causing the real-world test action to be performed on the physical test entity.