Transformer Feed-Forward Layers Build Predictions by Promoting Concepts in the Vocabulary Space

Mor Geva, Avi Caciularu, Kevin Ro Wang, Yoav Goldberg

Introduction

How do transformer-based language models (LMs) construct predictions? We study this question through the lens of the feed-forward network (FFN) layers, one of the core components in transformers Vaswani et al. (2017). Recent work showed that these layers play an important role in LMs, acting as memories that encode factual and linguistic knowledge Geva et al. (2021); Da et al. (2021); Meng et al. (2022). In this work, we investigate how outputs from the FFN layers are utilized internally to build predictions.

We begin by making two observations with respect to the representation of a single token in the input, depicted in Fig. 1. First, each FFN layer induces an additive update to the token representation (Fig. 1,A). Second, the token representation across the layers can be translated at any stage to a distribution over the output vocabulary Geva et al. (2021) (Fig. 1,B). We reason that the additive component in the update changes this distribution (§2), namely, FFN layers compute updates that can be interpreted in terms of the output vocabulary.

We then decompose the FFN update (§3), interpreting it as a collection of sub-updates, each corresponding to a column in the second FFN matrix (Fig. 1,C) that scales the token probabilities in the output distribution. Through a series of experiments, we find that (a) sub-update vectors across the entire network often encode a small-set of human-interpretable well-defined concepts, e.g. “breakfast” or “pronouns” (§4, Fig. 1,D), and (b) FFN updates rely primarily on token promotion (rather than elimination), namely, tokens in the top of the output distribution are those pushed strong enough by sub-updates (§5). Overall, these findings allow fine-grained interpretation of the FFN operation, providing better understanding of the prediction construction process in LMs.

Beyond interpretation, our findings also have practical utility. In §6.1, we show how we can intervene in the prediction process, in order to manipulate the output distribution in a direction of our choice. Specifically, we show that increasing the weight of only 10 sub-updates in GPT2 reduces toxicity in its generations by almost 50%. Also, in §6.2, we show that dominant sub-updates provide a useful signal for predicting an early exit point, saving 20% of the computation on average.

In conclusion, we investigate the mechanism in which FFN layers update the inner representations of transformer-based LMs. We propose that the FFN output can be viewed as a collection of updates that promote concrete concepts in the vocabulary space, and that these concepts are often interpretable for humans. Our findings shed light on the prediction construction process in modern LMs, suggesting promising research directions for interpretability, control, and efficiency.

Token Representations as Evolving Distributions Over the Vocabulary

To analyze the FFN updates, we read from the representation at any layer a distribution over the output vocabulary, by applying the same projection as in Eq. 1 Geva et al. (2021):

The FFN Output as a Collection of Updates to the Output Distribution

We now decompose the FFN output, and interpret it as a set of sub-updates in the vocabulary space.

Therefore, a FFN update can be viewed as a collection of sub-updates, each corresponding to a weighted value vector in the FFN output.

Terminology.

Interpreting Sub-Updates in the Vocabulary Space.

where ew\mathbf{e}_{w} is the embedding of ww, and Z\big{(}\cdot\big{)} is the constant softmax normalization factor.

In the next sections, we use these observations to answer two research questions of (a) What information is encoded in sub-updates and what tokens do they promote? (§4) and (b) How do FFN updates build the output probability distribution? (§5)

Sub-Updates Encode Concepts in the Vocabulary Space

We let experts (NLP graduate students) annotate concepts by identifying common patterns among the top-30 scoring tokens of each value vector. For a set of tokens, the annotation protocol includes three steps of: (a) Identifying patterns that occur in at least 4 tokens, (b) describing each recognized pattern, and (c) classifying each pattern as either “semantic” (e.g., mammals), “syntactic” (e.g., past-tense verbs), or “names”. The last class was added only for WikiLM (see below), following the observation that a large portion of the model’s vocabulary consists of names. Further details, including the complete instructions and a fully annotated example can be found in App. A.2.

Models.

We conduct our experiments over two auto-regressive decoder LMs: The model of Baevski and Auli (2019) (dubbed WikiLM), a 16-layer LM trained on the WikiText-103 corpus Merity et al. (2017) with word-level tokenization (∣V∣=267,744|\mathcal{V}|=267,744), and GPT2 Radford et al. (2019), a 12-layer LM trained on WebText Radford et al. (2019) with sub-word tokenization (∣V∣=50,257|\mathcal{V}|=50,257). GPT2 uses the GeLU activation function Hendrycks and Gimpel (2016), while WikiLM uses ReLU, and in contrast to GPT2, WikiLM does not apply layer normalization after FFN updates. WikiLM defines d=1024,dm=4096d=1024,d_{m}=4096 and GPT2 defines d=768,dm=3072d=768,d_{m}=3072 , resulting in a total of 65k65k and 36k36k value vectors, respectively. For our experiments, we sample 10 random vectors per layer from each model, yielding a total of 160 and 120 vectors to analyze from WikiLM and GPT2, respectively.

1 Projection of Sub-Updates is Meaningful

We validate our approach by comparing concepts in top-tokens of value vectors and 10 random vectors from a normal distribution with the empirical mean and standard deviation of the real vectors. We observe that a substantially higher portion of top-tokens were associated to a concept in value vectors compared to the random ones (Tab. 2): 55.1%55.1\% vs. 22.7%22.7\% in WikiLM, and 37%37\% vs. 16%16\% in GPT2. Also, in both models, the average number of concepts per vector was >1>1 in the value vectors compared to ∼0.5\sim 0.5 in the random ones. Notably, no semantic nor syntactic concepts were identified in WikiLM’s random vectors, and in GPT2, only 4%4\% of the tokens were marked as semantic concepts in the random vectors versus 24.9%24.9\% in the value vectors.

Updates vs. Sub-Updates.

We justify the FFN output decomposition by analyzing concepts in the top-tokens of 10 random FFN outputs per layer (Tab. 2). In WikiLM (GPT2), 39.4%39.4\% (46%46\%) of the tokens were associated with concepts, but for 19.7%19.7\% (34.2%34.2\%) the concept was “stopwords/punctuation”. Also, we observe very few concepts (<4%<4\%) in the last two layers of WikiLM. We account this to extreme sub-updates that dominate the layer’s output (§5.2). Excluding these concepts results in a considerably lower token coverage in projections of updates compared to those of sub-updates: 19.7%19.7\% vs. 55.1%55.1\% in WikiLM, and 11.8%11.8\% vs. 36.7%36.7\% in GPT2.

Overall, this shows that projecting sub-updates to the vocabulary provides a meaningful interface to the information they encode. Moreover, decomposing the FFN outputs is necessary for fine-grained interpretation of sub-updates.

2 Sub-Update Projections are Interpretable

Fig. 2 shows a breakdown of the annotations across layers, for WikiLM and GPT2. In both models and across all layers, a substantial portion (40%-70% in WikiLM and 20%-65% in GPT2) of the top-tokens were associated with well-defined concepts, most of which were classified as “semantic”. Also, we observe that the top-tokens of a single value vector were associated with 1.51.5 (WikiLM) and 1.11.1 (GPT2) concepts on average, showing that sub-updates across all layers encode a small-set of well-defined concepts. Examples are in Tab. 1.

These findings expand on previous results by Geva et al. (2021), who observed that value vectors in the upper layers represent next-token distributions that follow specific patterns. Our results, which hold across all the layers, suggest that these vectors represent general concepts rather than prioritizing specific tokens.

In practice, we find that this task is hard for humans,A sub-update annotation took 8.58.5 minutes on average. as it requires reasoning over a set of tokens without any context, while tokens often correspond to uncommon words, homonyms, or sub-words. Moreover, some patterns necessitate world knowledge (e.g. “villages in Europe near rivers”) or linguistic background (e.g. negative polarity items). This often leads to undetectable patterns, suggesting that the overall results are an underestimation of the true concept frequency. Providing additional context and token-related information are possible future directions for improving the annotation protocol.

Implication for Controlled Generation.

If sub-updates indeed encode concepts, then we can not only interpret their contribution to the prediction, but also intervene in this process, by increasing the weights of value vectors that promote tendencies of our choice. We demonstrate this in §6.1.

FFN Updates Promote Tokens in the Output Distribution

We showed that sub-updates often encode interpretable concepts (§4), but how do these concepts construct the output distribution? In this section, we show that sub-updates systematically configure the prediction via promotion of candidate tokens.

For the experiments, we use a random sample of 2000 examples from the validation set of WikiText-103,Data is segmented into sentences Geva et al. (2021). which both WikiLM and GPT2 did not observe during training. As the experiments do not involve human annotations, we use a larger GPT2 model with L=24,d=1024,dm=4096L=24,d=1024,d_{m}=4096.

We start by comparing the sub-updates’ scores to a reference token in two types of events:

We compute the mean, maximum, and minimum scores of the reference token by the 10 most dominant sub-updates in each event, and average over all the events. As a baseline, we compute the scores by 10 random sub-updates from the same layer.

Tab. 4 shows the results. In both models, tokens promoted to the top of the distribution receive higher maximum scores than tokens eliminated from the top position (1.2→0.51.2\rightarrow 0.5 in WikiLM and 8.5→4.08.5\rightarrow 4.0 in GPT2), indicating they are pushed strongly by a few dominant sub-updates. Moreover, tokens eliminated from the top of the distribution receive near-zero mean scores, by both dominant and random sub-updates, suggesting they are not being eliminated directly. In contrast to promoted tokens, where the maximum scores are substantially higher than the minimal scores (1.21.2 vs. −0.8-0.8 in WikiLM and 8.58.5 vs. −4.9-4.9 in GPT2), for eliminated tokens, the scores are similar in their magnitude (±0.5\pm 0.5 in WikiLM and 4.04.0 vs. −3.6-3.6 in GPT2). Last, scores by random sub-updates are dramatically lower in magnitude, showing that our choice of sub-updates is meaningful and that higher coefficients translate to greater influence on the output distribution.

This suggests that FFN updates work in a promotion mechanism, where top-candidate tokens are those being pushed by dominant sub-updates.

2 Sub-Updates Across Layers

Fig. 3 shows that, in both models, until the last few layers (23-24 in GPT2 and 14-16 in WikiLM), maximum and minimum scores are distributed around non-negative mean scores, with prominent peaks in maximum scores (layers 3-5 in GPT2 and layers 4-11 in WikiLM). This suggests that the token promotion mechanism generally holds across layers. However, scores diverge in the last layers of both models, with strong negative minimum scores, indicating that the probability of the top-candidate is pushed down by dominant sub-updates. We next show that these large deviations in positive and negative scores (Fig. 3, dashed lines) result from the operation of small sets of functional value vectors.

In both models, a small set of homogeneous clusters account for the extreme sub-updates shown in Fig. 3, which can be divided into two main groups of value vectors: Vectors in the upper layers that promote generally unlikely tokens (e.g. rare tokens), and vectors that are spread over all the layers and promote common tokens (e.g. stopwords). These clusters, which cover only a small fraction of the value vectors (1.7% in GPT2 and 1.1% in WikiLM), are mostly active for examples where the input sequence has ≤3\leq 3 tokens or when the target token can be easily inferred from the context (e.g. end-of-sentence period), suggesting that these value vectors might configure “easy” model predictions. More interestingly, the value vectors that promote unlikely tokens can be viewed as “saturation vectors”, which propagate the distribution without changing the top tokens. Indeed, these vectors are in the last layers, where often the model already stores its final prediction Geva et al. (2021).

Applications

We leverage our findings for controlled text generation (§6.1) and computation efficiency (§6.2).

LMs are known to generate toxic, harmful language that damages their usefulness Bender et al. (2021); McGuffie and Newhouse (2020); Wallace et al. (2019). We utilize our findings to create a simple, intuitive method for toxic language suppression.

If LMs indeed operate in a promotion mechanism, we reason that we can decrease toxicity by “turning on” non-toxic sub-updates. We find value vectors that promote safe, harmless concepts by extracting the top-tokens in the projections of all the value vectors and either (a) manually searching for vectors that express a coherent set of positive words (e.g. “safe” and “thank”), or (b) grading the tokens with the Perspective API and selecting non-toxic value vectors (see details in App. A.4). We turn on these value vectors by setting their coefficients to 3, a relatively high value according to Fig. 3. We compare our method with two baselines:

Self-Debiasing (SD) Schick et al. (2021): SD generates a list of undesired words for a given prompt by appending a self-debiasing input, which encourages toxic completions, and calculating which tokens are promoted compared to the original prompt. These undesired words’ probability are then decreased according to a decay constant λ\lambda, which we set to 50 (default).

WordFilter: We prevent GPT2 from generating words from a list of banned words by setting any logits that would result in a banned word completion to −∞-\infty Gehman et al. (2020).

Evaluation.

We evaluate our method on the challenging subset of RealToxicPrompts Gehman et al. (2020), a collection of 1,225 prompts that tend to yield extremely toxic completions in LMs, using the Perspective API, which grades text according to six toxicity attributes. A score of >0.5>0.5 indicates a toxic text w.r.t to the attribute. Additionally, we compute perplexity to account for changes in LM performance. We use GPT2 and, following Schick et al. (2021), generate continuations of 20 tokens.

Results.

Finding the non-toxic sub-updates manually was intuitive and efficient (taking <5<5 minutes). Tab. 5 shows that activation of only 10 value vectors (0.01%) substantially decreases toxicity (↓\downarrow47%), outperforming both SD (↓\downarrow37%) and WordFilter (↓\downarrow20%). Moreover, inducing sub-updates that promote “safety” related concepts is more effective than promoting generally non-toxic sub-updates. However, our method resulted in a perplexity increase greater than this induced by SD, though the increase was still relatively small.

2 Self-Supervised Early Exit Prediction

The recent success of transformer-based LMs in NLP tasks has resulted in major production cost increases Schwartz et al. (2020a), and thus has spurred interest in early-exit methods that reduce the incurred costs Xu et al. (2021). Such methods often use small neural models to determine when to stop the execution process Schwartz et al. (2020b); Elbayad et al. (2020); Hou et al. (2020); Xin et al. (2020, 2021); Li et al. (2021); Schuster et al. (2021).

In this section, we test our hypothesis that dominant FFN sub-updates can signal a saturation event (§5.2), to create a simple and effective early exiting method that does not involve any external model training. For the experiments, we use WikiLM, where saturation events occur across all layers (statistics for WikiLM and GPT2 are in App. A.5).

Baselines.

Evaluation.

Each method is evaluated by accuracy, i.e., the portion of examples for which exiting at the predicted layer yields the final model prediction, and by computation efficiency, measured by the amount of saved layers for examples with correct prediction. We run each method with five random seeds and report the average scores.

Results.

Tab. 6 shows that our method obtains a high accuracy of 94.1%, while saving 20% of computation on average without changing the prediction. Moreover, just by observing the dominant FFN sub-updates, it performs on-par with the prediction rules relying on the representation and FFN output vectors. This demonstrates the utility of sub-updates for predicting saturation events, and further supports our hypothesis that FFN updates play a functional role in the prediction (§5.2).

Related Work

The lack of interpretability of modern LMs has led to a wide interest in understanding their prediction construction process. Previous works mostly focused on analyzing the evolution of hidden representations across layers Voita et al. (2019), and probing the model with target tasks Yang et al. (2020); Clark et al. (2019); Tenney et al. (2019); Saphra and Lopez (2019). In contrast, our approach aims to interpret the model parameters and their utilization in the prediction process.

More recently, a surge of works have investigated the knowledge captured by the FFN layers Da et al. (2021); Jiang et al. (2020); Dai et al. (2022); Yao et al. (2022); Meng et al. (2022); Wallat et al. (2020). These works show that the FFN layers store various types of knowledge, which can be located in specific neurons and edited. Unlike these works, we focus on the FFN outputs and their contribution in the prediction construction process.

Last, our interpretation of FFN outputs as updates to the output distribution relates to recent works that interpreted groups of LM parameters in the discrete vocabulary space Geva et al. (2021); Khashabi et al. (2022), or viewed the representation as an information stream Elhage et al. (2021).

Conclusions

Understanding the inner workings of transformers is valuable for explainability to end-users, for debugging predictions, for eliminating undesirable behavior, and for understanding the strengths and limitations of NLP models. The FFN is an understudied core component of transformer-based LMs, which we focus on in this work.

We study the FFN output as a linear combination of parameter vectors, termed values, and the mechanism by which these vectors update the token representations. We show that value vectors often encode human-interpretable concepts and that these concepts are promoted in the output distribution.

Our analysis of transformer-based LMs provides a more detailed understanding of their internal prediction process, and suggests new research directions for interpretability, control, and efficiency, at the level of individual vectors.

Limitations

Our study focused on the operation of FFN layers in building model predictions. Future work should further analyze the interplay between these layers and other components in the network, such as attention-heads.

In our analysis, we decomposed the computation of FFN layers into smaller units, corresponding to single value vectors. However, it is possible that value vectors are compositional in the sense that combinations of them may produce new meanings. Still, we argue that analyzing individual value vectors is an important first step, since (a) the space of possible combinations is exponential, and (b) our analysis suggests that aggregation of value vectors is less interpretable than individual value vectors (§4.1). Thus, this approach opens new directions for interpreting the contribution of FFN layers to the prediction process in transformer LMs.

In addition, we chose to examine the broad family of decoder-based, auto-regressive LMs, which have been shown to be extremely effective for many NLP tasks, including few- and zero-shot tasks Wang et al. (2022). While these models share the same building blocks of all transformer-based LMs, it will be valuable to ensure that our findings still hold for other models, such as encoder-only LMs (e.g. RoBERTa Liu et al. (2019)) and models trained with different objective functions (e.g. masked language modeling Devlin et al. (2019)).

Finally, our annotation effort was made for the evaluation of our hypothesis that sub-updates encode human-interpretable concepts. Scaling our annotation protocol would enable a more refined map of the concepts, knowledge and structure captured by LMs. Furthermore, since our concept interpretation approach relies on manual inspection of sets of tokens, its success might depend on the model’s tokenization method. In this work, we analyzed models with two different commonly-used tokenizers, and future research could verify our method over other types of tokenizations as well.

Ethics Statement

Our work in understanding the role that single-values play in the inference that transformer-based LMs perform potentially improves their transparency, while also providing useful control applications that save energy (early-exit prediction) and increase model harmlessness (toxic language suppression). It should be made clear that our method for toxic language suppression only reduces the probability of toxic language generation and does not eliminate it. As such, this method (as well as our early-exit method) should not be used in the real world without further work and caution.

More broadly, our work suggests a general approach for modifying LM predictions in particular directions, by changing the weights of FFN sub-updates. While this is useful for mitigating biases, it also has the potential for abuse. It should be made clear that, as in the toxic language suppression application, our approach does not modify the information encoded in LMs, but only changes the intensity in which this information is exposed in the model’s predictions. Moreover, our work primarily proposes an interpretation for FFN sub-updates, which also could be used to identify abusive interventions. Regardless, we stress that LMs should not be integrated into critical systems without caution and monitoring.

Acknowledgements

We thank Shauli Ravfogel, Tal Schuster, and Jonathan Berant for helpful feedback and constructive suggestions. This project has received funding from the Computer Science Scholarship granted by the Séphora Berrebi Foundation, the PBC fellowship for outstanding PhD candidates in Data Science, and the European Research Council (ERC) under the European Union’s Horizon 2020 research and innovation programme, grant agreement No. 802774 (iEXTRACT).

References

Appendix A Appendix

Our interpretation method of sub-updates is based on directly projecting value vectors to the embedding matrix, i.e. for a value v\mathbf{v} and embedding matrix EE, we calculate EvE\mathbf{v} (§4). However, in some LMs like GPT2, value vectors in each layer are added to the token representation followed by a layer normalization (LN) Ba et al. (2016). This raises the question whether “reading” vectors that are normalized in the same manner as the representation would yield different concepts.

A.2 Concepts Annotation

We analyze the concepts encoded in sub-updates, by projecting their corresponding value vectors to the embedding matrix and identifying repeating patterns in the top-scoring 30 tokens (§3). Pattern identification was performed by experts (NLP graduate students), following the instructions presented in Fig. 5. Please note these are the instructions provided for annotations of WikiLM, which uses word-level tokenization. Thus, the terms “words” and “tokens” are equivalent in this case.

For value vectors in WikiLM, which uses a word-level vocabulary with many uncommon words, we additionally attached a short description field for each token that provides context about the meaning of the word. For the description of a token ww, we first try to extract the definition of ww from Wordnet.We use the NLTK python package. If ww does not exist in Wordnet, as often happens for names of people and places, we then search for ww in WikipediaUsing the wptools package https://pypi.org/project/wptools/. and extract a short (possibly noisy) description if the query was successful. A complete annotation example Tab. 7.

A.3 Sub-Update Contribution in FFN Outputs

Empirically, we observe that in some cases sub-updates with negative coefficients do appear as part of the 10 most dominant sub-updates in GPT2. We further attribute this to the success of GeLU in transformer models Shazeer (2020), as it increases the expressiveness of the model by allowing reversing the scores value vectors induce over the vocabulary.

Fig. 6 depicts the contribution of the top-10 dominant sub-updates per layer for WikiLM and GPT2, using 2000 random examples from the WikiText-103 validation set. Clearly, for all the layers, the contribution of the dominant sub-updates exceeds the contribution of random sub-updates. Observe that, even though they cover only 0.24% of the value vectors, the contribution of dominant sub-updates is typically around 5%, and in some layers (e.g. layers 8-16 in WikiLM and layer 1 in GPT2) it reaches over 10% of the total contribution. This demonstrates that analyzing the top-10 dominant sub-updates can shed light on the way predictions are built through the layers.

A.4 Toxic Language Suppression Details

The 10 manually selected value vectors were found by searching for non-toxic words, such as “safe” and “peace”, among the top-30 tokens in the vector projections to the vocabulary. We selected a small set of 10 value vectors whose top-scoring tokens were coherent and seemed to promote different kinds of non-toxic tokens. The list of manually picked vectors is provided in Tab. 8. Importantly, the search process of all vectors was a one-time effort that took <5<5 minutes in total. We chose the value vectors in a greedy-manner, without additional attempts to optimize our choice.

To select 10 non-toxic value vectors based on an automatic toxicity metric, we used the Perspective API. Concretely, we concatenated the top-30 tokens by each value vector and graded the resulting text with the toxicity score produced by the API. Then, we sampled 10 random vectors with a toxicity score <0.1<0.1 (a score of <0.5<0.5 indicates a non-toxic text).

A.5 Early Exit Details

This section provides further details and analysis regarding our early exit method and the baselines we implemented.

Baselines’ Implementation.

We train each binary classifier using 8k training examples, based on the standardized forms of each feature vector. We considered a hyperparameter sweep, using 8-fold cross-validation, with l2l2 or l1l1 regularization (lasso Tibshirani (1996) or ridge Hoerl and Kennard (1970)), regularization coefficients C∈{1e−3,1e−2,1e−1,1,1e1,1e2,1e3}{C\in\{1e^{-3},1e^{-2},1e^{-1},1,1e^{1},1e^{2},1e^{3}\}}, and took the best performing model for each layer. We also used a inversely proportional loss coefficient according to the class frequencies.

In order to achieve high accuracy, we further calibrate a threshold per classifier for reaching the maximal F1 score for each layer. This calibration is done after training each classifier, over a set of 1000 validation examples.

Frequency of Saturation Events.

We investigate the potential of performing early exit for WikiLM and GPT2. Tab. 9 and 10 depict the frequency of saturation events per layer, considering 10k examples from the WikiText-103 validation set, for WikiLM and GPT2, respectively. In GPT2, 34.15% of the examples require the full computation using all the model layers, while for WikiLM, this holds for only 15.22% of the examples. Notably, early fixation events in GPT2 are less common than in WikiLM, possibly due to the larger number of layers the prediction construction is spread over. Hence, we use WikiLM for our experiments, as it has significantly higher computation saving potential, as well as more saturation events per layer.