Star-Transformer

Qipeng Guo, Xipeng Qiu, Pengfei Liu, Yunfan Shao, Xiangyang Xue, Zheng Zhang

Introduction

Recently, the fully-connected attention-based models, like Transformer Vaswani et al. (2017), become popular in natural language processing (NLP) applications, notably machine translation Vaswani et al. (2017) and language modeling Radford et al. (2018). Some recent work also suggest that Transformer can be an alternative to recurrent neural networks (RNNs) and convolutional neural networks (CNNs) in many NLP tasks, such as GPT Radford et al. (2018), BERT Devlin et al. (2018), Transformer-XL Dai et al. (2019) and Universal Transformer Dehghani et al. (2018).

h1<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msub><mimathvariant="bold">h</mi><mn>2</mn></msub></mrow><annotationencoding="application/x−tex">h2</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.8444em;vertical−align:−0.15em;"></span><spanclass="mord"><spanclass="mordmathbf">h</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.3011em;"><spanstyle="top:−2.55em;margin−left:0em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">2</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>h3<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msub><mimathvariant="bold">h</mi><mn>4</mn></msub></mrow><annotationencoding="application/x−tex">h4</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.8444em;vertical−align:−0.15em;"></span><spanclass="mord"><spanclass="mordmathbf">h</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.3011em;"><spanstyle="top:−2.55em;margin−left:0em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">4</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>h5<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msub><mimathvariant="bold">h</mi><mn>6</mn></msub></mrow><annotationencoding="application/x−tex">h6</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.8444em;vertical−align:−0.15em;"></span><spanclass="mord"><spanclass="mordmathbf">h</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.3011em;"><spanstyle="top:−2.55em;margin−left:0em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">6</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>h7\mathbf{h}_{1}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi mathvariant="bold">h</mi><mn>2</mn></msub></mrow><annotation encoding="application/x-tex">\mathbf{h}_{2}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathbf">h</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:0em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">2</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>\mathbf{h}_{3}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi mathvariant="bold">h</mi><mn>4</mn></msub></mrow><annotation encoding="application/x-tex">\mathbf{h}_{4}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathbf">h</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:0em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">4</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>\mathbf{h}_{5}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi mathvariant="bold">h</mi><mn>6</mn></msub></mrow><annotation encoding="application/x-tex">\mathbf{h}_{6}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathbf">h</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:0em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">6</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>\mathbf{h}_{7}h8\mathbf{h}_{8} s<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msub><mimathvariant="bold">h</mi><mn>1</mn></msub></mrow><annotationencoding="application/x−tex">h1</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.8444em;vertical−align:−0.15em;"></span><spanclass="mord"><spanclass="mordmathbf">h</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.3011em;"><spanstyle="top:−2.55em;margin−left:0em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">1</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>h2<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msub><mimathvariant="bold">h</mi><mn>3</mn></msub></mrow><annotationencoding="application/x−tex">h3</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.8444em;vertical−align:−0.15em;"></span><spanclass="mord"><spanclass="mordmathbf">h</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.3011em;"><spanstyle="top:−2.55em;margin−left:0em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">3</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>h4<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msub><mimathvariant="bold">h</mi><mn>5</mn></msub></mrow><annotationencoding="application/x−tex">h5</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.8444em;vertical−align:−0.15em;"></span><spanclass="mord"><spanclass="mordmathbf">h</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.3011em;"><spanstyle="top:−2.55em;margin−left:0em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">5</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>h6<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msub><mimathvariant="bold">h</mi><mn>7</mn></msub></mrow><annotationencoding="application/x−tex">h7</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.8444em;vertical−align:−0.15em;"></span><spanclass="mord"><spanclass="mordmathbf">h</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.3011em;"><spanstyle="top:−2.55em;margin−left:0em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">7</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>h8\mathbf{s}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi mathvariant="bold">h</mi><mn>1</mn></msub></mrow><annotation encoding="application/x-tex">\mathbf{h}_{1}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathbf">h</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:0em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">1</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>\mathbf{h}_{2}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi mathvariant="bold">h</mi><mn>3</mn></msub></mrow><annotation encoding="application/x-tex">\mathbf{h}_{3}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathbf">h</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:0em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">3</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>\mathbf{h}_{4}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi mathvariant="bold">h</mi><mn>5</mn></msub></mrow><annotation encoding="application/x-tex">\mathbf{h}_{5}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathbf">h</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:0em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">5</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>\mathbf{h}_{6}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi mathvariant="bold">h</mi><mn>7</mn></msub></mrow><annotation encoding="application/x-tex">\mathbf{h}_{7}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathbf">h</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:0em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">7</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>\mathbf{h}_{8} Figure 1: Left: Connections of one layer in Transformer, circle nodes indicate the hidden states of input tokens. Right: Connections of one layer in Star-Transformer, the square node is the virtual relay node. Red edges and blue edges are ring and radial connections, respectively. More specifically, there are two limitations of the Transformer. First, the computation and memory overhead of the Transformer are quadratic to the sequence length. This is especially problematic with long sentences. Transformer-XL Dai et al. (2019) provides a solution which achieves the acceleration and performance improvement, but it is specifically designed for the language modeling task. Second, studies indicate that Transformer would fail on many tasks if the training data is limited, unless it is pre-trained on a large corpus. Radford et al. (2018); Devlin et al. (2018).

A key observation is that Transformer does not exploit prior knowledge well. For example, the local compositionality is already a robust inductive bias for modeling the text sequence. However, the Transformer learns this bias from scratch, along with non-local compositionality, thereby increasing the learning cost. The key insight is then whether leveraging strong prior knowledge can help to “lighten up” the architecture.

To address the above limitation, we proposed a new lightweight model named “Star-Transformer”. The core idea is to sparsify the architecture by moving the fully-connected topology into a star-shaped structure. Fig-1 gives an overview. Star-Transformer has two kinds of connections. The radial connections preserve the non-local communication and remove the redundancy in fully-connected network. The ring connections embody the local-compositionality prior, which has the same role as in CNNs/RNNs. The direct outcome of our design is the improvement of both efficiency and learning cost: the computation cost is reduced from quadratic to linear as a function of input sequence length. An inherent advantage is that the ring connections can effectively reduce the burden of the unbias learning of local and non-local compositionality and improve the generalization ability of the model. What remains to be tested is whether one shared relay node is capable of capturing the long-range dependencies.

We evaluate the Star-Transformer on three NLP tasks including Text Classification, Natural Language Inference, and Sequence Labelling. Experimental results show that Star-Transformer outperforms the standard Transformer consistently and has less computation complexity. An additional analysis on a simulation task indicates that Star-Transformer preserve the ability to handle with long-range dependencies which is a crucial feature of the standard Transformer.

In this paper, we claim three contributions as the following and our code is available on Github https://github.com/dmlc/dgl and https://github.com/fastnlp/fastNLP:

Compared to the standard Transformer, Star-Transformer has a lightweight structure but with an approximate ability to model the long-range dependencies. It reduces the number of connections from n2n^{2} to 2n2n, where nn is the sequence length.

The Star-Transformer divides the labor of semantic compositions between the radial and the ring connections. The radial connections focus on the non-local compositions and the ring connections focus on the local composition. Therefore, Star-Transformer works for modestly sized datasets and does not rely on heavy pre-training.

We design a simulation task “Masked Summation” to probe the ability dealing with long-range dependencies. In this task, we verify that both Transformer and Star-Transformer are good at handling long-range dependencies compared to the LSTM and BiLSTM.

Related Work

Recently, neural networks have proved very successful in learning text representation and have achieved state-of-the-art results in many different tasks.

A popular approach is to represent each word as a low-dimensional vector and then learn the local semantic composition functions over the given sentence structures. For example, Kim (2014); Kalchbrenner et al. (2014) used CNNs to capture the semantic representation of sentences, whereas Cho et al. (2014) used RNNs.

These methods are biased for learning local compositional functions and are hard to capture the long-term dependencies in a text sequence. In order to augment the ability to model the non-local compositionality, a class of improved methods utilizes various self-attention mechanisms to aggregate the weighted information of each word, which can be used to get sentence-level representations for classification tasks Yang et al. (2016); Lin et al. (2017); Shen et al. (2018a). Another class of improved methods augments neural networks with a re-reading ability or global state while processing each word Cheng et al. (2016); Zhang et al. (2018).

Modelling Non-Local Compositionality

There are two kinds of methods to model the non-local semantic compositions in a text sequence directly.

One class of models incorporate syntactic tree into the network structure for learning sentence representations Tai et al. (2015); Zhu et al. (2015).

Another type of models learns the dependencies between words based entirely on self-attention without any recurrent or convolutional layers, such as Transformer Vaswani et al. (2017), which has achieved state-of-the-art results on a machine translation task. The success of Transformer has raised a large body of follow-up work. Therefore, some Transformer variations are also proposed, such as GPT Radford et al. (2018), BERT Devlin et al. (2018), Transformer-XL Dai et al. (2019) , Universal Transformer Dehghani et al. (2018) and CN3 Liu et al. (2018a).

However, those Transformer-based methods usually require a large training corpus. When applying them on modestly sized datasets, we need the help of semi-supervised learning and unsupervised pretraining techniques Radford et al. (2018).

Graph Neural Networks

Star-Transformer is also inspired by the recent graph networks Gilmer et al. (2017); Kipf and Welling (2016); Battaglia et al. (2018); Liu et al. (2018b), in which the information fusion progresses via message-passing across the whole graph.

The graph structure of the Star-Transformer is star-shaped by introducing a virtual relay node. The radial and ring connections give a better balance between the local and non-local compositionality. Compared to the previous augmented models Yang et al. (2016); Lin et al. (2017); Shen et al. (2018a); Cheng et al. (2016); Zhang et al. (2018), the implementation of Star-Transform is purely based on the attention mechanism similar to the standard Transformer, which is simpler and well suited for parallel computation.

Due to its better parallel capacity and lower complexity, the Star-Transformer is faster than RNNs or Transformer, especially on modeling long sequences.

Model

The Star-Transformer consists of one relay node and nn satellite nodes. The state of ii-th satellite node represents the features of the ii-th token in a text sequence. The relay node acts as a virtual hub to gather and scatter information from and to all the satellite nodes.

Star-Transformer has a star-shaped structure, with two kinds of connections in the: the radial connections and the ring connections.

For a network of nn satellite nodes, there are nn radial connections. Each connection links a satellite node to the shared relay node. With the radial connections, every two non-adjacent satellite nodes are two-hop neighbors and can receive non-local information with a two-step update.

Ring Connections

Since text input is a sequence, we bake such prior as an inductive bias. Therefore, we connect the adjacent satellite nodes to capture the relationship of local compositions. The first and last nodes are also connected. Thus, all these local connections constitute a ring-shaped structure. Note that the ring connections allow each satellite node to gather information from its neighbors and plays the same role to CNNs or bidirectional RNNs.

With the radial and ring connections, Star-Transformer can capture both the non-local and local compositions simultaneously. Different from the standard Transformer, we make a division of labor, where the radial connections capture non-local compositions, whereas the ring connections attend to local compositions.

2 Implementation

The implementation of the Star-Transformer is very similar to the standard Transformer, in which the information exchange is based on the attention mechanism Vaswani et al. (2017).

where K=HWK,V=HWV\mathbf{K}=\mathbf{H}\mathbf{W}^{K},\mathbf{V}=\mathbf{H}\mathbf{W}^{V}, and WK,WV\mathbf{W}^{K},\mathbf{W}^{V} are learnable parameters.

To gather more useful information from H\mathbf{H}, similar to multi-channels in CNNs, we can use multi-head attention with kk heads.

where ⊕\oplus denotes the concatenation operation, and WiQ,WiK,WiV,WO\mathbf{W}_{i}^{Q},\mathbf{W}_{i}^{K},\mathbf{W}_{i}^{V},\mathbf{W}^{O} are learnable parameters.

Update

We initialize the state with H0=E\mathbf{H}^{0}=\mathbf{E} and s0=average(E)\mathbf{s}^{0}=average(\mathbf{E}).

The update of the Star-Transformer at step tt can be divided into two alternative phases: (1) the update of the satellite nodes and (2) the update of the relay node.

At the first phase, the state of each satellite node hi\mathbf{h}_{i} are updated from its adjacent nodes, including the neighbor nodes hi−1,hi+1\mathbf{h}_{i-1},\mathbf{h}_{i+1} in the sequence, the relay node st\mathbf{s}^{t}, its previous state, and its corresponding token embedding.

where Cit\mathbf{C}^{t}_{i} denotes the context information for the ii-th satellite node. Thus, the update of each satellite node is similar to the recurrent network, except that the update fashion is based on attention mechanism. After the information exchange, a layer normalization operation Ba et al. (2016) is used.

At the second phase, the relay node st\mathbf{s}^{t} summarizes the information of all the satellite nodes and its previous state.

By alternatively updating update the satellite and relay nodes, the Star-Transformer finally captures all the local and non-local compositions for an input text sequence.

Position Embeddings

To incorporate the sequence information, we also add the learnable position embeddings, which are added with the token embeddings at the first layer.

The overall update algorithm of the Star-Transformer is shown in the Alg-1.

3 Output

After TT rounds of update, the final states of HT\mathbf{H}^{T} and sT\mathbf{s}^{T} can be used for various tasks such as sequence labeling and classification. For different tasks, we feed them to different task-specific modules. For classification, we generate the fix-length sentence-level vector representation by applying a max-pooling across the final layer and mixing it with sT\mathbf{s}^{T}, this vector is fed into a Multiple Layer Perceptron (MLP) classifier. For the sequence labeling task, the HT\mathbf{H}^{T} provides features corresponding to all the input tokens.

Comparison to the standard Transformer

Since our goal is making the Transformer lightweight and easy to train with modestly sized dataset, we have removed many connections compared with the standard Transformer (see Fig-1). If the sequence length is nn and the dimension of hidden states is dd, the computation complexity of one layer in the standard Transformer is O(n2d)O(n^{2}d). The Star-Transformer has two phases, the update of ring connections costs O(5nd)O(5nd) (the constant 55 comes from the size of context information C\mathbf{C}), and the update of radial connections costs O(nd)O(nd), so the total cost of one layer in the Star-Transformer is O(6nd)O(6nd).

In theory, Star-Transformer can cover all the possible relationships in the standard Transformer. For example, any relationship hi→hj\mathbf{h}_{i}\rightarrow\mathbf{h}_{j} in the standard Transformer can be simulated by hi→s→hj\mathbf{h}_{i}\rightarrow\mathbf{s}\rightarrow\mathbf{h}_{j}. The experiment on the simulation task in Sec-5.1 provides some evidence to show the virtual node s\mathbf{s} could handle long-range dependencies. Following this aspect, we can give a rough analysis of the path length of dependencies in these models. As discussed in the Transformer paper Vaswani et al. (2017), the maximum dependency path length of RNN and Transformer are O(n)O(n), O(1)O(1), respectively. Star-Transformer can pass the message from one node to another node via the relay node so that the maximum dependency path length is also O(1)O(1), with a constant two comparing to Transformer.

Compare with the standard Transformer, all positions are processed in parallel, pair-wise connections are replaced with a “gather and dispatch” mechanism. As a result, we accelerate the Transformer 10 times on the simulation task and 4.5 times on real tasks. The model also preserves the ability to handle long input sequences. Besides the acceleration, the Star-Transformer achieves significant improvement on some modestly sized datasets.

Experiments

We evaluate Star-Transformer on one simulation task to probe its behavior when challenged with long-range dependency problem, and three real tasks (Text Classification, Natural Language Inference, and Sequence Labelling). All experiments are ran on a NVIDIA Titan X card. Datasets used in this paper are listed in the Tab-1. We use the Adam Kingma and Ba (2014) as our optimizer. On the real task, we set the embedding size to 300 and initialized with GloVe Pennington et al. (2014). And the symbol “Ours + Char” means an additional character-level pre-trained embedding JMT Hashimoto et al. (2017) is used. Therefore, the total size of embedding should be 400 which as a result of the concatenation of GloVe and JMT. We also fix the embedding layer of the Star-Transformer in all experiments.

Since semi- or unsupervised model is also a feasible solution to improve the model in a parallel direction, such as the ELMo Peters et al. (2018) and BERT Devlin et al. (2018), we exclude these models in the comparison and focus on the relevant architectures.

The evaluation metric is the Mean Square Error (MSE), and the generated dataset has (10k/10k/10k) samples in (train/dev/test) sets. The Fig-2 show a case of the masked summation task.

The mask summation task asks the model to recognize the mask value and gather columns in different positions. When the sequence length nn is significantly higher than the number of the columns kk, the model will face the long-range dependencies problem. The Fig-3a shows the performance curves of models on various lengths. Although the task is easy, the performance of LSTM and BiLSTM dropped quickly when the sequence length increased. However, both Transformer and Star-Transformer performed consistently on various lengths. The result indicates the Star-Transformer preserves the ability to deal with the non-local/long-range dependencies.

Besides the performance comparison, we also study the speed with this simulation task since we could ignore the affection of padding, masking, and data processing. We also report the inference time in the Fig-3b, which shows that Transformer is faster than LSTM and BiLSTM a lot, and Star-Transformer is faster than Transformer, especially on the long sequence.

2 Text Classification

Text classification is a basic NLP task, and we select two datasets to observe the performance of our model in different conditions, Stanford Sentiment Treebank(SST) dataset Socher et al. (2013) and MTL-16 Liu et al. (2017) consists of 16 small datasets on various domains. We truncate the sequence which its length higher than 256 to ensure the standard Transformer can run on a single GPU card.

For classification tasks, we use the state of the relay node sT\mathbf{s}^{T} plus the feature of max pooling on satellite nodes max⁡(HT)\max(\mathbf{H}^{T}) as the final representation and feed it into the softmax classifier. The description of hyper-parameters is listed in Tab-1 and Appendix.

Results on SST and MTL-16 datasets are listed in Tab-2,3, respectively. On the SST, the Star-Transformer achieves 2.5 points improvement against the standard Transformer and beat the most models.

Also, on the MTL-16, the Star-Transformer outperform the standard Transformer in all 16 datasets, the improvement of the average accuracy is 4.2. The Star-Transformer also gets better results compared with existing works. As we mentioned in the introduction, the standard Transformer requires large training set to reveal its power. Our experiments show the Star-Transformer could work well on the small dataset which only has 1400 training samples. Results of the time-consuming show the Star-Transformer could be 4.5 times fast than the standard Transformer on average.

3 Natural Language Inference

Natural Language Inference (NLI) asks the model to identify the semantic relationship between a premise sentence and a corresponding hypothesis sentence. In this paper, we use the Stanford Natural Language Inference (SNLI) Bowman et al. (2015) for evaluation. Since we want to study how the model encodes the sentence as a vector representation, we set Star-Transformer as a sentence vector-based model and compared it with sentence vector-based models.

In this experiment, we follow the previous work Bowman et al. (2016) to use concat(r1,r2,∥r1−r2∥,r1−r2)\text{concat}(\mathbf{r}_{1},\mathbf{r}_{2},\|\mathbf{r}_{1}-\mathbf{r}_{2}\|,\mathbf{r}_{1}-\mathbf{r}_{2}) as the classification feature. The r1,r2\mathbf{r}_{1},\mathbf{r}_{2} are representations of premise and hypothesis sentence, it is calculated by sT+max⁡(HT)\mathbf{s}^{T}+\max(\mathbf{H}^{T}) which is same with the classification task. See Appendix for the detail of hyper-parameters.

As shown in Tab-4, the Star-Transformer outperforms most typical baselines (DiSAN, SPINN) and achieves comparable results compared with the state-of-the-art model. Notably, our model beats standard Transformer by a large margin, which is easy to overfit although we have made a careful hyper-parameters’ searching for Transformer.

The SNLI dataset is not a small dataset in NLP area, so improving the generalization ability of the Transformer is a significant topic.

The best result in Tab-4 Yoon et al. (2018) using a large network and fine-tuned hyper-parameters, they get the best result on SNLI but an undistinguished result on SST, see Tab-2.

4 Sequence Labelling

To verify the ability of our model in sequence labeling, we choose two classic sequence labeling tasks: Part-of-Speech (POS) tagging and Named Entity Recognition (NER) task.

Three datasets are used as our benchmark: one POS tagging dataset from Penn Treebank (PTB) Marcus et al. (1993), and two NER datasets from CoNLL2003 Sang and Meulder (2003), CoNLL2012 Pradhan et al. (2012). We use the final state of satellite nodes HT\mathbf{H}^{T} to classify the label in each position. Since we believe that the complex neural network could be an alternative of the CRF, we also report the result without CRF layer.

As shown in Tab-5, Star-Transformer achieves the state-of-the-art performance on sequence labeling tasks. The “Star-Transformer + Char” has already beat most of the competitors. Star-Transformer could achieve such results without CRF, suggesting that the model has enough capability to capture the partial ability of the CRF. The Star-Transformer also outperforms the standard Transformer on sequence labeling tasks with a significant gap.

5 Ablation Study

In this section, we perform an ablation study to test the effectiveness of the radial and ring connections.

We test two variants of our models, the first variants (a) remove the radial connections and only keep the ring connections. Without the radial connections, the maximum path length of this variant becomes O(n)O(n). The second variant (b) removes the ring connections and remains the radial connections. Results in Tab-6 give some insights, the variant (a) loses the ability to handle long-range dependencies, so it performs worse on both the simulation and real tasks. However, the performance drops on SNLI and CoNLL03 is moderate since the remained ring connections still capture the local features. The variant (b) still works on the simulation task since the maximum path length stays unchanged. Without the ring connections, it loses its performance heavily on real tasks. Therefore, both the radial and ring connections are necessary to our model.

Conclusion and Future Works

In this paper, we present Star-Transformer which reduce the computation complexity of the standard Transformer by carefully sparsifying the topology. We compare the standard Transformer with other models on one toy dataset and 21 real datasets and find Star-Transformer outperforms the standard Transformer and achieves comparable results with state-of-the-art models.

This work verifies the ability of Star-Transformer by excluding the factor of unsupervised pre-training. In the future work, we will investigate the ability of Star-Transformer by unsupervised pre-training on the large corpus. Moreover, we also want to introduce more NLP prior knowledge into the model.

Acknowledgments

We would like to thank the anonymous reviewers for their valuable comments. The research work is supported by Shanghai Municipal Science and Technology Commission (No. 17JC1404100 and 16JC1420401), National Key Research and Development Program of China (No. 2017YFB1002104), and National Natural Science Foundation of China (No. 61672162 and 61751201).

References