Large Transformer Model Inference Optimization
Large transformer models set state-of-the-art results but are expensive to train and use, with extremely high inference costs in time and memory. Lilian Weng's blog post discusses approaches for making transformer inference more efficient. It cites Pope et al. 2022, noting large memory footprint and low parallelizability as two main factors. Parameters and intermediate states, such as the KV cache, must be stored in memory; for batch size 512 and context length 2048, the KV cache totals 3TB, 3x the model size. Attention cost scales quadratically with sequence length, and autoregressive decoding is hard to parallelize.
The post outlines goals: reduce memory footprint, reduce computation complexity (FLOPs), and reduce latency. It covers parallelism, memory offloading, smart batching, network compression techniques (pruning, quantization, distillation), and architecture-specific improvements, especially for attention layers.
Distillation uses a smaller student model to mimic a teacher. DistilBERT reduces BERT parameters by 40% while maintaining 97% of BERT's performance on fine-tuned downstream tasks and running 71% faster, with loss combining soft distillation, masked language modeling, and cosine embedding. Distillation can combine with quantization, pruning, or sparsification.
Quantization approaches include post-training quantization and quantization-aware training. Challenges arise from high dynamic ranges of activations and outliers. Mixed-precision quantization, fine-grained quantization, second-order methods like GPTQ and Q-BERT, outlier smoothing with SmoothQuant, and quantization-aware training with distillation are discussed. Pruning can be unstructured or structured, with a routine workflow of training, pruning, and optional retraining, motivated by the Lottery Ticket Hypothesis.
Based on reporting from the original publisher. Visit the source for full context and later updates.
Publisher excerpt
[Updated on 2023-01-24: add a small section on Distillation .] Large transformer models are mainstream nowadays, creating SoTA results for a variety of tasks. They are powerful but very expensive to train and use. The extremely high inference cost, in both time and memory, is a big bottleneck for adopting a powerful transformer for solving real-world tasks at scale. Why is it hard to run inference for large transformer models? Besides the increasing size of SoTA models, there are two main factors contributing to th