Machine Theory of Mind

Neil C. Rabinowitz, Frank Perbet, H. Francis Song, Chiyuan Zhang, S. M. Ali Eslami, Matthew Botvinick

Introduction

For all the excitement surrounding deep learning and deep reinforcement learning at present, there is a concern from some quarters that our understanding of these systems is lagging behind. Neural networks are regularly described as opaque, uninterpretable black-boxes. Even if we have a complete description of their weights, it’s hard to get a handle on what patterns they’re exploiting, and where they might go wrong. As artificial agents enter the human world, the demand that we be able to understand them is growing louder.

Let us stop and ask: what does it actually mean to “understand” another agent? As humans, we face this challenge every day, as we engage with other humans whose latent characteristics, latent states, and computational processes are almost entirely inaccessible. Yet we function with remarkable adeptness. We can make predictions about strangers’ future behaviour, and infer what information they have about the world; we plan our interactions with others, and establish efficient and effective communication.

A salient feature of these “understandings” of other agents is that they make little to no reference to the agents’ true underlying structure. We do not typically attempt to estimate the activity of others’ neurons, infer the connectivity of their prefrontal cortices, or plan interactions with a detailed approximation of the dynamics of others’ hippocampal maps. A prominent argument from cognitive psychology is that our social reasoning instead relies on high-level models of other agents Gopnik & Wellman (1992). These models engage abstractions which do not describe the detailed physical mechanisms underlying observed behaviour; instead, we represent the mental states of others, such as their desires, beliefs, and intentions. This ability is typically described as our Theory of Mind Premack & Woodruff (1978). While we may also, in some cases, leverage our own minds to simulate others’ (e.g. Gordon, 1986; Gallese & Goldman, 1998), our ultimate human understanding of other agents is not measured by a 1-1 correspondence between our models and the mechanistic ground truth, but instead by how much these models afford for tasks such as prediction and planning Dennett (1991).

In this paper, we take inspiration from human Theory of Mind, and seek to build a system which learns to model other agents. We describe this as a Machine Theory of Mind. Our goal is not to assert a generative model of agents’ behaviour and an algorithm to invert it. Rather, we focus on the problem of how an observer could learn autonomously how to model other agents using limited data Botvinick et al. (2017). This distinguishes our work from previous literature, which has relied on hand-crafted models of agents as noisy-rational planners – e.g. using inverse RL Ng et al. (2000); Abbeel & Ng (2004), Bayesian inference Lucas et al. (2014); Evans et al. (2016), Bayesian Theory of Mind Baker et al. (2011); Jara-Ettinger et al. (2016); Baker et al. (2017) or game theory Camerer et al. (2004); Yoshida et al. (2008); Camerer (2010); Lanctot et al. (2017). In contrast, we learn the agent models, and how to do inference on them, from scratch, via meta-learning.

Building a rich, flexible, and performant Machine Theory of Mind may well be a grand challenge for AI. We are not trying to solve all of this here. A main message of this paper is that many of the initial challenges of building a ToM can be cast as simple learning problems when they are formulated in the right way. Our work here is an exercise in figuring out these simple formulations.

There are many potential applications for this work. Learning rich models of others will improve decision-making in complex multi-agent tasks, especially where model-based planning and imagination are required Hassabis et al. (2013); Hula et al. (2015); Oliehoek & Amato (2016). Our work thus ties in to a rich history of opponent modelling Brown (1951); Albrecht & Stone (2017); within this context, we show how meta-learning could be used to furnish an agent with the ability to build flexible and sample-efficient models of others on the fly. Such models will be important for value alignment Hadfield-Menell et al. (2016) and flexible cooperation Nowak (2006); Kleiman-Weiner et al. (2016); Barrett et al. (2017); Kris Cao , and will likely be an ingredient in future machines’ ethical decision making Churchland (1996). They will also be highly useful for communication and pedagogy Dragan et al. (2013); Fisac et al. (2017); Milli et al. (2017), and will thus likely play a key role in human-machine interaction. Exploring the conditions under which such abilities arise can also shed light on the origin of our human abilities Carey (2009). Finally, such models will likely be crucial mediators of our human understanding of artificial agents.

Lastly, we are strongly motivated by the goals of making artificial agents human-interpretable. We attempt a novel approach here: rather than modifying agents architecturally to expose their internal states in a human-interpretable form, we seek to build intermediating systems which learn to reduce the dimensionality of the space of behaviour and re-present it in more digestible forms. In this respect, the pursuit of a Machine ToM is about building the missing interface between machines and human expectations Cohen et al. (1981).

We consider the challenge of building a Theory of Mind as essentially a meta-learning problem Schmidhuber et al. (1996); Thrun & Pratt (1998); Hochreiter et al. (2001); Vilalta & Drissi (2002). At test time, we want to be able to encounter a novel agent whom we have never met before, and already have a strong and rich prior about how they are going to behave. Moreover, as we see this agent act in the world, we wish to be able to collect data (i.e. form a posterior) about their latent characteristics and mental states that will enable us to improve our predictions about their future behaviour.

To do this, we formulate a meta-learning task. We construct an observer, who in each episode gets access to a set of behavioural traces of a novel agent. The observer’s goal is to make predictions of the agent’s future behaviour. Over the course of training, the observer should get better at rapidly forming predictions about new agents from limited data. This “learning to learn” about new agents is what we mean by meta-learning. Through this process, the observer should also learn an effective prior over the agents’ behaviour that implicitly captures the commonalities between agents within the training population.

We introduce two concepts to describe components of this observer network and their functional role. We distinguish between a general theory of mind – the learned weights of the network, which encapsulate predictions about the common behaviour of all agents in the training set – and an agent-specific theory of mind – the “agent embedding” formed from observations about a single agent at test time, which encapsulates what makes this agent’s character and mental state distinct from others’. These correspond to a prior and posterior over agent behaviour.

This paper is structured as a sequence of experiments of increasing complexity on this Machine Theory of Mind network, which we call a ToMnet. These experiments showcase the idea of the ToMnet, exhibit its capabilities, and demonstrate its capacity to learn rich models of other agents incorporating canonical features of humans’ Theory of Mind, such as the recognition of false beliefs.

Some of the experiments in this paper are directly inspired by the seminal work of Baker and colleagues in Bayesian Theory of Mind, such as the classic food-truck experiments Baker et al. (2011; 2017). We have not sought to directly replicate these experiments as the goals of this work differ. In particular, we do not immediately seek to explain human judgements in computational terms, but instead we emphasise machine learning, scalability, and autonomy. We leave the alignment to human judgements as future work. Our experiments should nevertheless generalise many of the constructions of these previous experiments.

In Section 3.1, we show that for simple, random agents, the ToMnet learns to approximate Bayes-optimal hierarchical inference over agents’ characteristics.

In Section 3.2, we show that the ToMnet learns to infer the goals of algorithmic agents (effectively performing few-shot inverse reinforcement learning), as well as how they balance costs and rewards.

In Section 3.3, we show that the ToMnet learns to characterise different species of deep reinforcement learning agents, capturing the essential factors of variations across the population, and forming abstract embeddings of these agents. We also show that the ToMnet can discover new abstractions about the space of behaviour.

In Section 3.4, we show that when the ToMnet is trained on deep RL agents acting in POMDPs, it implicitly learns that these agents can hold false beliefs about the world, a core component of humans’ Theory of Mind.

In Section 3.5, we show that the ToMnet can be trained to predict agents’ belief states as well, revealing agents’ false beliefs explicitly. We also show that the ToMnet can infer what different agents are able to see, and what they therefore will tend to believe, from their behaviour alone.

Model

Here we describe the formalisation of the task. We assume we have a family of partially observable Markov decision processes (POMDPs) M=⋃jMj\mathcal{M}=\bigcup_{j}\mathcal{M}_{j}. Unlike the standard formalism, we associate the reward functions, discount factors, and conditional observation functions with the agents rather than with the POMDPs. For example, a POMDP could be a gridworld with a particular arrangement of walls and objects; different agents, when placed in the same POMDP, might receive different rewards for reaching these objects, and be able to see different amounts of their local surroundings. The POMDPs are thus tuples of state spaces SjS_{j}, action spaces AjA_{j}, and transition probabilities TjT_{j} only, i.e. Mj=(Sj,Aj,Tj)\mathcal{M}_{j}=(S_{j},A_{j},T_{j}). In this work, we only consider single-agent POMDPs, though the extension to the multi-agent case is simple. When agents have full observability, we use the terms MDP and POMDP interchangeably. We write the joint state space over all POMDPs as S=⋃jSjS=\bigcup_{j}S_{j}.

Separately, we assume we have a family of agents A=⋃iAi\mathcal{A}=\bigcup_{i}\mathcal{A}_{i}, with corresponding observation spaces Ωi\Omega_{i}, conditional observation functions ωi(⋅):S→Ωi\omega_{i}(\cdot):S\rightarrow\Omega_{i}, reward functions RiR_{i}, discount factors γi\gamma_{i}, and resulting policies πi\pi_{i}, i.e. Ai=(Ωi,ωi,Ri,γi,πi)\mathcal{A}_{i}=(\Omega_{i},\omega_{i},R_{i},\gamma_{i},\pi_{i}). These policies might be stochastic (as in Section 3.1), algorithmic (as in Section 3.2), or learned (as in Sections 3.3–3.5). We do not assume that the agents’ policies πi\pi_{i} are optimal for their respective tasks. The agents may be stateful – i.e. with policies parameterised as πi(⋅∣ωi(st),ht)\pi_{i}(\cdot|\omega_{i}(s_{t}),h_{t}) where hth_{t} is the agent’s (Markov) hidden state – though we assume agents’ hidden states do not carry over between episodes.

In turn, we consider an observer who makes potentially partial and/or noisy observations of agents’ trajectories, via a state-observation function ω(obs)(⋅):S→Ω(obs)\omega^{(obs)}(\cdot):S\rightarrow\Omega^{(obs)}, and an action-observation function α(obs)(⋅):A→A(obs)\alpha^{(obs)}(\cdot):A\rightarrow A^{(obs)}. Thus, if agent Ai\mathcal{A}_{i} follows its policy πi\pi_{i} on POMDP Mj\mathcal{M}_{j} and produces trajectory τij={(st,at)}t=0T\tau_{ij}=\{(s_{t},a_{t})\}_{t=0}^{T}, the observer would see τij(obs)={(xt(obs),at(obs))}t=0T\tau_{ij}^{(obs)}=\{(x_{t}^{(obs)},a_{t}^{(obs)})\}_{t=0}^{T}, where xt(obs)=ω(obs)(st)x_{t}^{(obs)}=\omega^{(obs)}(s_{t}) and at(obs)=α(obs)(at)a_{t}^{(obs)}=\alpha^{(obs)}(a_{t}). For all experiments we pursue here, we set ω(obs)(⋅)\omega^{(obs)}(\cdot) and α(obs)(⋅)\alpha^{(obs)}(\cdot) as identity functions, so that the observer has unrestricted access to the MDP state and overt actions taken by the agents; the observer does not, however, have access to the agents’ parameters, reward functions, policies, or identifiers.

The observer must learn to predict the behaviour of many agents, whose rewards, parameterisations, and policies may vary considerably; in this respect, the problem resembles the one-shot imitation learning setup recently introduced in Duan et al. (2017) and Wang et al. (2017). However, the problem statement differs from imitation learning in several crucial ways. First, the observer need not be able to execute the behaviours itself: the behavioural predictions may take the form of atomic actions, options, trajectory statistics, or goals or subgoals. The objective here is not to imitate, but instead to form predictions and abstractions that will be useful for a range of other tasks. Second, there is an informational asymmetry, where the “teacher” (i.e. the agent Ai\mathcal{A}_{i}) may conceivably know less about the environment state sts_{t} than the “student” (i.e. the observer), and it may carry systematic biases; its policy, πi\pi_{i}, may therefore be far from optimal. As a result, the observer may need to factor in the likely knowledge state of the agent and its cognitive limitations when making behavioural predictions. Finally, as a ToM needs to operate online while observing a new agent, we place a high premium on the speed of inference. Rather than using the computationally costly algorithms of classical inverse reinforcement learning (e.g. Ng et al., 2000; Ramachandran & Amir, 2007; Ziebart et al., 2008; Boularias et al., 2011), or Bayesian ToM (e.g. Baker et al., 2011; Nakahashi et al., 2016; Baker et al., 2017), we drive the ToMnet to amortise its inference through neural networks (as in Kingma & Welling, 2013; Rezende et al., 2014; Ho & Ermon, 2016; Duan et al., 2017; Wang et al., 2017).

2 The architecture

To solve these tasks, we designed the ToMnet architecture shown in Fig 1. The ToMnet is composed of three modules: a character net, a mental state net, and a prediction net.

3 Agents and environments

We deploy the ToMnet to model agents belonging to a number of different “species” of agent. In Section 3.1, we consider species of agents with random policies. In Section 3.2, we consider species of agents with full observability over MDPs, which plan using value iteration. In Sections 3.3 – 3.5, we consider species of agents with different kinds of partial observability (i.e. different functions ωi(⋅)\omega_{i}(\cdot)), with policies parameterised by feed-forward nets or LSTMs. We trained these agents using a version of the UNREAL deep RL framework Jaderberg et al. (2017), modified to include an auxiliary belief task of estimating the locations of objects within the MDP. Crucially, we did not change the core architecture or algorithm of the ToMnet observer to match the structure of the species, only the ToMnet’s capacity.

The POMDPs we consider here are all gridworlds with a common action space (up/down/left/right/stay), deterministic dynamics, and a set of consumable objects, as described in the respective sections and in Appendix C. We experimented with these POMDPs due to their simplicity and ease of control; our constructions should generalise to richer domains too. We parameterically generate individual Mj\mathcal{M}_{j} by randomly sampling wall, object, and initial agent locations.

Experiments

To demonstrate its essential workings, we tested the ToMnet observer on a simple but illustrative toy problem. We created a number of different species of random agents, sampled agents from them, and generated behavioural traces on a distribution of random 11×1111\times 11 gridworlds (e.g. Fig 2a). Each agent had a stochastic policy defined by a fixed vector of action probabilities πi(⋅)=πi\pi_{i}(\cdot)=\bm{\pi_{i}}. We defined different species based on how sparse its agents’ policies were: within a species S(α)\mathcal{S}(\alpha), each πi\bm{\pi_{i}} was drawn from a Dirichlet distribution with concentration parameter α\alpha. At one extreme, we created a species of agents with near-deterministic policies by drawing πi∼Dir(α=0.01)\bm{\pi_{i}}\sim Dir(\alpha=0.01); here a single agent might overwhelmingly prefer to always move left, and another to always move up. At the other extreme, we created a species of agent with far more stochastic policies, by drawing πi∼Dir(α=3)\bm{\pi_{i}}\sim Dir(\alpha=3).

When the ToMnet observer is trained on a species S(α)\mathcal{S}(\alpha), it learns how to approximate Bayes-optimal, online inference about agents’ policies πi(⋅)=πi∼Dir(α)\pi_{i}(\cdot)=\bm{\pi_{i}}\sim Dir(\alpha). Fig 3a shows how the ToMnet’s estimates of action probability increase with the number of past observations of that action, and how training the ToMnet on species with lower α\alpha makes it apply priors that the policies are indeed sparser. We can also see how the ToMnet specialises to a given species by testing it on agents from different species (Fig 3c): the ToMnet makes better predictions about novel agents drawn from the species which it was trained on. Moreover, the ToMnet easily learns how to predict behaviour from mixtures of species (Fig 3d): when trained jointly on species with highly deterministic (α=0.01\alpha=0.01) and stochastic (α=3\alpha=3) policies, it implicitly learns to expect this bimodality in the policy distribution, and specialises its inference accordingly. We note that it is not learning about two agents, but rather two species of agents, which each span a spectrum of individual parameters.

There should be nothing surprising about seeing the ToMnet learn to approximate Bayes-optimal online inference; this should be expected given more general results about inference and meta-learning with neural networks MacKay (1995); Finn & Levine (2017). Our point here is that a very first step in reasoning about other agents is an inference problem. The ToMnet is just an engine for learning to do inference and prediction on other agents.

In summary, without any changes to its architecture, a ToMnet learns a general theory of mind that is specialised for the distribution of agents it encounters in the world, and estimates an agent-specific theory of mind online for each individual agent that captures the sufficient statistics of its behaviour.

2 Inferring goal-directed behaviour

An elementary component of humans’ theory of other agents is an assumption that agents’ behaviour is goal-directed. There is a wealth of evidence showing that this is a core component of our model from early infancy Gergely et al. (1995); Woodward (1998; 1999); Buresh & Woodward (2007), and intelligent animals such as apes and corvids have been shown to have similar expectations about their conspecifics Call & Tomasello (2008); Ostojić et al. (2013). Inferring the desires of others also takes a central role in machine learning in imitation learning, most notably in inverse RL Ng et al. (2000); Abbeel & Ng (2004).

We demonstrate here how the ToMnet observer learns how to infer the goals of reward-seeking agents. We defined species of agents who acted within gridworlds with full observability (Fig 2a). Each gridworld was 11×1111\times 11 in size, had randomly-sampled walls, and contained four different objects placed in random locations. Consuming an object yielded a reward for the agent and caused the episode to terminate. Each agent, Ai\mathcal{A}_{i}, had a unique, fixed reward function, such that it received reward ri,a∈(0,1)r_{i,a}\in(0,1) when it consumed object aa; the vectors ri\mathbf{r_{i}} were sampled from a Dirichlet distribution with concentration parameter α=0.01\alpha=0.01. Agents also received a negative reward of −0.01-0.01 for every move taken, and a penalty of 0.050.05 for walking into walls. In turn, the agents planned their behaviour through value iteration, and hence had optimal policies πi∗\pi_{i}^{*} with respect to their own reward functions.

We trained the ToMnet to observe behaviour of these agents in randomly-sampled “past” MDPs, and to use this to predict the agents’ behaviour in a “current” MDP. We detail three experiments below; these explore the range of capabilities of the ToMnet in this domain.

First, we provided the ToMnet with a full trajectory of an agent on a single past MDP (Fig 4a). In turn, we queried the ToMnet with the initial state of a current MDP (Fig 4b) and asked for a set of predictions: the next action the agent would take (Fig 4c top), what object the agent would consume by the end of the episode (Fig 4c bottom), and a set of statistics about the agent’s trajectory in the current MDP, the successor representation (SR; the expected discounted state occupancy; Dayan, 1993, Fig 4). The ToMnet’s predictions qualitatively matched the agents’ true behaviours.

We note that unlike the approach of inverse RL, the ToMnet is not constrained to explicitly infer the agents’ reward functions in service of its predictions. Nevertheless, in this simple task, using a 2-dimensional character embedding space renders this information immediately legible (Fig 5d). This is also true when the only behavioural prediction is next-step action.

3 Learning to model deep RL agents

The previous experiments demonstrate the ToMnet’s ability to learn models of simple, algorithmic agents which have full observability. We next considered the ToMnet’s ability to learn models for a richer population of agents: those with partial observability and neural network-based policies, trained using deep reinforcement learning. In this section we show how the ToMnet learns how to do inference over the kind of deep RL agent it is observing, and show the specialised predictions it makes as a consequence.

This domain begins to capture the complexity of reasoning about real-world agents. So long as the deep RL agents share some overlap in their tasks, structure, and learning algorithms, we expect that they should exhibit at least some shared behavioural patterns. These patterns should also diverge systematically from each other as the aforementioned factors vary, and individual agents may also exhibit idiosyncratic behaviour as they settle in local minima while optimising their respective policies. There are thus opportunities to learn rich general and agent-specific theories of mind for such populations. Moreover, as the tasks and networks become more complex, hand-crafting a Machine Theory of Mind to parse behaviour based on our human knowledge (e.g. Baker et al., 2011; Nakahashi et al., 2016; Baker et al., 2017; Lake et al., 2017) becomes increasingly intractable; instead we seek here a path towards machines which learn how to model others’ minds autonomously Botvinick et al. (2017).

We trained three different species of agents on a modified version of the gridworlds, described below in Section 3.4. In brief, agents received maximum reward for reaching a subgoal location first, then consuming a preferred object that differed from agent to agent. Consuming any of the non-subgoal objects terminated the episode. All agents were based on the UNREAL architecture Jaderberg et al. (2017), with details given in Appendix D. One species of agent (“blind”) was unable to observe the maze state at all, and could only observe its previous action (at−1a_{t-1}) and reward (rt−1r_{t-1}), which it could integrate over time through its LSTM state. The second species had partial observability (“sighted”), but was stateless: these agents could observe the gridworld within a 5×55\times 5 window centred at their current location, with the rest of the maze shrouded by a binary mask; the agents’ policies however were purely reactive, implemented via feedforward networks without any memory. The third species shared the benefits of the other two, being both sighted (with partial observability) and stateful (with an LSTM-based policy).

The ToMnet, on the other hand, observed the behaviour of these agents with full observability of the POMDP state. We constructed a training set for the ToMnet by generating behaviour from 120 trained agents (3 species ×\times 4 preferred objects ×\times 10 neural nets trained with different random seeds). We held out a test set of a further 120 trained agents (i.e. trained using 10 additional random seeds) for evaluating the trained ToMnet. All results below are shown for the test set, i.e. predicting behaviour for novel agents which the ToMnet had never seen before.

Unlike previous experiments, these agents’ behaviour depended on both their individual characteristics and their state; the ToMnet thus needed both a character net and a mental net to make the best predictions (Fig 7a).

Qualitative evaluations of the ToMnet’s predictions show how it learned the expected behaviour of the three species of agents. Fig 7b shows the ToMnet’s predictions of future state occupancy for the same query state, but given different past observations of how the agent behaves. Without being given the species label, the ToMnet implicitly infers it, and maps out where the agent will go next: blind agents continue until they hit a wall, then turn; sighted but stateless agents consume objects opportunistically; sighted, stateful agents explore the interior and seek out the subgoal. Thus the ToMnet develops general models for the three different species of agents in its world.

4 Acting based on false beliefs

It has long been argued that a core part of human Theory of Mind is that we recognise that other agents do not base their decisions directly on the state of the world, but rather on an internal representation of the state of the world Leslie (1987); Gopnik & Astington (1988); Wellman (1992); Baillargeon et al. (2016). This is usually framed as an understanding that other agents hold beliefs about the world: they may have knowledge that we do not; they may be ignorant of something that we know; and, most dramatically, they may believe the world to be one way, when we in fact know this to be mistaken. An understanding of this last possibility – that others can have false beliefs – has become the most celebrated indicator of a rich Theory of Mind, and there has been considerable research into how much children, infants, apes, and other species carry this capability Baron-Cohen et al. (1985); Southgate et al. (2007); Clayton et al. (2007); Call & Tomasello (2008); Krupenye et al. (2016); Baillargeon et al. (2016).

Here, we sought to explore whether the ToMnet would also learn that agents may hold false beliefs about the world. To do so, we first needed to generate a set of POMDPs in which agents could indeed hold incorrect information about the world (and act upon this). To create these conditions, we allowed the state of the environment to undergo random changes, sometimes where the agents couldn’t see them. In the subgoal maze described above in Section 3.3, we included a low probability (p=0.1p=0.1) state transition when the agent stepped on the subgoal, such that the four other objects would randomly permute their locations instantaneously (Fig 10a-b). These swap events were only visible to the agent insofar as the objects’ positions were within the agent’s current field of view; when the swaps occurred entirely outside its field of view, the agent’s internal state and policy at the next time step remained unaffected (policy changes shown in Fig 10c, right side), a signature of a false belief. As agents were trained to expect these low-probability swap events, they learned to produce corrective behaviour as their policy was rolled out over time (Fig 10d, right side). While the trained agents were competent at the task, they were not optimal.

Our goal was to determine whether the ToMnet would learn a general theory of mind that included an element of false beliefs. However, the ToMnet, as described, does not have the capacity to explicitly report agents’ (latent) belief states, only the ability to report predictions about the agents’ overt behaviour. To proceed, we took inspiration from the literature on human infant and ape Theory of Mind Call & Tomasello (2008); Baillargeon et al. (2016). Here, experimenters have often utilised variants of the classic “Sally-Anne test” Wimmer & Perner (1983); Baron-Cohen et al. (1985) to probe subjects’ models of others. In the classic test, the observer watches an agent leave a desired object in one location, only for it to be moved, unseen by the agent. The subject, who sees all, is asked where the agent now believes the object lies. While infants and apes have limited ability to explicitly report such inferences about others’ mental states, experimenters have nevertheless been able to measure these subjects’ predictions of where the agents will actually go, e.g. by measuring anticipatory eye movements, or surprise when agents behave in violation of subjects’ expectations Call & Tomasello (2008); Krupenye et al. (2016); Baillargeon et al. (2016). These experiments have demonstrated that human infants and apes can implicitly model others as holding false beliefs.

We used the swap events to construct a gridworld Sally-Anne test. We hand-crafted scenarios where an agent would see its preferred blue object in one location, but would have to move away from it to reach a subgoal before returning to consume it (Fig 11a). During this time, the preferred object might be moved by a swap event, and the agent may or may not see this occur, depending on how far away the subgoal was. We forced the agents along this trajectory (off-policy), and measured how a swap event affected the agent’s probability of moving back to the preferred object. As expected, when the swap occurred within the agent’s field of view, the agent’s likelihood of turning back dropped dramatically; when the swap occurred outside its field of view, the policy was unchanged (Fig 11b, left).

In turn, we presented these demonstration trajectories to the ToMnet (which had seen past behaviour indicating the agent’s preferred object). Crucially, the ToMnet was able to observe the entire POMDP state, and thus was aware of swaps when the agent was not. To perform this task properly, the ToMnet needs to have implicitly learned to separate out what it itself knows, and what the agent can plausibly know, without relying on a hand-engineered, explicit observation model for the agent. Indeed, the ToMnet predicted the correct behavioural patterns (Fig 11b, right): specifically, the ToMnet predicts that when the world changes far away from an agent, that agent will persist with a policy that is founded on false beliefs about the world.

This test was a hand-crafted scenario. We validated its results by looking at the ToMnet’s predictions for how the agents responded to all swap events in the distribution of POMDPs. We sampled a set of test mazes, and rolled out the agents’ policies until they consumed the subgoal, selecting only episodes where the agents had seen their preferred object along the way. At this point, we created a set of counterfactuals: either a swap event occurred, or it didn’t.

We measured the ground truth for how the swaps would affect the agent’s policy, via the average Jensen-Shannon divergence (DJSD_{JS}) between the agent’s true action probabilities in the no-swap and swap conditionsFor a discussion of why we used the DJSD_{JS} measure, see Appendix F.2.. As before, the agent’s policy often changed when a swap was in view (for these agents, within a 2 block radius), but wouldn’t change when the swap was not observable (Fig 12a, left).

The ToMnet learned that the agents’ policies were indeed more sensitive to local changes in the POMDP state, but were relatively invariant to changes that occurred out of sight (Fig 12a, right). The ToMnet did not, however, learn a hard observability boundary, and was more liberal in predicting that far-off changes could affect agent policy. The ToMnet also correctly predicted that the swaps would induce corrective behaviour over longer time periods, even when they were not initially visible (Fig 12b).

These patterns were even more pronounced when we trained the ToMnet on mixed populations of agents with different fields of view. In this task, the ToMnet had to infer what each agent could see (from past behaviour alone) in order to predict each agent’s behaviour in the future. The ToMnet’s predictions reveal an implicit grasp of how different agents’ sensory abilities render them differentially vulnerable to acquire false beliefs (Fig 12c-d).

Most surprising of all, we found that the ToMnet learned these statistics even if the ToMnet had never seen swap events during training: the curves in Fig 12 were qualitatively identical for the ToMnet under such conditions (Fig A1).

On the one hand, we were impressed that the ToMnet learns a general theory of mind that incorporates an implicit understanding that agents act based on their own persistent representations of the world, even if they are mistaken. On the other hand, we should not attribute this cognitive ability to a special feature of the ToMnet architecture itself, which is indeed very straightforward. Rather, this work demonstrates that representational Theory of Mind can arise simply by observing competent agents acting in POMDPs.

5 Explicitly inferring belief states

We have demonstrated that the ToMnet learns that agents can act based on false beliefs. This is limited, though, in that the ToMnet cannot explicitly report what these agents know and don’t know about the world. Moreover, it is difficult to extract any beliefs that will not manifest immediately in overt behaviour.

We therefore extended the ToMnet to be able to make declarative statements about agents’ beliefs. We achieved this by constructing a supervised dataset of belief states in the sample gridworld. We trained the UNREAL agents to report their beliefs about the locations of the four objects and the subgoal at every time step, alongside their policy. To do this, we added a head to the LSTM that output a posterior over each object’s current location on the grid (or whether it was absent). During training, the agents learned to report their best estimate of each object’s current location, based on its observations so far during the episode. Example belief states for the query MDP states in Fig 10a-b are shown in Fig 13a. Note that these reported beliefs are not causal to the agents’ policy; they are just a readout from the same LSTM hidden state that serves the policy.

In turn, we trained a single ToMnet on agents with a range of different fields of view. This ToMnet had to observe only agents’ overt behaviour, and use it to predict, for a given query state, an agent’s policy, consumptions, successor representations, and reported beliefs (i.e., what the agent would say about the world if it were asked; e.g. Fig 13b).

As shown in Fig 13c, the ToMnet learns agent-specific theories of mind for the different subspecies that grasp the essential differences between their belief-forming tendencies: agents with less visibility of changes in their world are more likely to report false beliefs; and behave according to them too (as in Fig 13c).

Last of all, we included an additional variational information bottleneck penalty, to encourage low-dimensional abstract embeddings of agent types. As with the agent characterisation in Fig 7, the character embeddings of these agents separated along the factors of variation (field of view and preferred object; Fig 14). Moreover, these embeddings show the ToMnet’s ability to distinguish different agents’ visibility: blind and 3×33\times 3 agents are easily distinguishable, whereas there is little in past behaviour to separate 7×77\times 7 agents from 9×99\times 9 agents (or little benefit in making this distinction).

We note that this particular construction of explicit belief inference will likely not scale in its current form. Our method depends on two assumptions that break down in the real world. First, it requires access to others’ latent belief states for supervision. We assume here that the ToMnet gets access to these via a rich communication channel; as humans, this channel is likely much sparser. It is an empirical question as to whether the real-world information stream is sufficient to train such an inference network. We do, however, have privileged access to some of our own mental states through meta-cognition; though this data may be biased and noisy, it might be sufficiently rich to learn this task. Second, it is intractable to predict others’ belief states about every aspect of the world. As humans, we nevertheless have the capacity to make such predictions about arbitrary variables as the need arises. This may require creative solutions in future work, such as forming abstract embeddings of others’ belief states that can be queried.

Discussion

In this paper, we used meta-learning to build a system that learns how to model other agents. We have shown, through a sequence of experiments, how this ToMnet learns a general model for agents in the training distribution, as well as how to construct an agent-specific model online while observing a new agent’s behaviour. The ToMnet can flexibly learn such models over a range of different species of agents, whilst making few assumptions about the generative processes driving these agents’ decision making. The ToMnet can also discover abstractions within the space of behaviours.

We note that the experiments we pursued here were simple, and designed to illustrate the core ideas and capabilities of such a system. There is much work to do to scale the ToMnet to richer domains.

First, we have worked entirely within gridworlds, due to the control such environments afford. We look forward to extending these systems to operate within complex 3D visual environments, and within other POMDPs with rich state spaces.

Second, we did not experiment here with limiting the observability of the observer itself. This is clearly an important challenge within real-world social interaction, e.g. when we try to determine what someone else knows that we do not. This is, at its heart, an inference problem Baker et al. (2017); learning to do this robustly is a future challenge for the ToMnet.

Third, there are many other dimensions over which we may wish to characterise agents, such as whether they are animate or inanimate Scholl & Tremoulet (2000), prosocial or adversarial Ullman et al. (2009), reactive or able to plan Sutton & Barto (1998). Potentially more interesting is the possibility of using the ToMnet to discover new structure in the behaviour of either natural or artificial populations, i.e. as a kind of machine anthropology.

Fourth, a Theory of Mind is important for social beings as it informs our social decision-making. An important step forward for this research is to situate the ToMnet inside artificial agents, who must learn to perform multi-agent tasks.

In pursuing all these directions we anticipate many future needs: to enrich the set of predictions a ToMnet must make; to introduce gentle inductive biases to the ToMnet’s generative models of agents’ behaviour; and to consider how agents might draw from their own experience and cognition in order to inform their models of others. Addressing these will be necessary for advancing a Machine Theory of Mind that learns the rich capabilities of responsible social beings.

Acknowledgements

We’d like to thank the many people who provided feedback on the research and the manuscript, including Marc Lanctot, Jessica Hamrick, Ari Morcos, Agnieszka Grabska-Barwinska, Avraham Ruderman, Christopher Summerfield, Pedro Ortega, Josh Merel, Doug Fritz, Nando de Freitas, Heather Roff, Kevin McKee, and Tina Zhu.

References

Appendix A Model description: architectures

Here we describe the precise details of the architectures used in the main text.

We note that we did not optimise our results by tweaking architectures or hyperparameters in any systematic or substantial way. Rather, we simply picked sensible-looking values. We anticipate that better performance could be obtained by improving these decisions, but this is beyond the scope of this work.

Pre-processing. Both the character net and the mental state net consume trajectories, which are sequences of observed state/action pairs, τij(obs)={(xt(obs),at(obs))}t=0T\tau_{ij}^{(obs)}=\{(x_{t}^{(obs)},a_{t}^{(obs)})\}_{t=0}^{T}, where ii is the agent index, and jj is the episode index. The observed states in our experiments, xt(obs)x_{t}^{(obs)}, are always tensors of shape (11×11×K)(11\times 11\times K), where KK is the number of feature planes (comprising one feature plane for the walls, one for each object, and one for the agent). The observed actions, at(obs)a_{t}^{(obs)}, are always vectors of length 5. We combine these data through a spatialisation-concatenation operation, whereby the actions are tiled over space into a (11×11×5)(11\times 11\times 5) tensor, and concatenated with the states to form a single tensor of shape (11×11×(K+5))(11\times 11\times(K+5)).

Training. All ToMnets were trained with the Adam optimiser, with learning rate 10−410^{-4}, using batches of size 16. We trained the ToMnet for 40k minibatches for random agents (Section 3.1), and for 2M minibatches otherwise.

A.2 ToMnet for random agents (Section 3.1)

A.3 ToMnet for inferring goals (Section 3.2)

Data. Character embedding formed from a single past episode, comprising a full trajectory on a single MDP. Query state is the initial state of a new MDP, so no mental state embedding required.

Action prediction head. From the torso output: a 1-layer convnet with 32 channels and ReLUs, followed by average pooling, and a fully-connected layer to 5-dim logits, followed by a softmax. This gives the predicted policy, π^\hat{\pi}.

Consumption prediction head. From the torso output: a 1-layer convnet with 32 channels and ReLUs, followed by average pooling, and a fully-connected layer to 4-dims, followed by a sigmoid. This gives the respective Bernoulli probabilities that each of the four objects will be consumed by the end of the episode, c^\hat{c}.

Successor representation prediction head. From the torso output: a 1-layer convnet with 32 channels and ReLUs, then a 1-layer convnet with 3 channels, followed by a softmax over each channel independently. This gives the predicted normalised SRs for the three discount factors, γ=0.5,0.9,0.99\gamma=0.5,0.9,0.99.

A.3.2 Experiment 2: many past MDPs, only a single snapshot each

A.3.3 Experiment 3: greedy agents

A.4 ToMnet for modelling deep RL agents (Section 3.3)

A.5 ToMnet for false beliefs (Sections 3.4–3.5)

The ToMnet architecture was the same as described above in Appendix A.4. The experiments in Section 3.5 also included an additional belief prediction head to the prediction net.

Belief prediction head. For each object, this head outputs a 122-dim discrete distribution (the predicted belief that the object is in each of the 11×1111\times 11 locations on the map, or whether the agent believes the object is absent altogether). From the torso output: a 1-layer convnet with 32 channels and ReLU, branching to (a) another 1-layer convnet with 5 channels for the logits for the predicted beliefs that each object is at the 11×1111\times 11 locations on the map, as well as to (b) a fully-connected layer to 5-dims for the predicted beliefs that each object is absent. We unspatialise and concatenate the outputs of (a) and (b) in each of the 5 channels, and apply a softmax to each channel.

Appendix B Loss function

Here we describe the components of the loss function used for training the ToMnet.

For each agent, Ai\mathcal{A}_{i}, we sample past and current trajectories, and form predictions for the query POMDP at time tt. Each prediction provides a contribution to the loss, described below. We average the respective losses across each of the agents in the minibatch, and give equal weighting to each loss component.

Action prediction. The negative log-likelihood of the true action taken by the agent under the predicted policy:

Consumption prediction. For each object, kk, the negative log-likelihood that the object is/isn’t consumed:

Successor representation prediction. For each discount factor, γ\gamma, we define the agent’s empirical successor representation as the normalised, discounted rollout from time tt onwards, i.e.:

where ZZ is the normalisation constant such that ∑sSRγ(s)=1\sum_{s}SR_{\gamma}(s)=1. The loss here is then the cross-entropy between the predicted successor representation and the empirical one:

Belief prediction. The agent’s belief states for each object kk is a discrete distribution over 122 dims (the 11×1111\times 11 locations on the map, plus an additional dimension for an absent object), denoted bk(s)b_{k}(s). For each object, kk, the loss is the cross-entropy between the ToMnet’s predicted belief state and the agent’s true belief state:

Deep Varational Information Bottleneck. In addition to these loss components, where DVIB was used, we included an additional term for the β\beta-weighted KLs between posteriors and the priors

Appendix C Gridworld details

The POMDPs Mj\mathcal{M}_{j} were all 11×1111\times 11 gridworld mazes. Mazes in Sections 3.1–3.2 were sampled with between 0 and 4 random walls; mazes in Sections 3.3–3.5 were sampled with between 0 and 6 random walls. Walls were defined between two randomly-sampled endpoints, and could be diagonal.

Each Mj\mathcal{M}_{j} contained four terminal objects. These objects could be consumed by the agent walking on top of them. Consuming these objects ended an episode. If no terminal object was consumed after 31 steps (random and algorithmic agents; Sections 3.1–3.2) or 51 steps (deep RL agents; Sections 3.3–3.5), the episodes terminated automatically as a time-out. The sampled walls may trap the agent, and make it impossible for the agent to terminate the episode without timing out.

Deep RL agents (Sections 3.3–3.5) acted in gridworlds that contained an additional subgoal object. Consuming the subgoal did not terminate the episode.

Reward functions for the agents were as follows:

Random agents (Section 3.1.) No reward function.

Algorithmic agents (Section 3.2). For a given agent, the reward function over the four terminal objects was drawn randomly from a Dirichlet with concentration parameter 0.01. Each agent thus has a sparse preference for one object. Penalty for each move: 0.01. Penalty for walking into a wall: 0.05. Greedy agents’ penalty for each move: 0.5. These agents planned their trajectories using value iteration, with a discount factor of 1. When multiple moves of equal value were available, these agents sampled from their best moves stochastically.

Deep RL agents (Sections 3.3–3.5). Penalty for each move: 0.005. Penalty for walking into a wall: 0.05. Penalty for ending an episode without consuming a terminal object: 1.

For each deep RL agent species (e.g. blind, stateless, 5×55\times 5, …), we trained a number of canonical agents which received a reward of 1 for consuming the subgoal, and a reward of 1 for consuming a single preferred terminal object (e.g. the blue one). Consuming any other object yielded zero reward (though did terminate the episode). We artifically enlarged this population of trained agents by a factor of four, by inserting permutations into their observation functions, ωi\omega_{i}, that effectively permuted the object channels. For example, when we took a trained blue-object-preferring agent, and inserted a transformation that swapped the third object channel with the first object channel, this agent behaved as a pink-object-preferring agent.

Appendix D Deep RL agent training and architecture

Deep RL agents were based on the UNREAL architecture Jaderberg et al. (2017). These were trained with over 100M episode steps, using 16 CPU workers. We used the Adam optimiser with a learning rate of 10−510^{-5}, and BPTT, unrolling over the whole episode (50 steps). Policies were regularised with an entropy cost of 0.005 to encourage exploration.

We trained a total of 660 agents, spanning 33 random seeds ×\times 5 fields of view ×\times 2 architectures (feedforward/convolutional LSTM) ×\times 2 depths (4 layer convnet or 2 layer convnet, both with 64 channels). We selected the top 20 agents per condition (out of 33 random seeds), by their average return. We randomly partitioned these sets into 10 training and 10 test agents per condition. With the reward permutations described above in Appendix C, this produced 40 training and 40 test agents per condition.

Observations. Agents received an observation at each time step of nine 11×1111\times 11 feature planes – indicating, at each location, whether a square was empty, a wall, one of the five total objects, the agent, or currently unobservable.

Beliefs. We also trained agents with the auxiliary task of predicting the current locations of all objects in the map. To do this, we included an additional head to the Convolutional LSTMs, in addition to the policy (πt\pi_{t}) and baseline (VtV_{t}) heads. This head output a posterior for each object’s location in the world, bkb_{k} (i.e. a set of five 122-dim discrete distributions, over the 11×1111\times 11 maze size, including an additional dimension for a prediction that that the object is absent). For the belief head, we used a 3-layer convnet with 32 channels and ReLU nonlinearities, followed by a softmax. This added a term to the training loss: the cross entropy between the current belief state and the true current world state. The loss for the belief prediction was scaled by an additional hyperparameter, swept over the values 0.5, 2, and 5.

Appendix E Additional results

Appendix F Additional notes

In Fig 12c, the policies of agents with 3×33\times 3 fields of view are seen to be considerably more sensitive to swap events that occur adjacent to the agent than the agents with 9×99\times 9 fields of view. Agents with 5×55\times 5 and 7×77\times 7 had intermediate sensitivities.

We did not perform a systematic analysis of the policy differences between these agents, but we speculate here as to the origin of this phenomenon. As we note in the main text, the agents were competent at their respective tasks, but not optimal. In particular, we noted that agents with larger fields of view were often sluggish to respond behaviourally to swap events. This is evident in the example shown on the left hand side of Fig 10. Here an agent with a 5×55\times 5 field of view does not respond to the sudden appearance of its preferred blue object above it by immediately moving upwards to consume it; its next-step policy does shift some probability mass to moving upwards, but only a small amount (Fig 10c). It strongly adjusts its policy on the following step though, producing rollouts that almost always return directly to the object (Fig 10d). We note that when a swap event occurs immediately next to an agent with a relatively large field of view (5×55\times 5 and greater), such an agent has the luxury of integrating information about the swap events over multiple timesteps, even if it navigates away from this location. In contrast, agents with 3×33\times 3 fields of view might take a single action that results in the swapped object disappearing altogether from their view. There thus might be greater pressure on these agents during learning to adjust their next-step actions in response to neighbouring swap events.

F.2 Use of Jensen-Shannon Divergence

In Sections 3.4–3.5, we used the Jensen-Shannon Divergence (DJSD_{JS}) to measure the effect of swap events on agents’ (and the ToMnet’s predicted) behaviour (Figs 12-13). We wanted to use a standard metric for changes to all the predictions (policy, successors, and beliefs), and we found that the symmetry and stability of DJSD_{JS} was most suited to this. We generally got similar results when using the KL-divergence, but we typically found more variance in these estimates: DKLD_{KL} is highly sensitive the one of the distributions assigning little probability mass to one of the outcomes. This was particularly problematic when measuring changes in the successor representations and belief states, which were often very sparse. While it’s possible to tame the the KL by adding a little uniform probability mass, this involves an arbitrary hyperparameter which we preferred to just avoid.

Appendix G Version history

v2: 12 Mar 2018

Added missing references to opponent modelling in introduction.

Typographical error in citation Dayan (1993).