RNNs can generate bounded hierarchical languages with optimal memory

John Hewitt, Michael Hahn, Surya Ganguli, Percy Liang, Christopher D. Manning

Introduction

Recurrent neural networks (RNNs; Elman (1990)) trained on large datasets have demonstrated a grasp of natural language syntax Karpathy et al. (2015); Kuncoro et al. (2018). While considerable empirical work has studied RNN language models’ ability to capture syntactic properties of language Linzen et al. (2016); Marvin and Linzen (2018); Hewitt and Manning (2019); van Schijndel and Linzen (2018), their success is not well-understood theoretically. In this work, we provide theoretical insight into RNNs’ syntactic success, proving that they can efficiently generate a family of bounded hierarchical languages. These languages form the scaffolding of natural language syntax.

Hierarchical structure characterized by long distance, nested dependencies, lies at the foundation of natural language syntax. This motivates, for example, context-free languages Chomsky (1956), a fundamental paradigm for describing natural language syntax. A canonical family of context-free languages (CFLs) is Dyck-kk, the language of balanced brackets of kk types, since any CFL can be constructed via some Dyck-kk Chomsky and Schützenberger (1959).

However, while context-free languages like Dyck-kk describe arbitrarily deep nesting of hierarchical structure, in practice, natural languages exhibit bounded nesting. This is clear in, e.g., bounded center-embedding Karlsson (2007); Jin et al. (2018) (Figure 1). To reflect this, we introduce and study Dyck-(kk,mm), which adds a bound mm on the number of unclosed open brackets at any time. Informally, the ability to efficiently generate Dyck-(kk,mm) suggests the foundation of the ability to efficiently generate languages with the syntactic properties of natural language. See §(1.1) for further motivation for bounded hierarchical structure.

In our main contribution, we prove that RNNs are able to generate Dyck-(kk,mm) as memory-efficiently as any model, up to constant factors. Since Dyck-(kk,mm) is a regular (finite-state) language, the application of general-purpose RNN constructions trivially proves that RNNs can generate the language Merrill (2019). However, the best construction we are aware of uses O(km2)O(k^{\frac{m}{2}}) hidden units Horne and Hush (1994); Indyk (1995), where kk is the vocabulary size and mm is the nesting depth, which is exponential.For hard-threshold neural networks, the lower-bound would be Ω(km2)\Omega(k^{\frac{m}{2}}) if Dyck-(kk,mm) were an arbitrary regular language. We provide an explicit construction proving that a Simple (Elman; Elman (1990)) RNN can generate any Dyck-(kk,mm) using only 6m⌈log⁡k⌉−2m=O(mlog⁡k)6m\lceil\log k\rceil-2m=O(m\log k) hidden units, an exponential improvement. This is not just a strong result relative to RNNs’ general capacity; we prove that even computationally unbounded models generating Dyck-(kk,mm) require Ω(mlog⁡k)\Omega(m\log k) hidden units, via a simple communication complexity argument.

Our proofs provide two explicit constructions, one for the Simple RNN and one for the LSTM, which detail how these networks can use their hidden states to simulate stacks in order to efficiently generate Dyck-(kk,mm). The differences between the constructions exemplify how LSTMs can use exclusively their gates, ignoring their Simple RNN subcomponent entirely, to reduce the memory by a factor of 2 compared to the Simple RNN.We provide implementations at https://github.com/john-hewitt/dyckkm-constructions/.

We prove these results under a theoretical setting that aims to reflect the realistic settings in which RNNs have excelled in NLP. First, we assume finite precision; the value of each hidden unit is represented by p=O(1)p=O(1) bits. This has drastic implications compared to existing work Siegelmann and Sontag (1992); Weiss et al. (2018); Merrill (2019); Merrill et al. (2020). It implies that only regular languages can be generated by any machine with dd hidden units, since they can take on only 2pd2^{pd} states Korsky and Berwick (2019). This points us to focus on whether languages can be implemented memory-efficiently. Second, we consider networks as language generators, not acceptors;Acceptors consume a whole string and then decide whether the string is in the language; generators must decide which tokens are possible continuations at each timestep. informally, RNNs’ practical successes have been primarily as generators, like in language modeling and machine translation Karpathy et al. (2015); Wu et al. (2016).

Finally, we include a preliminary study in learning Dyck-(kk,mm) with LSTM LMs from finite samples, finding for a range of kk and mm that learned LSTM LMs extrapolate well given the hidden sizes predicted by our theory.

In summary, we prove that RNNs are memory optimal in generating a family of bounded hierarchical languages that forms the scaffolding of natural language syntax by describing mechanisms that allow them to do so; this provides theoretical insight into their empirical success.

Hierarchical structure is central to human language production and comprehension, showing up in grammatical constraints and semantic composition, among other properties Chomsky (1956). Agreement between subject and verb in English is an intuitive example:

... The Dyck-kk languages—well-nested brackets of kk types—are the prototypical languages of hierarchical structure; by the Chomsky-Schützenberger Theorem Chomsky and Schützenberger (1959), they form the scaffolding for any context-free language. They have a simple structure:

..... However, human languages are unlike Dyck-kk and other context-free languages in that they exhibit bounded memory requirements. Dyck-kk requires storage of an unboundedly long list of open brackets in memory. In human language, as the center-embedding depth grows, comprehension becomes more difficult, like in our example sentence above Miller and Chomsky (1963). Empirically, center-embedding depth of natural language is rarely greater than 33 Jin et al. (2018); Karlsson (2007). However, it does exhibit long-distance, shallow hierarchical structure:

.. Our Dyck-(kk,mm) language puts a bound on depth in Dyck-kk, capturing the long-distance hierarchical structure of natural language as well as its bounded memory requirements.For further motivation, we note that center-embedding directly implies bounded memory requirements in arc-eager left-corner parsers Resnik (1992).

Related Work

This work contributes primarily to the ongoing theoretical characterization of the expressivity of RNNs. Siegelmann and Sontag (1992) proved that RNNs are Turing-complete if provided with infinite precision and unbounded computation time. Recent work in NLP has taken an interest in the expressivity of RNNs under conditions more similar to RNNs’ practical uses, in particular assuming one “unrolling” of the RNN per input token, and precision bounded to be logarithmic in the sequence length. In this setting, Weiss et al. (2018) proved that LSTMs can implement simplified counter automata; the implications of which were explored by Merrill (2019, 2020). In this same regime, Merrill et al. (2020) showed a strict hierarchy of RNN expressivity, proving among other results that RNNs augmented with an external stack Grefenstette et al. (2015) can recognize hierarchical (Context-Free) languages like Dyck-kk, but LSTMs and RNNs cannot.

Korsky and Berwick (2019) prove that, given infinite precision, RNNs can recognize context-free languages. Their proof construction uses the floating point precision to simulate a stack, e.g., implementing push by dividing the old floating point value by 22, and pop by multiplying by 22. This implies that one can recognize any language requiring a bounded stack, like our Dyck-(kk,mm), by providing the model with precision that scales with stack depth. In contrast, our work assumes that the precision cannot scale with the stack depth (or vocabulary size); in practice, neural networks are used with a fixed precision Hubara et al. (2017).

Our work also connects to empirical studies of what RNNs can learn given finite samples. Considerable evidence has shown that LSTMs can learn languages requiring counters (but Simple RNNs do not) Weiss et al. (2018); Sennhauser and Berwick (2018); Yu et al. (2019); Suzgun et al. (2019), and neither Simple RNNs nor LSTMs can learn Dyck-kk. In our work, this conclusion is foregone because Dyck-kk requires unbounded memory while RNNs have finite memory; we show that LSTMs extrapolate well on Dyck-(kk,mm), the memory-bounded variant of Dyck-kk. Once augmented with an external (unbounded) memory, RNNs have been shown to learn hierarchical languages Suzgun et al. (2019); Hao et al. (2018); Grefenstette et al. (2015); Joulin and Mikolov (2015). Finally, considerable study has gone into what RNN LMs learn about natural language syntax Lakretz et al. (2019); Khandelwal et al. (2018); Gulordava et al. (2018); Linzen et al. (2016).

Preliminaries and definitions

A formal language L\mathcal{L} is a set of strings L⊆Σ∗ω\mathcal{L}\subseteq\Sigma^{*}\omega over a fixed vocabulary, Σ\Sigma (with the end denoted by special symbol ω\omega). We denote an arbitrary string as w1:T∈Σ∗ωw_{1:T}\in\Sigma^{*}\omega, where TT is the string length.

The Dyck-kk language is the language of nested brackets of kk types, and so has 2k2k words in its vocabulary: \Sigma=\{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i},{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}\}_{i\in[k]}. Any string in which brackets are well-nested, i.e., each {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} is closed by its corresponding {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}, is in the language. Formally, we can write it as the strings generated by the following context-free grammar:

where ϵ\epsilon is the empty string.And ω\omega appended to the end. The memory necessary to generate any string in Dyck-kk is proportional to the number of unclosed open brackets at any time. We can formalize this simply by counting how many more open brackets than close brackets there are at each timestep:

where count(w1:t,a)\texttt{count}(w_{1:t},a) is the number of times aa occurs in w1:tw_{1:t}. We can now define Dyck-(kk,mm) by combining Dyck-kk with a depth bound, as follows: {restatable}[Dyck-(kk,mm) ]defndefn_dyckkm For any positive integers k,mk,m, Dyck-(kk,mm) is the set of strings

2 Recurrent neural networks

We now provide definitions of recurrent neural networks as probability distributions over strings, and define what it means for an RNN to generate a formal language. We start with the most basic form of RNN we consider, the Simple (Elman) RNN: {restatable}[Simple RNN (generator)]defndefnsimplernn A Simple RNN (generator) with dd hidden units is a probability distribution fθf_{\theta} of the following form:

The Long Short-Term Memory (LSTM) model Hochreiter and Schmidhuber (1997) is a popular extension to the Simple RNN, intended to ease learning by resolving the vanishing gradient problem. In this work, we’re not concerned with learning but with expressivity. We study whether the LSTM’s added complexity enables it to generate Dyck-(kk,mm) using less memory. {restatable}[LSTM (generator)]defndefnlstmgenerator An LSTM (generator) with dd hidden units is a probability distribution fθf_{\theta} of the following form:

3 Formal language generation

With this definition of RNNs as generators, we now define what it means for an RNN (a distribution) to generate a language (a set). Intuitively, since a formal language is a set of strings L\mathcal{L}, our definition should be such that a distribution generates L\mathcal{L} if its probability mass on the set of all strings Σ∗ω\Sigma^{*}\omega is concentrated on the set L\mathcal{L}. So, we first define the set of strings on which a probability distribution concentrates its mass. The key intuition is to control the local token probabilities fθ(wt∣w1:t−1)f_{\theta}(w_{t}|w_{1:t-1}), not the global fθ(w1:T)f_{\theta}(w_{1:T}), which must approach zero with sequence length. {restatable}[locally ϵ\epsilon-truncated support]defndefnepsilontruncated Let fθf_{\theta} be a probability distribution over Σ∗ω\Sigma^{*}\omega, with conditional probabilities fθ(wt∣w1:t−1)f_{\theta}(w_{t}|w_{1:t-1}). Then the locally ϵ\epsilon-truncated support of the distribution is the set

This is the set of strings such that the model assigns at least ϵ\epsilon probability mass to each token conditioned on the prefix leading up to that token. A distribution generates a language, then, if there exists an ϵ\epsilon such that the locally truncated support of the distribution is equal to the language:We also note that any fθf_{\theta} generates multiple languages, since one can vary the parameter ϵ\epsilon; for example, any softmax-defined distribution must generate Σ∗\Sigma^{*} with ϵ\epsilon small because they assign positive mass to all strings.

A probability distribution fθf_{\theta} over Σ∗\Sigma^{*} generates a language L⊆Σ∗\mathcal{L}\subseteq\Sigma^{*} if there exists ϵ>0\epsilon>0 such that the locally ϵ\epsilon-truncated support of fθf_{\theta} is L\mathcal{L}.

Formal results

For our results, we first present two theorems for the Simple RNN and LSTM that use O(mk)O(mk) hidden units by simulating a stack of mm O(k)O(k)-dimensional vectors, which are useful for discussing the constructions. Then we show how to reduce to O(mlog⁡k)O(m\log k) via an efficient encoding of kk symbols in O(log⁡k)O(\log k) space.

Using the same mechanisms as in the proofs above but using an efficient encoding of each stack element in O(log⁡k)O(\log k) units, we achieve the following.

While we have emphasized expressive power under memory constraints—what functions can be expressed, not what is learned in practice—neural networks are frequently intentionally overparameterized to aid learning Zhang et al. (2017); Shwartz-Ziv and Tishby (2017). Even so, known constructions for Dyck-(kk,mm) would require a number of hidden units far beyond practicality. Consider if we were to use a vocabulary size of 100,000100{,}000, and a practical depth bound of 33. Then if we were using a km+1k^{m+1} hidden unit construction to generate Dyck-(kk,mm), we would need 100,0004=1020100{,}000^{4}=10^{20} hidden units. By using our LSTM construction, however, we would need only 3×3×⌈log⁡2(100,000)⌉−1×3=1503\times 3\times\lceil\log_{2}(100{,}000)\rceil-1\times 3=150 hidden units, suggesting that networks of the size commonly used in practice are large enough to learn these languages.

Lower bound.

Stack constructions in Simple RNNs

The memory needed to close all the brackets in a Dyck-(kk,mm) prefix w1:tw_{1:t} can be represented as a stack of (yet unclosed) open brackets [{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}], m′≤mm^{\prime}\leq m; reading each new parenthesis either pushes or pops from this stack. Informally, all of our efficient RNN constructions generate Dyck-(kk,mm) by writing to and reading from an implicit stack that they encode in their hidden states. In this section, we present some challenges in a naive approach, and then describe a construction to solve these challenges. We provide only the high-level intuition; rigorous proofs are provided in the Appendix.

We start by describing what is achievable with an extended model family, second-order RNNs Rabusseau et al. (2019); Lee et al. (1986), which allow their recurrent matrix WW to be chosen as a function of the input (unlike any of the RNNs we consider.) Under such a model, we show how to store a stack of up to mm kk-dimensional vectors in mkmk memory. Such a representation can be thought of as the concatenation of mm kk-dimensional vectors in the hidden state, like this:

We call each kk-dimensional component a “stack slot”. If we want the first stack slot to always represent the top of the stack, then there’s a natural way to implement pop and push operations. In a push, we want to shift all the slots toward the bottom, so there’s room at the top for a new element. We can do this with an off-diagonal matrix WpushW_{\text{push}}:Note that only needing to store mm things means that when we push, there should be nothing in slot mm; otherwise, we’d be pushing element m+1m+1.

This would implement the Wht−1Wh_{t-1} part of the Simple RNN equation. We can then write the new element (given by UxtUx_{t}) to the first slot. If we wanted to pop, we could do so with another off-diagonal matrix WpopW_{\text{pop}}, shifting everything towards the top to get rid of the top element:

This won’t work for a Simple RNN because it only has one WW.

2 A Simple RNN Stack in 2​m​k2𝑚𝑘2mk memory

Our construction gets around the limitation of only having a single WW matrix in the Simple RNN by doubling the space to 2mk2mk. Splitting the space hh into two mkmk-sized partitions, we call one hpoph_{\text{pop}}, the place where we write the stack if we see a pop operation, and the other hpushh_{\text{push}} analogously for the push operation. If one of hpoph_{\text{pop}} or hpushh_{\text{push}} is empty (equal to ) at any time, we can try reading from both of them, as follows:

Our WW matrix is actually the concatenation of two of the WpopW_{\text{pop}} and WpushW_{\text{push}} matrices. Now we have two candidates, both hpushh_{\text{push}} and hpoph_{\text{pop}}; but we only want the one that corresponds to push if w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, or pop if w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}. We can zero out the unobserved option using the term UxtUx_{t}, adding a large negative value to every hidden unit in the stack that doesn’t correspond to push if xtx_{t} is an open bracket {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, or pop if xtx_{t} is a close bracket {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}:

So Ux_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}}=[e_{i};0,\dots,-\beta-1,\dots], where eie_{i} is a one-hot representation of {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, and Ux_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}}=[-\beta-1,\dots,-\beta-1,0,\dots], where β\beta is the very large value we specified in our finite-precision arithmetic. Thus, when we apply the sigmoid to the new Wht−1+Uxt+bWh_{t-1}+Ux_{t}+b, whichever of {hpop,hpush}\{h_{\text{pop}},h_{\text{push}}\} doesn’t correspond to the true new state is zeroed out.For whichever of htmp∈{hpop,hpush}h_{\text{tmp}}\in\{h_{\text{pop}},h_{\text{push}}\} that is not zeroed out, σ(htmp)≠htmp\sigma(h_{\text{tmp}})\not=h_{\text{tmp}}. Hence, we scale all of UU and WW to be large, such that σ(Wht−1)∈{0,1}\sigma(Wh_{t-1})\in\{0,1\}.

Stack construction in LSTMs

We could implement our 2mk2mk Simple RNN construction in an LSTM, but its gating functions suggest more flexibility in memory management, and Levy et al. (2018) claim that the LSTM’s modeling power stems from its gates. With an LSTM, we achieve the mkmk of the oracle we described, all while exclusively using its gates.

To implement a stack using the LSTM’s gates, we use the same intuitive description of the stack as before: mm stack slots, each of dimensionality kk. However while the top of the stack is the first slot in the Simple RNN, the bottom of the stack is the first slot in the LSTM. Before we discuss mechanics, we introduce the memory dynamics of the model. Working through the example in Figure 2, when we push the first open bracket, it’s assigned to slot 11; then the second and third open brackets are assigned to slots 22 and 33. Then a close bracket is seen, so slot 33 is erased. In general, the stack is represented in a contiguous sequence of slots, where the first slot represents the bottom of the stack. Thus, the top of the stack could be at any of the mm stack slots. So to allow for ease of linearly reading out information from the stack, we store the full stack only in the cell state ctc_{t}, and let only the slot corresponding to the top of the stack, which we’ll refer to as the top slot, through the output gate to the hidden state hth_{t}.

pop mechanics.

To implement a pop operation, the forget gate finds the top slot, which is the slot farthest from slot 11 that isn’t empty (that is, that encodes some {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}.) In practice, we do this by guaranteeing that the forget gate is equal to 11 for all stack slots before (but excluding) the last non-empty stack slot. Since this last non-empty stack slot encodes the top of the stac,, for it and all subsequent (empty) stack slots, the forget gate is set to .The input gate and the output gate are both set to 0\mathbf{0}. This erases the element at the top of the stack, summarized in the following diagram:

output mechanics.

We’ve so far described how the new cell state ctc_{t} is determined. So that it’s easy to tell what symbol is at the top of the stack, we only want the top slot of the stack passing through the output gate. We do this by guaranteeing that the output gate is equal to for all stack slots from the first slot to the top slot (exclusive). The output gate is then set to 11 for this top slot (and all subsequent empty slots,) summarized in the following diagram:

Defining the generating distribution

So far, we’ve discussed how to implement implicit stack-like memory in the Simple RNN and the LSTM. However, the formal claims we make center around RNNs’ ability to generate Dyck-(kk,mm).

Assume that at any timestep tt, our stack mechanism has correctly pushed and popped each open and close bracket, as encoded in ht,cth_{t},c_{t}. We still need to prove that our probability distribution,

assigns greater than ϵ\epsilon probability only to symbols that constitute continuations of some string in Dyck-(kk,mm), by specifying the parameters VV and bvb_{v}.

If and only if fewer than mm elements are on the stack, all {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} must be assigned ≥ϵ\geq\epsilon probability. This encodes the depth bound. In our constructions, mm elements are on the stack if and only if stack slot mm is non-zero. So, each row V_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}} is zeros except for slot mm, where each dimension is a large negative number, while the bias term bvb_{v} is positive.

Observing the end of the string ω𝜔\omega.

If and only if elements are on the stack, the string can end. The row VωV_{\omega} detects if any stack slot is non-empty.In particular, the bias term bωb_{\omega} is positive, but the sum Vωht−1+bωV_{\omega}h_{t-1}+b_{\omega} is negative if the stack is not empty.

The close bracket {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i} can be observed if and only if the top stack slot encodes {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}. Both the Simple RNN construction and the LSTM construction make it clear which stack slot encodes the top of the stack. In the Simple RNN, it’s always the first slot. In the LSTM, it’s the only non-empty slot in hth_{t}. In our stack constructions, we assumed each stack slot is a kk-dimensional one-hot vector eie_{i} to encode symbol {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}. So in the Simple RNN, V_{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}{}_{i} reads the top of the stack through a one-hot vector eie_{i} in slot 11, while in the LSTM it does so through eie_{i} in all mm slots. This ensures that V_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}}^{\top}h_{t} is positive if and only if {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} is at the top of the stack.

2 Generation in O​(m​log⁡k)𝑂𝑚𝑘O(m\log k) memory

Thus, we can always detect which symbol {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} is encoded by setting b_{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}{}_{i}=\log k-0.5.In actuality, we use a slightly less compact encoding, spending log⁡k−1\log k-1 more hidden units set to 11, to incrementally subtract log⁡k−1\log k-1 from all logits. Then the bias terms b_{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}{}_{i} are set to 0.50.5, avoiding possible precision issues with representing the float ⌈log⁡k⌉−0.5\lceil\log k\rceil-0.5.

Experiments

Our proofs have concerned constructing RNNs that generate Dyck-(kk,mm); now we present a short study connecting our theory to learning Dyck-(kk,mm) from finite samples. In particular, for k∈{2,8,32,128}k\in\{2,8,32,128\} and m∈{3,5}m\in\{3,5\}, we use our theory to set the hidden dimensionality of LSTMs to 3m⌈log⁡k⌉−m3m\lceil\log k\rceil-m, and train them as LMs on samples from a distributionDefined in Appendix H over Dyck-(kk,mm). For space, we provide an overview of the experiments, with details in the Appendix. We evaluate the models’ abilities to extrapolate to unseen lengths by setting a maximum length of 8484 for m=3m=3, and 180180 for m=5m=5, and testing on sequences longer than those seen at training time.Up to twice as long as the training maximum. For our evaluation metric, let pjp_{j} be the probability that the model predicts the correct closing bracket given that jj tokens separate it from its open bracket. We report meanjpj\text{mean}_{j}p_{j}, to evaluate the model’s bracket-closing memory.

For all configurations, we find that the LSTMs using our memory limit achieve error less than 10−410^{-4} when trained on 20 million tokens. Strikingly, this is despite the fact that for large mm and kk, a small fraction of the possible stack states is seen at training time;See Table 3. this shows that the LSTMs are not simply learning kmk^{m} structureless DFA states. Learning curves are provided in Figure 3.

Discussion and conclusion

We proved that finite-precision RNNs can generate Dyck-(kk,mm), a canonical family of bounded-depth hierarchical languages, in O(mlog⁡k)O(m\log k) memory, a result we also prove is tight. Our constructions provide insight into the mechanisms that RNNs and LSTMs can implement.

The Chomsky hierarchy puts all finite memory languages in the single category of regular languages. But humans generating natural language have finite memory, and context-free languages are known to be both too expressive and not expressive enough Chomsky (1959); Joshi et al. (1990). We thus suggest the further study of what structure networks can encode in their memory (here, stack-like) as opposed to (just) their position in the Chomsky hierarchy. While we have settled the representation question for Dyck-(kk,mm), many open questions still remain: What broader class of bounded hierarchical languages can RNNs efficiently generate? Our experiments point towards learnability; what class of memory-bounded languages are efficiently learnable? We hope that answers to these questions will not just demystify the empirical success of RNNs but ultimately drive new methodological improvements as well.

Code for running our experiments is available at https://github.com/john-hewitt/dyckkm-learning. An executable version of the experiments in this paper is on CodaLab at https://worksheets.codalab.org/worksheets/0xd668cf62e9e0499089626e45affee864.

Acknowledgements

The authors would like to thank Nelson Liu, Amita Kamath, Robin Jia, Sidd Karamcheti, and Ben Newman. JH was supported by an NSF Graduate Research Fellowship, under grant number DGE-1656518. Other funding was provided by a PECASE award.

References

Appendix A Appendix outline

This Appendix has the following order. In (§B), we provide a definition of Dyck-(kk,mm) equivalent to that in the main text but more useful for our proofs. In (§C), we state preliminary definitions and assumptions, and prove the lower-bound of Ω(mlog⁡k)\Omega(m\log k) hidden units to generate Dyck-(kk,mm). In (§D), we formally introduce our Simple RNN stack construction, and prove its correctness in a lemma. In (§E), we formally introduce our LSTM stack construction, and prove its correctness in a lemma. In (§F), we prove that a linear (+softmax) decoder on the Simple RNN and LSTM hidden states can be used to generate Dyck-(kk,mm) in O(mk)O(mk) hidden units using a 11-hot encoding of stack elements. In this section we also prove that a general RNN construction of DFAs allows for generation of Dyck-(kk,mm) in O(km+1)O(k^{m+1}) hidden units. In (§G), we provide an alternative encoding of elements in our stack constructions for the Simple RNN and LSTM that uses O(log⁡k)O(\log k) space per element, and prove that our stack constructions still hold using this encoding. We provide a linear (+softmax) decoder on the hidden states of the Simple RNN and LSTM (when using the O(log⁡k)O(\log k) stack element encoding) that can be used as a drop-in replacement for the decoder from the 11-hot representations, thus generating Dyck-(kk,mm). This proving that both the Simple RNN and LSTM generate Dyck-(kk,mm) in O(mlog⁡k)O(m\log k) hidden units.

Appendix B The Dyck-(k𝑘k,m𝑚m) languages

To better understand the success of neural networks on natural language syntax, we aim for a formal language that models the unbounded recursiveness of natural language while also reflecting its bounded memory requirements. We thus introduce the Dyck-(kk,mm) languages, corresponding to sequences of balanced brackets of kk types with a maximal number of mm unclosed brackets at any point in the sequences (yielding a bound stack depth of mm to parse such sentences).

Though in the main text we defined Dyck-(kk,mm) by intersecting Dyck-kk with a language that simply bounds the difference between the number of open brackets and the number of close brackets, here we provide an equivalent definition that will aid in our proofs. Each Dyck-(kk,mm) language, specified by fixing a value of mm and kk, is defined using a deterministic finite automaton (DFA). Here, we provide a general description of any Dyck-(kk,mm) DFA.

Formally, we define each language by the deterministic finite automaton Dm,k=(Q,Σ,δ,F,q0)\mathcal{D}_{m,k}=(Q,\Sigma,\delta,F,q_{0}). The vocabulary Σ\Sigma consists of kk types of open brackets: \{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}\}_{i=1,\dots,k} corresponding closing brackets \{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}\}_{i=1,\dots,k}. Strings over the vocabulary are w1:T∈Σ∗ωw_{1:T}\in\Sigma^{*}\omega, where ω∉Σ\omega\not\in\Sigma is a special symbol representing the end of the sequence. Σ∪{ω}\Sigma\cup\{\omega\} are collectively referred to as symbols.

This slightly nonstandard requirement allows for a natural connection with language models, which must estimate the probability that a string ends at any given token.But crucially, since ω∉Σ\omega\not\in\Sigma, the language model need not define a distribution after ω\omega is seen; it can only be the last token. Overloading notation, we’ll also use Dm,k\mathcal{D}_{m,k} to refer to the language (the set of strings) itself, defined as the strings accepted by the DFA.

We now define the states q∈Qq\in Q. First, we define reject state rr, and accept state [ω][\omega]. Each other state is uniquely identified by a list of open bracket symbols of length up to mm; thus the full set of states is provided by:

The number of states is thus km+1+1k^{m+1}+1, where all but two states reflect a list of open brackets. We will denote each list [{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}}{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{2}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}], 0≤m′≤m0\leq m^{\prime}\leq m as a stack state with m′m^{\prime} elements, and the value of {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}} as the top element of the stack. We let q_{0}=[]=[{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}}{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{2}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}], m′=0m^{\prime}=0, the empty stack state.

B.2 Transition function, δ𝛿\delta

We now define the transition function, δ\delta.

The state [][] can transition either to the accept state or to another stack state:

while any other symbol transitions to the reject state.

For any state of the form [{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}}{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{2}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}], where m′<mm^{\prime}<m, an open bracket can be pushed to the list (since m′<mm^{\prime}<m), or the last open bracket can be removed, by observing its corresponding close bracket.

All other symbols transition to the reject state.

For any state of the form [{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}}{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{2}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m}}], that is, the list is of length mm and thus full, the close bracket of the top of the stack removes that bracket,

All other symbols transition to the reject state.

The reject state rr transitions to itself for all symbols.

No transitions from the accept state need be defined, since only δ([],ω)=[ω]\delta([],\omega)=[\omega], and ω\omega must be the last symbol of any string in the universe, and can only occur once.

This accounts for the transition function from all states in QQ, and completes our definition of the DFA Dm,kD_{m,k}. We are now prepared to formally define the language:

Overloading notation, for any w1:T∈Σ∗ωw_{1:T}\in\Sigma^{*}\omega, for any t∈1,…,Tt\in 1,\dots,T, we’ll say qt=Dm,k(w1:t)q_{t}=D_{m,k}(w_{1:t}), denoting the state that Dm,kD_{m,k} is in after consuming w1:tw_{1:t}.

Appendix C Preliminaries

* Through the notion of ϵ\epsilon-truncated support, a language model specifies which tokens are allowable continuations of each string prefix by assigning them greater than ϵ\epsilon probability. With this, we’re ready to connect RNNs and formal languages:

A probability distribution fθf_{\theta} generates a formal language L\mathcal{L} if there exists ϵ>0\epsilon>0 such that the ϵ\epsilon-truncated support Lfθ\mathcal{L}_{f_{\theta}} is equal to L\mathcal{L}.

Under our finite-precision setting, a reasonable assumption about the properties of the sigmoid and hyperbolic tangent functions make our claim considerably simpler. The sigmoid function, σ(x)=11+e−x\sigma(x)=\frac{1}{1+e^{-x}}, has range (0,1)(0,1), excluding its boundaries {0,1}\{0,1\}. However, in a finite-precision arithmetic, σ(x)\sigma(x) cannot become arbitrarily close to or 11. In fact, in popular deep learning library PyTorchUnder the float datatype, PyTorch v1.3.1 Paszke et al. (2019)., σ(x)\sigma(x) is exactly equal to 11 for all x>6x>6. We define the floating point sigmoid function to equal or 11 if the absolute value of its input is greater or equal in absolute value to some threshold β\beta:

Similarly for the hyperbolic tangent function, we define:

For the rest of this paper, we will refer to σfp\sigma_{\text{fp}} as σ\sigma, and tanhfp\text{tanh}_{\text{fp}} as tanh.

C.2 Lower bound of Ω​(m​log⁡k)Ω𝑚𝑘\Omega(m\log k) for generation of Dyck-(k𝑘k,m𝑚m)

We provide a communication complexity argument. Assume for contradiction that there exists such a machine with d<mlog⁡kpd<\frac{m\log k}{p}. Consider any string w_{1:m}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m}} of mm open brackets, where each is one of kk types. There are kmk^{m} such strings, but because d<mlog⁡kpd<\frac{m\log k}{p}, we have that the model has 2dp<km2^{dp}<k^{m} possible representations, so at least two such strings must share the same representation. Let such a pair be w≠w′w\not=w^{\prime}, where likewise w′w^{\prime} is defined as {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i^{\prime}_{1}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i^{\prime}_{m}}. Consider the sequence of open bracket indices that defines ww: i1,…,imi_{1},\dots,i_{m}. Let ss be the string {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i_{m}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i_{1}} of close brackets in the reverse order of the indices of ww, and likewise s′s^{\prime} for w′w^{\prime}.

The string w::sw::s, where :::: denotes concatenation, is in Dyck-(kk,mm). However, w′::sw^{\prime}::s is not in Dyck-(kk,mm) as it breaks the well-nested brackets condition wherever {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{j}}\not={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i^{\prime}_{j}}. This must occur at least once, since w≠w′w\not=w^{\prime}. Since the model assigns them the same representation even though they must be distinguished to generate Dyck-(kk,mm), the model does not generate Dyck-(kk,mm).

Appendix D Proving Simple RNN stack correspondence in 2​m​k2𝑚𝑘2mk units

In this section, we provide a formal description of our Simple RNN stack construction, and introduce and prove the stack correspondence lemma, to guarantee its correctness. \defnsimplernn*

where [… ][\dots] denotes concatenation of vectors into a kmkm-dimensional vector. We note that RR is injective, that is, the inverse map R−1R^{-1} is well-defined on all vectors in the image of RR.

Later, when constructing the efficient O(mlog⁡k)O(m\operatorname{log}k) memory encoding in Section G.1, we will swap the one-hot vectors eie_{i} for other vectors that will also have entries in {0,1}\{0,1\}.Looking ahead to the more complex LSTM construction, we’ll introduce notation to denote these vectors e_{i}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}), and look to replace ψ\psi when we develop the O(mlog⁡k)O(m\log k) construction.

If h=S(σ,q)h=\mathcal{S}(\sigma,q), then the definition guarantees that Q(h)=q\mathcal{Q}(h)=q.

Defining Transition Matrices

To define the RNN (apart from VV, bvb_{v}), we have to specify the matrices W,U,E,W,U,E,, and the vector bb.

To specify the matrices EE and UU, we need to fix an assignment to the integers 1,…,2k1,\dots,2k to the symbols in Σ\Sigma. We will assign the integers 1,…,k1,\dots,k to the opening brackets (i.e., {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{1}\mapsto 1,\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{k}\mapsto k), and the integers 1,…,k1,\dots,k to the closing brackets (i.e., {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{1}\mapsto k+1,\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{k}\mapsto 2k).

Here, we write Ik×kI_{k\times k} for the identity matrix, and 1k×k{\bf 1}_{k\times k} for a matrix filled entirely with ones.

Proof of Correctness

To prove correctness of the construction, the first step is to show that the transition dynamics of the Simple RNN correctly simulates the dynamics of the stack when consuming one symbol.

Let \phi({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})=push, \phi({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i})=pop.

Assume that the hidden state hth_{t} encodes a stack qq:

Let wt+1w_{t+1} be the new input symbol, and assume that w1:t+1w_{1:t+1} is a valid prefix of a word in Dyck-(kk,mm). Let ht+1h_{t+1} be the next hidden state, after reading wt+1w_{t+1}. Then

Before showing this lemma, we note the following property of the activation function σ\sigma, under the finite precision assumption:

We can write ht=[hpush,hpop]h_{t}=[h_{push},h_{pop}]. Set h′=hpush+hpoph^{\prime}=h_{push}+h_{pop}. By construction of S\mathcal{S},

At this point, we note that the only relevant property of the encodings hi′h^{\prime}_{i} for this proof is that their entries are in {0,1}\{0,1\}, not that they are one-hot vectors. This will make it possible to plug in a more efficient O(log⁡k)O(\log k) encoding later in Section G.1.

Case 1:

Assume w_{t+1}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}. We have

which is equal to S(push,\delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})).

Case 2:

Assume w_{t+1}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}. We have

which is equal to S(push,\delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i})). ∎

From the previous lemma, we can derive the following Stack Correspondence lemma, which asserts that the Simple RNN correctly reflects stack dynamics over an entire input string:

For all strings w1:Tw_{1:T} in Dyck-(kk,mm), for all t=1,...,Tt=1,...,T, let qtq_{t} be the state of Dm,kD_{m,k} after consuming prefix w1:tw_{1:t}. Let hth_{t} be the RNN hidden state after consuming w1:tw_{1:t}. Then if qt∉{[ω]}q_{t}\not\in\{[\omega]\},

The claim is shown by induction over tt, applying the previous lemma in each step. We will show the claim for all t=0,1,…,Tt=0,1,\dots,T, setting h0h_{0} to be the zero vector 02km{\bf 0}_{2km}.

As Q(h0)=q0\mathcal{Q}(h_{0})=q_{0}, this proves the claim for t=0t=0. To prove the inductive step, we assume Q(ht)=qt\mathcal{Q}(h_{t})=q_{t} has already been shown.

Noting qt+1=δ(qt,wt+1)q_{t+1}=\delta(q_{t},w_{t+1}), we obtain

As noted in the definition of Q\mathcal{Q} above (D.3), this entails

Appendix E Proving LSTM stack correspondence in m​k𝑚𝑘mk units

In this section, we prove our lemma describing the stack construction implemented by an LSTM. To do so, we first define the LSTM: \defnlstmgenerator*

First, we provide notation for describing the structure of the memory cell (ctc_{t}) that fθf_{\theta} will maintain. Next, we proceed to describe precisely, but without referring to the LSTM equations, how this memory changes over time. We then describe how these high-level dynamics are implemented by the LSTM equations, assuming that we’re able to set the values of the intermediate values (gates and new cell candidate) as desired. We then explicitly construct LSTM weight matrices that provide the desired intermediate values. This completes the proof of the memory dynamics, which we formalize in a stack correspondence lemma once we’ve introduced the notation.

Applying this function to each stack slot allows us to map from st,1,…,st,ms_{t,1},\dots,s_{t,m} to stack states:

where m′≤mm^{\prime}\leq m is the maximal integer such that ψ(st,m′)≠∅\psi(s_{t,m^{\prime}})\not=\varnothing. Intuitively, this means we filter out any empty stack slots at the end of the sequence (those slots st,js_{t,j} where j>m′j>m^{\prime}.) We’ll also make use of the inverse, Q−1\mathcal{Q}^{-1}, mapping DFA stack states to stack slot lists,

where slots sm′+1,…,sms_{m^{\prime}+1},\dots,s_{m} are the zero vector 0\mathbf{0} because only the first m′m^{\prime} slots encode a symbol {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, and ψ(0)=∅\psi(\mathbf{0})=\varnothing.

These one-hot encodings (eie_{i}) of symbols ({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}) allow for clear exposition, but are not fundamental to the construction; in (§ G), we replace these kk-dimensional one-hot encodings given by ψ\psi with O(log⁡k)O(\log k)-dimensional encodings. Throughout our description of the stack construction, we’ll explicitly rely on the following property, so we can easily replace this particular ψ\psi in (§ G) by providing another ψ′\psi^{\prime} with the same property.

that is, the sum of all dimensions of the representation is equal to 11, and

that is, the representation takes on values only in and 11, and

that is, the empty stack slot is encoded by the zero vector, and

that is, encodings of symbols {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} are unique, and none are equal to to the encoding of the empty symbol ∅\varnothing.

The encoding we’ve so far provided, \psi^{-1}(e_{i})={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, obeys these properties: the sum of a single one-hot vector is 11, no two such vectors are the same, one-hot vectors only take on values in {0,1}\{0,1\}, and we let ψ(0)=∅\psi(\mathbf{0})=\varnothing.

E.2 Description of stack state dynamics

Before we describe the memory dynamics, we introduce a useful property of the stack slots.

A cell state ctc_{t} composed of stack slots st,1,…,st,ms_{t,1},\dots,s_{t,m} is jj-top if there exists i∈{1,…,k}i\in\{1,\dots,k\} such that s_{t,j}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}), and for all j′∈{j+1,…,m}j^{\prime}\in\{j+1,\dots,m\}, st,j′=0s_{t,j^{\prime}}=\mathbf{0}.

Intuitively, we’ll enforce the constraint that a cell state that is jj-top encodes the element {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{m^{\prime}} at the top of the stack [{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{m^{\prime}}] in stack slot j=m′j=m^{\prime}.

We’re now ready to describe the memory dynamics of fθf_{\theta}. The memory ctc_{t} is initialized to 0\mathbf{0}. So for all j∈{1,…,m}j\in\{1,\dots,m\},

The recurrent equation of stack slots is then

Intuitively, the first case of Equation E.10 implements pop, the second case implements push (if pushing to an otherwise empty stack) and the third case implement push (if pushing to a non-empty stack). The fourth case specifies that slots that are not pushed to or popped are maintained from timestep t−1t-1 to timestep tt.

Intuitively, this means only the top of the stack is stored in the hidden state, and it’s stored in whichever units correspond to the jj-top slot in ctc_{t}. Since \psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}) takes on values in {0,1}\{0,1\}, \text{tanh}(\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})) takes on values in {0,tanh(1)}\{0,\text{tanh}(1)\}.

We’re now ready to formalize the first part of our proof in a lemma that encompasses how fθf_{\theta} structures and manipulates its memory.

For all strings w1:Tw_{1:T} in Dyck-(kk,mm), for all t=1,...,Tt=1,...,T, let qtq_{t} be the state of Dm,kD_{m,k} after consuming prefix w1:tw_{1:t}. Let ctc_{t} be the hidden state of the LSTM after consuming w1:tw_{1:t}. Then if qt=∉{[ω]}q_{t}=\not\in\{[\omega]\},

and letting q_{t}=[{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}] for m′≤mm^{\prime}\leq m without loss of generality,

where if qt=[]q_{t}=[], then m′=0m^{\prime}=0 and ht,j=0h_{t,j}=\mathbf{0} for all jj.

Note again that the value of ψ−1\psi^{-1} is not specified by the lemma; whereas so far we’ve used \psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})=e_{i}, a one-hot encoding, in (§ G) we’ll replace ψ−1\psi^{-1} while maintaining the lemma.

Intuitively, Equation E.12 states that fθf_{\theta} keeps an exact representation of the stack of unclosed open brackets so far, a statement solely concerning the cell state ctc_{t}. However, all intermediate values in the LSTM equations are functions of the hidden state hth_{t}, not the cell state directly; Equation E.13 states how the top of the stack is represented in the hidden state, hth_{t}. We state these as one lemma since it is convenient to prove them by induction together.

E.3 Proof of the stack correspondence lemma assuming stack slot dynamics

We first prove Lemma 4 under the assumption that the stack slot dynamics given in Equation E.10 hold. We proceed by induction on the prefix length.

When the prefix length is , the DFA is in state [][], and s0,j=0s_{0,j}=\mathbf{0} for all jj, so Q(s0,1,…,s0,m)=[]\mathcal{Q}(s_{0,1},\dots,s_{0,m})=[], as required. Further, we have h0=0h_{0}=\mathbf{0}, as required. This completes the base case.

Assume for the sake of induction that Lemma 4 holds for strings up to length tt. Now consider a prefix w1:t+1w_{1:t+1}. From it, we can take the prefix w1:tw_{1:t}. Running fθf_{\theta}, we have st,1,…,st,ms_{t,1},\dots,s_{t,m}. By the induction hypothesis, we have that \mathcal{Q}(s_{t,1},\dots,s_{t,m})=q_{t}=[{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}], (Equation E.12) as well as that

by Equation E.13. Now, the symbol wt+1w_{t+1} can be one of 2k+12k+1 symbols: any of the kk open brackets {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, the kk close brackets {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}, or ω\omega. The proof now proceeds by each of these three cases.

Without loss of generality, we let q_{t}=[{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}] for m′<mm^{\prime}<m. Strict inequality is guaranteed since if m′=mm^{\prime}=m, then \delta(q_{t},{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})=r, meaning w1:t+1w_{1:t+1} cannot be a prefix of any string in Dyck-(kk,mm), contradicting a premise of the lemma. Through the inverse of the stack mapping Equation E.3, we have

the stack slots of fθf_{\theta} at timestep tt. Since s_{t,m^{\prime}}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{j}}) and st,m′+1=0s_{t,m^{\prime}+1}=\mathbf{0}, we have that ctc_{t} is m′m^{\prime}-top. Thus, since w_{t+1}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, the second or third condition in Equation E.10 (depending on whether m′=0m^{\prime}=0) dictate that s_{t+1,m^{\prime}+1}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}) (where the ii is because w_{t+1}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}). All other stack slots fall under the condition st+1,j=st,js_{t+1,j}=s_{t,j}.

With this, we can reason about the DFA state encoded by ct+1c_{t+1} at timestep w1:t+1w_{1:t+1}:

which is the DFA’s state at timestep t+1t+1, as required.

Finally, we reason about ht+1,jh_{t+1,j}. We’ve just shown that that ct+1c_{t+1} is (m′+1)(m^{\prime}+1)-top, and that s_{t+1,m^{\prime}+1}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}). By Equation E.11, we have that h_{t+1,m^{\prime}+1}=\text{tanh}(\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})), and ht+1,j=0h_{t+1,j}=\mathbf{0} for all j≠m′+1j\not=m^{\prime}+1 as required, completing this case.

We have q_{t}=[{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}] for 0<m′0<m^{\prime}. this inequality is strict since if m′=0m^{\prime}=0, then qt=[]q_{t}=[], and \delta(q_{t},{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i})=r, meaning w1:t+1w_{1:t+1} is not a prefix of any string in Dm,kD_{m,k}. Thus, we have that ctc_{t} is m′m^{\prime}-top. Since w_{t+1}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}, we have that st+1,m′=0s_{t+1,m^{\prime}}=\mathbf{0}, and st+1,j=st,js_{t+1,j}=s_{t,j} for all j≠m′j\not=m^{\prime}. Thus, we have

and thus the DFA state corresponding to the stack slots is:

Which is \delta(q_{t},{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}), as required. By the same reasoning as for case w_{t+1}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, since ct+1c_{t+1} is (m′−1)(m^{\prime}-1)-top and that s_{t+1,m^{\prime}-1}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{m^{\prime}-1}), we have that h_{t+1,m^{\prime}-1}=\text{tanh}(\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})) and ht+1,j=0h_{t+1,j}=\mathbf{0} for all j≠m′−1j\not=m^{\prime}-1, as required. This completes the case.

𝑡1𝜔w_{t+1}=\omega: If wt+1=ωw_{t+1}=\omega, then by the definition of δ\delta, δ(qt,ω)∈{[ω],r}\delta(q_{t},\omega)\in\{[\omega],r\}; hence the premises of Lemma 4 don’t hold, so the lemma vacuously holds.

Summary

In this section we’ve proved the stack correspondence lemma for our LSTM construction, assuming we can implement the dynamics of the model as specified in Equations E.10 E.11; we have yet to show that the LSTM update can implement the equations as promised.

E.4 Implementation of stack state dynamics in LSTM equations assuming intermediate values

In this subsection, we define how the dynamics defined in Equations E.10, E.11 are implemented in the LSTM equations, referring to intermediate gate values, assuming such values can be reached with some setting of LSTM parameters.

How these fulfill Equations E.10, E.11? We can set the gate values to do so as follows. First, the forget gate is used to detect the condition under which slot jj should be erased:

Second, the new cell candidate expression is used to give the option of writing a new open bracket, ψ−1(wt)\psi^{-1}(w_{t}) to any stack slot:

Third, the input gate determines which stack slot is the first non-empty slot in the sequence, and only allows the information of the new open bracket to be written to that slot:

Finally, the output gate is set to identify the top of the stack and not let any other non-0\mathbf{0} (i.e., all stack slots below the top) through:

Many of these conditions refer to values, like ctc_{t} and ct−1c_{t-1}, not available when computing the gates (as only ht−1h_{t-1} and xtx_{t} are available); it is convenient to refer to these conditions and later show how they are implemented with the available values.

E.5 Proof of stack state dynamics given gate and new cell candidate values

In this section, we prove that Equations E.10, E.11 are implemented by the LSTM fθf_{\theta} given the gate and new cell candidate values defined in the previous subsection.

We start with Equation E.10, the definition of the stack slot dynamics, proceeding by each of the four cases for defining st,js_{t,j} in Equation E.10.

where the last equality hold because of the condition that ct−1c_{t-1} is (j−1)(j-1)-top, which implies st−1,j=0s_{t-1,j}=\mathbf{0}.

Proving Equation E.11 is implemented

Next, we prove that Equation E.11 is implemented by the LSTM fθf_{\theta}, given that Equation E.10 is. First, we must show ht,jh_{t,j} is equal to tanh(ei)\text{tanh}(e_{i}) if ctc_{t} is jj-top. If ctc_{t} is jj-top, it must be that st,j=eis_{t,j}=e_{i} for some ii. Since ctc_{t} is jj-top, we have ot,j=1o_{t,j}=\mathbf{1}, and st,j=eis_{t,j}=e_{i}. Plugging these values into the hidden state expression in Equation E.15, we get:

If that condition does not hold, we must have that ctc_{t} not jj-top. thus, ctc_{t} is either jj-top for some j′<jj^{\prime}<j or j′>jj^{\prime}>j, or ct=0c_{t}=\mathbf{0}.

If ctc_{t} is j′j^{\prime}-top for j′<jj^{\prime}<j, then we have that st,j=0s_{t,j}=\mathbf{0}, by the definition of jj-top. In this case, we have

If ctc_{t} is j′j^{\prime}-top for j′>jj^{\prime}>j, then we have that ot,j=0o_{t,j}=\mathbf{0}, by Equation E.19. In this case, we have

Finally, if ct=0c_{t}=\mathbf{0}, meaning it is not jj-top for any jj, then

Summary

So far, we’ve proved the stack correspondence lemma’s induction step for LSTMs assuming that we can provide parameters of the network such that the values assumed in Equations E.16, E.17, E.18, E.19, that is, the gates and new cell candidate, are achieved. We have yet to provide the settings of parameters to do so.

E.6 Construction and proof of LSTM parameters providing gate and new cell candidate values

The equation defining ot,jo_{t,j} is as follows:

and likewise for the second condition. The third condition relies on the fact that ψ−1(∅)=0\psi^{-1}(\varnothing)=\mathbf{0}.

that is, each row of UfU_{f} is equal to the sum, over all close brackets, of the embedding of that close bracket (transposed.) Since all x_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}} are orthogonal and unit norm (by the definition of EE), this gives the following intermediate value:

that is, all units of the bias are equal to the same value.

As we stated earlier, Equation E.19 refers to conditions on ctc_{t}, which do not show up in the computation of oto_{t}; instead, oto_{t} is a function of ht−1h_{t-1}, and xtx_{t}. As such, we re-write Equation E.19 in terms of these values, in particular splitting the condition on ctc_{t} being j′j^{\prime}-top for j′>jj^{\prime}>j into two separate conditions,

Recall that h_{t,j^{\prime}}=\text{tanh}(\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})) indicates that ct−1c_{t-1} is j′j^{\prime}-top, by the induction hypothesis. With our equation above we outline two different conditions that we’ll show together are necessary and sufficient to guarantee that the top slot is located at some j′>jj^{\prime}>j. Under each condition, we set the output gate to because of that (since if the top slot is at j′>jj^{\prime}>j, then slot jj is not equal to , and must not be let through the gate.)

While we cannot condition directly on the value of ctc_{t}, we know that ctc_{t} is j′j^{\prime}-top for j′>jj^{\prime}>j under two conditions on ht−1h_{t-1}, and thus on ct−1c_{t-1}. First, ct−1c_{t-1} can be (j+2)(j+2)-top or greater; since only one element can be popped at once, ctc_{t} can be no less than (j+1)(j+1)-top. Second, ct−1c_{t-1} can be jj-top or (j+1)(j+1)-top, and the input wtw_{t} is not some {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}; that is, it does not pop the top element off of the stack. Because wtw_{t} is not {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}, it must be an open bracket, {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, pushing to the stack. Hence, under these conditions, ctc_{t} must be at least (j+1)(j+1)-top. If neither of these conditions hold, then ctc_{t} cannot be j′j^{\prime}-top for j′>jj^{\prime}>j, since if ct−1c_{t-1} is jj-top or (j−1)(j-1)-top and w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}, then ctc_{t} must be at most jj-top. And if ct−1c_{t-1} is j′j^{\prime}-top for j′<jj^{\prime}<j, then it is impossible since only one bracket can be pushed at a time. Thus Equation E.27 is equivalent to Equation E.19, as required.

We now prove that parameters (Wo,Uo,bo)(W_{o},U_{o},b_{o}) implement Equation E.19, by implementing Equation E.27. In the first condition, h_{t-1,j^{\prime}}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}) for some j′>j+1j^{\prime}>j+1 and ii; we want to show that ot,j=0o_{t,j}=\mathbf{0}. Based on our construction of parameters, we have that ot,jo_{t,j} is upper-bounded by the following, when w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}:

where the last equality holds under our finite precision σ\sigma by setting the scaling factor λ\lambda such that −0.5λγ<−β-0.5\lambda\gamma<-\beta, that is λ>2βγ\lambda>\frac{2\beta}{\gamma}.

In the next condition, we have that ht−1,j′=eih_{t-1,j^{\prime}}=e_{i}for j′∈{j,j+1}j^{\prime}\in\{j,j+1\}, and ii, and x_{t}\not={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}. In this case, we can write out the value of ot,jo_{t,j} as follows:

Finally, if neither of those conditions hold, we have that h_{t-1,j^{\prime}}\not=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}) for any j′≥jj^{\prime}\geq j, and so ht−1,j′=tanh(0)=0h_{t-1,j^{\prime}}=\text{tanh}(\mathbf{0})=\mathbf{0}. In this case, we can lower-bound the value of ot,jo_{t,j}:

as required. This completes the proof of the output gate.

The equation defining ft,jf_{t,j} is as follows:

We define Wf,j,j′W_{f,j,j^{\prime}} as follows:

again relying on the embedding-validity property.

that is, each row of UfU_{f} is equal to the negative of the sum, over all close brackets, of the embedding of that close bracket (transposed.) Since all x_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}} are orthogonal and unit norm, this gives the following intermediate value:

that is, all units of the bias are equal to the same value.

We now prove that, as given, the parameters (Wf,Uf,bf)(W_{f},U_{f},b_{f}) implement Equation E.16 assuming the induction hypothesis of Lemma 4.

Equation E.16 has two cases. In the first case, ct−1c_{t-1} is jj-top and w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}, and we need to prove ft,j=0f_{t,j}=\mathbf{0}. From the induction hypothesis, this means that h_{t-1,j}=\text{tanh}(\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})), and all other slots ht−1,j′=0h_{t-1,j^{\prime}}=\mathbf{0}. The value of the gate is as follows:

where the last equality holds by setting λ\lambda as already stated: λ>2βγ\lambda>\frac{2\beta}{\gamma}.

The second case is whenever the conditions of the first case don’t hold, in which case we must prove ft,j=1f_{t,j}=\mathbf{1}. Thus, we either have ct−1c_{t-1} not jj-top, in which case,

or we have w_{t}\not={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}, in which case,

The equation defining it,ji_{t,j} is as follows:

We define Wi,j,j′W_{i,j,j^{\prime}} as follows:

once again relying only on the encoding-validity property.

Since all x_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}} are orthogonal and unit norm, this gives the following intermediate value:

We now prove that the parameters (Wi,Ui,bi)(W_{i},U_{i},b_{i}), when plugged into fθf_{\theta}, implement Equation E.18 assuming the induction hypothesis of Lemma 4. Equation E.18 has three cases.

The condition of the first case is that ct−1=0c_{t-1}=\mathbf{0}, j=1j=1, and w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}; we need to show that it,1=1i_{t,1}=\mathbf{1}. Since ct−1=0c_{t-1}=\mathbf{0}, we have that ht−1=0h_{t-1}=\mathbf{0}. We can compute it,1i_{t,1} as follows:

note the bias term −0.5λγ-0.5\lambda\gamma is only set that way for j=1j=1.

The condition of the second case is that ct−1c_{t-1} is (j−1)(j-1)-top, and w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}; we need to show that it,j=1i_{t,j}=\mathbf{1}. From the induction hypothesis, we have that h_{t-1,j-1}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i^{\prime}}) for some i′i^{\prime}, and ht−1h_{t-1} is elsewhere. We calculate it,ji_{t,j} as follows,

where we do not have to consider the case it,1i_{t,1} since ct−1c_{t-1} cannot be -top, so this case cannot hold.

Finally, in the third case, none of the above conditions hold, and we need to prove it,j=0i_{t,j}=\mathbf{0}. There are a number of possibilities here to enumerate. First, let ct−1=0c_{t-1}=\mathbf{0} and j=1j=1, but w_{t}\not={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}. Then,

Next, we let ct−1=0c_{t-1}=\mathbf{0} and w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, but j>1j>1. Then,

Next, we let j=1j=1, w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, but ct−1≠0c_{t-1}\not=\mathbf{0}. Thus ct−1c_{t-1} is j′j^{\prime}-top for some j′j^{\prime}.

If the second case of Equation E.18 doesn’t hold, then either ct−1c_{t-1} is not (j−1)(j-1)-top, or w_{t}\not={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}; we need to prove it,j=0i_{t,j}=\mathbf{0}. If ct−1c_{t-1} is not (j−1)(j-1)-top and j>1j>1, then we can upper-bound the value of it,ji_{t,j} as follows:

If ct−1c_{t-1} is (j−1)(j-1)-top but w_{t}\not={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, then the value of it,ji_{t,j} is as follows,

Summary

We’ve now specified all parameters of the LSTM and completed the induction step in our proof of the stack correspondence lemma, Lemma 4; this completes the proof of the lemma.

Appendix F Proving generation in O​(m​k)𝑂𝑚𝑘O(mk) hidden units

For both the Simple RNN and the LSTM, we’ve proven stack correspondence lemmas, Lemmas 3, 4. These guarantee that for all prefixes of strings in Dyck-(kk,mm), we can rely on properties of the hidden states of each model to be perfectly informative about the state of the DFA stack.

The two constructions differ in where they make information accessible; we describe those differences here and then provide a general proof that is agnostic to which construction is used.

Recall that we must show that the ϵ\epsilon-truncated support of fθf_{\theta} (for both Simple RNN and LSTM) is the same set as Dyck-(kk,mm). The token-level probability distribution conditioned on the history is given as follows:

and we denote this distribution pfθp_{f_{\theta}}. We first specify VV by row; a single row exists for each of the 2k+12k+1 words in the vocabulary. We let vw,jv_{w,j} refer to the row of word ww, and the kk columns that will participate in the dot product with the rows of stack slot jj in ht−1h_{t-1}.

We start with the only difference in the softmax matrix between the Simple RNN and the LSTM. For the Simple RNN:

because the top of the stack is guaranteed to be in the first stack slot, and where ζ\zeta is a positive scaling constant we’ll define later. And for the LSTM,

for all jj, since exactly one slot is guaranteed to be non-empty (if the stack is non-empty) and it could be at any of the mm slots. The rest of the softmax construction is common between the Simple RNN and the LSTM:

So that we can swap out VV when we swap out ψ\psi for a more efficient encoding, we state here a properties of VV we rely on (once ψ\psi is fixed):

where uu specifies the stack slot where the construction (Simple RNN or LSTM) stores the element at the top of the stack; for the Simple RNN, u=1u=1, for the LSTM, uu is such that ctc_{t} is uu-top, as defined. This ensures that the softmax matrix correctly distinguishes between the symbol, {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, that is encoded in the top of the stack, from any other symbol. Further, we rely on

The softmax matrix we’ve provided obeys the softmax validity property with respect to ψ\psi since ζei⊤ei=ζ\zeta e_{i}^{\top}e_{i}=\zeta, and ζei⊤ej=0\zeta e_{i}^{\top}e_{j}=0, j≠ij\not=i. Further, −ζ1⊤ei=−ζ-\zeta\mathbf{1}^{\top}e_{i}=-\zeta for all ii.

We’re now prepared to state the final lemma in our proof of Theorems 4, 4.

Let w1:T∈Dm,kw_{1:T}\in D_{m,k}. For t=1,...,tt=1,...,t, let qt=Dm,k(w1:T)q_{t}=D_{m,k}(w_{1:T}). Then for all w∈Σw\in\Sigma,

Intuitively, Lemma 5 states that fθf_{\theta} assigns greater than ϵ\epsilon probability mass in context to tokens wtw_{t} such that the prefix w1:t+1w_{1:t+1} is the prefix of some member of Dm,kD_{m,k}.

We proceed in cases by the state, qq. We’ll show that a lower-bound on the probabilities of allowed symbols is greater than an upper-bound on the probabilities of disallowed symbols. We should note, however, that this is effectively a technicality to ensure our construction is fully constructive – that is, we provide concrete values for each parameter in the model; else, we could simply indicate that as ζ\zeta grows, the probability mass on ww such that δ(q,w)=r\delta(q,w)=r converges to , while the mass on all other symbols converges 1 over the number of such symbols.

Case q=[]𝑞q=[]

First, consider the case that qt−1=[]q_{t-1}=[]. For all {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}, \delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i})=r, and for all other symbols ww, δ(q,w)≠r\delta(q,w)\not=r. By the stack correspondence lemmas, Lemma 4 3, we have that Q(ht−1)=qt−1\mathcal{Q}(h_{t-1})=q_{t-1} (Simple RNN) or Q(ct−1)=qt−1\mathcal{Q}(c_{t-1})=q_{t-1} (LSTM), and ht−1=0h_{t-1}=\mathbf{0}.

For all {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}, we have \delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i})=r, and we have logits:

as required. For all w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, we have \delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})\not=r, and we have logits:

Finally, for ω\omega, δ(q,ω)≠r\delta(q,\omega)\not=r, and we have logits:

as required. With the logits specified for our whole vocabulary, we can compute the partition function of the softmax function by summing over all kk open brackets, all kk close brackets, and the END bracket, to determine probabilities for each token,

while the probability for any {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} and ω\omega is p_{f_{\theta}}(w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}|h_{t-1})=e^{0.5\zeta\gamma}/Z_{[]}. The partition function is computed as e0.5ζγe^{0.5\zeta\gamma} multiplied by the number of words ww such that δ(q,w)≠r\delta(q,w)\not=r (in this case, kk open brackets and the end bracket), plus a quantity upper bounded by e−0.5ζγe^{-0.5\zeta\gamma} multiplied by the number of words ww such that δ(q,w)=r\delta(q,w)=r (in this case, kk close brackets.) All probabilities under this model are of this form.

Next, consider the case that q=[{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}},...,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}] for 0<m′<m0<m^{\prime}<m. For the Simple RNN, we have by Lemma 3 that h_{t-1,1}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{m}^{\prime}) (the top of the stack.) Likewise for the LSTM, we have by Lemma 3, that ht−1,m′=tanh(eim′)h_{t-1,m^{\prime}}=\text{tanh}(e_{i_{m^{\prime}}}), and ht−1,j≠m′=0h_{t-1,j\not=m^{\prime}}=\mathbf{0}. Either way, we have v_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}}^{\top}h_{t-1}\geq\zeta\gamma by the softmax validity property.In particular, for the Simple RNN, this expression is equal to ζ\zeta due to the softmax validity property; because of the tanh in ht=ot⊙tanh(ct)h_{t}=o_{t}\odot\text{tanh}(c_{t}) in the LSTM expression, it’s equal to ζγ<ζ\zeta\gamma<\zeta where γ=tanh(1)\gamma=\text{tanh}(1).

As with the first case, we construct the logits for each of the symbols. First, we have \delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{m^{\prime}})\not=r; this is the top element of the stack, which can be popped with a close bracket. We have logits:

as required. For w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i} where i≠im′i\not=i_{m^{\prime}}, we have \delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})=r (since {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} is not the top of the stack, seeing close bracket {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i} would break the well-balancing requirement), and the logits:

again by the softmax validity property. For w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, we have \delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})\not=r, and the logits:

as required, since m′<mm^{\prime}<m, and by the softmax validity property, v_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}}^{\top}h_{t,m^{\prime}<m}=\mathbf{0}. Finally, for ω\omega, we have δ(q,ω)=r\delta(q,\omega)=r, and the logits

by the softmax validity property, as required. To reason about the probabilities, we again need to construct the partition function of the softmax

where ZpZ_{p} stands for ZZ-“partial”, for a partially full stack. (kk open brackets are allowed, plus 1 close bracket; k−1k-1 close brackets are disallowed, as well as the ω\omega word.)

Next, consider the case that q=[{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}},...,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m}}], that is, a full stack with mm elements. For the Simple RNN, by Lemma 3, we have that h_{t-1,1}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}), and h_{t-1,m}=\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{j}) for some jj. (The top of the stack is {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, and the stack is full.) And for the LSTM, by Lemma 4, we have that h_{t-1,m}=\text{tanh}(\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m}})), and ht−1,j≠m=0h_{t-1,j\not=m}=\mathbf{0}. The logits for each symbol are as follows. Identical to the last case, we have \delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i_{m}})\not=r, and p_{f_{\theta}}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i_{m}}|h_{t-1})\propto e^{0.5\zeta\gamma} as required; this is the element at the top of the stack. Also identically, all other brackets have \delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i\not=i_{m}})=r, and the values for those symbols are ≤−0.5ζγ\leq-0.5\zeta\gamma, as required. For w_{t}={\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, we have \delta(q,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})=r; this is because with mm elements, the mm-bound means no more elements can be pushed. The logits are:

by the softmax validity property, as required. Finally, for ω\omega, identically to the previous case, we have δ(q,ω)=r\delta(q,\omega)=r, and logits e−0.5ζγe^{-0.5\zeta\gamma}, as required. The partition function is as follows:

where ZfZ_{f} stands for ZZ-“full”, for a full stack.

Case q=[ω]𝑞delimited-[]𝜔q=[\omega]

Last, consider the case that q=[ω]q=[\omega]. In this case, if at timestep tt, then symbol wt=ωw_{t}=\omega, since the only transition to state [ω][\omega] is δ([],ω)=[ω]\delta([],\omega)=[\omega]. Because of this and the definition of the universe of strings as Σ∗ω\Sigma^{*}\omega, fθf_{\theta} need not be defined on strings progressing from q=[ω]q=[\omega]. Intuitively, this is because ω\omega indicates that the string has terminated. Hence this case vacuously holds.

We thus have three partition functions, ZpZ_{p}, Z[]Z_{[]}, and ZfZ_{f}, whose values we only have bounds for. However, we can see that as we scale ζ\zeta large, they converge to Z[]=(k+1)e0.5ζγZ_{[]}=(k+1)e^{0.5\zeta\gamma}, Zf=ke0.5ζγZ_{f}=ke^{0.5\zeta\gamma}, Zp=(k+1)e0.5ζγZ_{p}=(k+1)e^{0.5\zeta\gamma} because the contributions from the disallowed symbols converges to zero. Thus, the probability assigned to the lowest-probability allowed symbol converges to 1k+1\frac{1}{k+1} as ζ\zeta grows large, while the probability assigned to the highest-probability disallowed symbol converges to . So, we choose ϵ=12(k+1)\epsilon=\frac{1}{2(k+1)},and let ζ>2.4γ\zeta>\frac{2.4}{\gamma}, so that e0.5ζγ>10e−0.5ζγe^{0.5\zeta\gamma}>10e^{-0.5\zeta\gamma}. Under this, the smallest probability assigned to any allowed symbol is lower-bounded by e0.5ζγ(k+1)e0.5ζγ+ke−0.5ζγ>1(k+1)+0.1k>ϵ\frac{e^{0.5\zeta\gamma}}{(k+1)e^{0.5\zeta\gamma}+ke^{-0.5\zeta\gamma}}>\frac{1}{(k+1)+0.1k}>\epsilon, and the largest probability assigned to any disallowed symbol is upper-bounded by e−0.5ζγke+0.5ζγ≤0.1k=110k<ϵ\frac{e^{-0.5\zeta\gamma}}{ke^{+0.5\zeta\gamma}}\leq\frac{0.1}{k}=\frac{1}{10k}<\epsilon. This completes the proof of Lemma 5.

F.2 Completing proofs of generation (Theorems 4, 4)

Now that we’ve proved the probability correctness lemma, we’re ready to complete our proof that our Simple RNN and LSTM constructions generate Dyck-(kk,mm).

Recall that we’ve overloaded notation, calling Dm,k⊂σ∗D_{m,k}\subset\sigma^{*} the set of strings defining the language Dyck-(kk,mm). We must show that the ϵ\epsilon-truncated support of fθf_{\theta}, which we’ll call Lfθ\mathcal{L}_{f_{\theta}} is equal to Dm,kD_{m,k}. We show both inclusions.

Let w∈Dm,kw\in D_{m,k}. We’ll show w∈Lfθw\in\mathcal{L}_{f_{\theta}}. For all prefixes w1:tw_{1:t}, t=1…,Tt=1\dots,T, let qt−1=Dm,k(w1:t−1)q_{t-1}=D_{m,k}(w_{1:t-1}), the state of the DFA after consuming all tokens of the prefix except the last. For the Simple RNN, by Lemma 3, we have that Q(ht−1)=qt−1\mathcal{Q}(h_{t-1})=q_{t-1}. For the LSTM, by Lemma 4, we have that Q(st−1,1,…,st−1,m)=qt−1\mathcal{Q}(s_{t-1,1},\dots,s_{t-1,m})=q_{t-1}. Given this, for either construction, by Lemma 5, we have that pfθ(wt∣w1:t−1)>ϵp_{f_{\theta}}(w_{t}|w_{1:t-1})>\epsilon. Since this is true for all t=1,…,Tt=1,\dots,T, we have that ww is in the ϵ\epsilon-truncated support of fθf_{\theta}, and so w∈Lfθw\in\mathcal{L}_{f_{\theta}}.

Let w1:T∈Lfθw_{1:T}\in\mathcal{L}_{f_{\theta}}. We’ll show w1:T∈Dm,kw_{1:T}\in D_{m,k} by proving the contrapositive.

Let w1:T∉Dm,kw_{1:T}\not\in D_{m,k}. We have that wT=ωw_{T}=\omega by definition. From the transition function δ\delta, we know that δ(q,ω)∈{[ω],r}\delta(q,\omega)\in\{[\omega],r\}, that is, ω\omega transitions from any state either to the accept state or the reject state. Because w1:T∉Dm,kw_{1:T}\not\in D_{m,k}, it must be that δ(qT−1,ω)=r\delta(q_{T-1},\omega)=r, that is, qT=rq_{T}=r. We have no guarantee about fθ(wT∣w1:T−1)f_{\theta}(w_{T}|w_{1:T-1}), however, since we don’t know whether wT−1w_{T-1} is a prefix of some string in Dyck-(kk,mm). However, we do know that there must be some t′t^{\prime} such that the first time qt=rq_{t}=r is for t=t′t=t^{\prime}, that is, the first timestep in which a disallowed symbol is seen and Dm,kD_{m,k} transitions to the reject state rr (after which it self-loops in rr by definition.)

Consider then the prefix w1:t′−1w_{1:t^{\prime}-1}. We know that qt′−1≠rq_{t^{\prime}-1}\not=r. So without loss of generality, let q_{t^{\prime}-1}=[{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{1}},\dots,{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i_{m^{\prime}}}]. We then construct a string:

which simply closes all the brackets on the stack represented by qt′q_{t^{\prime}}. By recursive application of δ\delta, we have that qT′−1=[]q_{T^{\prime}-1}=[], and thus qT′=[ω]q_{T^{\prime}}=[\omega], meaning w1:T′∈Dm,kw_{1:T^{\prime}}\in D_{m,k}.

Coming back to our prefix wt:t′−1w_{t:t^{\prime}-1}, we now know that it is a prefix of a string w1:T′∈Dm,kw_{1:T^{\prime}}\in D_{m,k}. Thus, for our Simple RNN, by Lemma 3, we have that Q(ht′−1)=qt′−1\mathcal{Q}(h_{t^{\prime}-1})=q_{t^{\prime}-1}. Likewise for our LSTM, by Lemma 4, we have that Q(ct′−1)=qt′−1\mathcal{Q}(c_{t^{\prime}-1})=q_{t^{\prime}-1}. And by Lemma 5, since δ(qt′−1,wt′)=r\delta(q_{t^{\prime}-1},w_{t^{\prime}})=r, we have pfθ(wt′∣w1:t′−1)<ϵp_{f_{\theta}}(w_{t^{\prime}}|w_{1:t^{\prime}-1})<\epsilon. And so, w1:T∉Lfθw_{1:T}\not\in\mathcal{L}_{f_{\theta}}. This completes the proof of Theorems 4, 4. ∎

𝑚1O(k^{m+1}) As a corrolary of the above, we now formalize our proof that a general DFA construction in RNNs, using O(km+1)O(k^{m+1}) hidden units to simulate the DFA Dk,mD_{k,m} of Dyck-(kk,mm), permits an RNN construction that generates Dyck-(kk,mm). Formally, \thmnaivegeneration*

A general construction of any DFA in an RNN has ∣Q∣∣Σ∣|Q||\Sigma| states, Merrill (2019); Giles et al. (1990). For Dyck-(kk,mm), ∣Q∣∣Σ∣∈O(km+1)|Q||\Sigma|\in O(k^{m+1}). Each state of the DFA q∈Qq\in Q is represented Σ\Sigma times, once for each word in the vocabulary. If qt=δ(qt−1,w)q_{t}=\delta(q_{t-1},w), then the hidden state is hqt,wh_{q_{t},w}, a 1-hot vector, equal to one at an index specified by and unique to (qt−1,w)(q_{t-1},w). By defining the mapping Q(hq,w)=q\mathcal{Q}(h_{q,w})=q, this construction obeys a stack correspondence lemma. Finally, since the state of the RNN specifies the DFA state as a 1-hot encoding, we can define the softmax matrix VV as follows. For each state qq, we can simply set all rows of VV corresponding to state qq to explicitly encode the log-probabilities of probability distributions we just proved in (§ F). ∎

Appendix G Extending to generation in O​(m​log⁡k)𝑂𝑚𝑘O(m\log k) hidden units

In this section, we prove an O(mlog⁡k)O(m\log k) upper bound on the number of hidden units necessary to capture Dyck-(kk,mm) with an RNN, matching the Ω(mlog⁡k)\Omega(m\log k) lower-bound. This is accomplished by defining a new mapping ψ\psi from slots sts_{t} to open brackets {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} such that we can encode kk open brackets using just 3log⁡k−13\log k-1 hidden units while maintaining the stack correspondence lemmas and Dyck-(kk,mm) generation properties. Formally, for the Simple RNN: \thmsimplernnmlogk* Next, for the LSTM: \thmlstmmlogk*

Intuitively, a simple way to encode kk elements in log⁡k\log k space without making use of floating-point precision is to assign each element one of the 2log⁡k2^{\log k} binary configurations of {0,1}log⁡k\{0,1\}^{\log k}. Our construction will build off of this.

Where the semicolon (;)(;) denotes concatenation. This can be efficiently implemented in our simple RNN construction by modifying the UU matrix to encode each ψ−1\psi^{-1}. Each row of the softmax matrix VV is as follows:

where jj is for all j∈[1,…,m]j\in[1,\dots,m], and all slots not specified are equal to 0\mathbf{0}.

Now it suffices to prove the stack correspondence lemma and probability correctness lemmas using ψ∗\psi_{*} and VV.

As noted in (§ \refappendixsecsimplernn)(\S~{}\ref{appendix_sec_simple_rnn}), the only relevant property of the encodings eie_{i} used for symbols {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} in Lemma 3 is that it take on values in {0,1}\{0,1\}. This is also true of ψ∗−1\psi^{-1}_{*}, so the lemma still holds.

Proof of probability correctness lemma, Lemma 5.

It suffices to show that VV obeys the softmax validity property (Definition 7) with respect to ψ∗\psi_{*} to prove Lemma 5 using new encoding ψ∗\psi_{*}. First, we have that

Where v_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{j},1} is specified since the top of the stack is always in slot 11, ensuring the first requirement. Intuitively, this dot product simply counts up the number of bits that agree between the row jj and the symbol ii encoded; if they’re the same, all bits agree; if not, at least 1 must disagree.

Counting hidden units

Since ψ∗\psi_{*} encodes each symbol in d=3⌈log⁡k⌉−1d=3\lceil\log k\rceil-1 space, and the Simple RNN construction constructs a stack in 2md2md space, we have that Simple RNNs can generate Dyck-(kk,mm) in 2m(3⌈log⁡k⌉−1)=6m⌈log⁡k⌉−2m2m(3\lceil\log k\rceil-1)=6m\lceil\log k\rceil-2m space. This proves Theorem 4.

G.2 O​(log⁡k)𝑂𝑘O(\log k) vocabulary encoding in the LSTM

We’ll use a slight variation of the construction we used for the Simple RNN, made possible since the LSTM can encode the value −1-1 due to the hyperbolic tangent. First, for the encoding, we swap the last ⌈log⁡k⌉−1\lceil\log k\rceil-1 values 11 for −1-1:

And negate the corresponding values of the softmax matrix, as follows:

where jj is for all j∈[1,…,m]j\in[1,\dots,m], and all slots not specified are equal to 0\mathbf{0}.

Now it suffices to prove the stack correspondence lemma and probability correctness lemmas using ψ∗\psi_{*} and VV.

It suffices to show that ψ∗\psi_{*} obeys the encoding-validity property (Definition 5). First, for all ii, we have,

as required. This was possible because we could use the −1-1 value in the encoding ψ∗\psi_{*}. Second, we see that \psi_{*}^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i})\in\{0,1\}^{3\lceil\log k\rceil-1}, as required. Third, we have still that ψ∗(0)=0\psi_{*}(\mathbf{0})=\mathbf{0}. Fourth and finally, we have by construction that all encodings \psi_{*}^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}) are distinct and none are equal to zero, by construction, since each is assigned a different bit configuration (and each bit configuration is concatenated to its negation.)

So, the encoding-validity property holds on ψ∗\psi_{*}, and the stack correspondence lemma, Lemma 4 holds.

Proof of probability correctness lemma, Lemma 5.

The proof of the probability correctness lemma holds as an immediate corollary of the proof of the same lemma for the Simple RNN using ψ∗\psi_{*}. In particular, the only difference between the ψ∗\psi_{*} and VV used for the Simple RNN and that used for the LSTM is the swapping of a factor of −1-1 from a span of vv to that of ψ∗\psi_{*}, so we still have

where uu picks out the top of the stack from the stack correspondence lemma. We identically have

again from the proof for the Simple RNN, and likewise for v_{\omega,j}^{\top}\psi^{-1}({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}) for all j∈[m]j\in[m]. Finally, we have v_{{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i},m^{\prime}}=\mathbf{0}, for m′<mm^{\prime}<m, again from the proof of the simple RNN. This proves that ψ∗\psi_{*} obeys the softmax validity properties, so Lemma 5 still holds for our LSTM construction.

Counting hidden units

Since ψ∗\psi_{*} encodes each symbol in d=3⌈log⁡k⌉−1d=3\lceil\log k\rceil-1 space, and the LSTM construction constructs a stack in mdmd space, we have that LSTM can generate Dyck-(kk,mm) in m(3⌈log⁡k⌉−1)=3m⌈log⁡k⌉−mm(3\lceil\log k\rceil-1)=3m\lceil\log k\rceil-m space. This proves Theorem 4.

G.3 Accounting of finite precision.

Appendix H Experiment Details

In this section, we provide detail on our preliminary study on LSTM LMs learning Dyck-(kk,mm) from samples.

We run experiments on Dyck-(kk,mm) for k∈{2,8,32,128}k\in\{2,8,32,128\} and m∈{3,5}m\in\{3,5\}.

As Dyck-(kk,mm) is an (infinite) set, in order to train language models on it, we must define a probability distribution over it. Intuitively, we sample tokens conditioned on the current DFA (stack) state. Depending on the stack state, one or two of the actions {\{push {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}, pop, end}\} are possible. For example, in the empty stack state, push {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} and end are possible. We sample uniformly at random from the possible actions. Conditioned on choosing the push {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} action, we sample uniformly at random from the open brackets {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i}. Conditioned on choosing the pop action, we generate with probability 1 the {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i} that corresponds to the {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\langle}_{i} on the top of the DFA stack (to obey the well-balancing condition. Upon choosing the end action, we generate ω\omega and terminate.

We choose this distribution to provide the following statistical properties. Consider the markov chain with nodes 11 to mm representing ∣qt∣|q_{t}|, the number of symbols on the DFA state stack at timestep tt. (This is a Markov chain because qtq_{t} is sufficient to describe the probability distribution over all suffixes after timestep tt.) From the probability distribution we’ve defined, for any timestep where the stack is neither empty nor full, that is, 0<∣qt∣<m0<|q_{t}|<m there is probability 1/1/ of advancing in the markov chain towards mm, that is ∣qt+1∣=∣qt∣+1|q_{t+1}|=|q_{t}|+1, and probability 1/21/2 of retreating towards , that is, ∣qt+1∣=∣qt∣−1|q_{t+1}|=|q_{t}|-1. This hitting time of state mm from state is the expected number of timesteps it takes for the generation sequence to start at the empty stack state qt=[]q_{t}=[] and end up at a full stack state ∣qt′∣=m|q_{t^{\prime}}|=m. Because of the 1/21/2 probability of advancing or retreating along the markov chain, the hitting time is O(m2)O(m^{2}).

Length conditions.

Suzgun et al. (2018) showed that the choice of formal language training lengths has a significant effect on the generalization properties of trained LSTMs. We choose training lengths carefully, keeping in mind the O(m2)O(m^{2}) hitting time from empty to full stack states, to empirically ensure that in expectation, the longest training sequences traverse from the empty state to a full stack state ∣qt∣=m|q_{t}|=m and back to the empty stack state at least three times. This ensures that models are not able to use simple heuristics at training time, like remembering the first open bracket in the sequence to close the last close bracket. The exact length statistics are provided in Table 2. Length statistics for training and testing sets are shown in Figure 5..

DFA state analysis.

Since Dyck-(kk,mm) is a regular language, it is reasonable to believe it may learn equivalences between strings that result in the same DFA state, but fail to generalize to DFA states not seen during training time. In Table 3, we see that for kk equal to 22 an 88, all DFA states are seen at training time. Equivalently, every possible stack configuration of 22 or 88 brackets, of stack sizes up to 33 or 55, are seen during training time. For k=32k=32, however, only 15%15\% of all possible DFA states are seen at training time, and only 58%58\% of DFA states seen at testing time are also seen at training time. For k=128k=128, the numbers are even more stark, where 0.3%0.3\% of all possible states, and 31%31\% of states seen at testing time are also seen at training time. Thus, the ability of models to generalize to the test set for kk equal to 3232 and 128128 shows that the learned LSTMs are not simply memorizing DFA states from training time.There are over 34 billion possible DFA states for k=128,m=5k=128,m=5. Instead, we speculate that they’re performing stack-like operations, and leave further investigation to future work.

Sample counts

To test sample efficiency of learning, we study four dataset sizes: {2k,20k,200k,2m,20m}\{2k,20k,200k,2m,20m\} tokens for training for each k,mk,m combination. In all training settings, we use identical development and test sets of size 20k20k and 300k300k tokens, respectively. The development set is sampled from the training distribution.

H.2 Models

We use LSTMs, defined as in the main text with a linear readout layer, and implemented in PyTorch Paszke et al. (2019). We set the hidden dimensionality of the LSTM to 3mlog⁡(k)−m3m\log(k)-m, and the input dimensionality to 2k+102k+10.

H.3 Training

We use the default LSTM initialization provided by PyTorch. We train using Adam Kingma and Ba (2014), using a starting learning rate of 0.010.01 for all training sets less than 2m2m tokens. Based on hand hyperparameter optimization, we found that for k=128k=128, at 2m2m tokens it was better to use a starting learning rate of 0.0010.001. For training sets of size 20m20m tokens, we use a starting learning rate of 0.0010.001 for all settings of kk. We use a batch size of 1010 for all experiments. We evaluate perplexity on the development set after every epoch, restarting Adam with a learning rate decayed by 0.50.5 if the development perplexity does not achieve a new minimum. After three consecutive epochs without a new minimum development perplexity, we stop training. We use no explicit regularization.

H.4 Evaluation

We’re interested in evaluating the behavior of the LSTM LMs in hierarchical memory management, that is, in remembering what type of bracket {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i} can come next. The model cannot possibly do better than random at predicting the next open bracket. Other aspects of the language, like whether the string can end, or if the stack is full, can be solved easily with a counter, which LSTMs are known to implement Weiss et al. (2018); Suzgun et al. (2019). We thus evaluate whether, for each observed close bracket, the LM is confident about which close bracket (that is, ii for {\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{i}) it is, since this is deterministic; it must be the close bracket corresponding to the open bracket at the top of the stack of qt−1q_{t-1}. We do this by normalizing the probability assigned by the LM to the correct close bracket by the sum of probabilities assigned to any close bracket:

We evaluate whether the model is confident, when we define as p({\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle}_{j}|{\color[rgb]{0,0.453125,0.8515625}\definecolor[named]{pgfstrokecolor}{rgb}{0,0.453125,0.8515625}\rangle})>0.8.

We train with three independent seeds for each training setting (training set size, kk, and mm), and report the median across seeds of our bracket closing metric in Figure 3.