共计 2662 个字符,预计需要花费 7 分钟才能阅读完成。
在自然语言处理领域,BERT 等 Transformer 模型的出现带来了革命性的变化。其中,CLS Token 作为一个看似简单却至关重要的设计,经常被开发者们使用却未必完全理解。本文将带大家从 BERT 原始论文出发,深入探讨 CLS Token 的原理和实际应用。

CLS Token 的设计初衷
BERT 论文(《BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding》)中首次提出了 CLS Token 的概念。这个特殊的标记被添加在每个输入序列的开头,它的设计初衷主要有两个:
- 作为整个序列的 ” 聚合表示 ”,特别适用于分类任务
- 在预训练阶段服务于下一句预测 (NSP) 任务
论文中提到:”The first token of every sequence is always a special classification token ([CLS]). The final hidden state corresponding to this token is used as the aggregate sequence representation for classification tasks.”
CLS Token 与其他 Pooling 方法的对比
在实际应用中,除了使用 CLS Token 外,常见的序列表示方法还有:
- 平均池化(Average Pooling):取所有 token 嵌入的平均值
- 最大池化(Max Pooling):取所有 token 嵌入的最大值
- 首尾拼接:将第一个和最后一个 token 的嵌入拼接起来
相比之下,CLS Token 的优势在于:
- 它是通过自注意力机制动态聚合了整个序列的信息
- 专门针对分类任务进行了优化(通过预训练)
- 计算效率高,只需取第一个位置的输出
不过,对于某些特定任务(如需要保留更多位置信息的任务),其他 Pooling 方法可能更合适。
PyTorch 代码实战
下面我们通过一个完整的 PyTorch 示例,展示如何使用 CLS Token 进行文本分类:
import torch
from transformers import BertModel, BertTokenizer
# 初始化模型和分词器
model_name = 'bert-base-uncased'
tokenizer = BertTokenizer.from_pretrained(model_name)
model = BertModel.from_pretrained(model_name)
# 示例文本
text = "This is a sample text for classification."
# 文本预处理
inputs = tokenizer(text, return_tensors="pt", padding=True, truncation=True)
# 前向传播
with torch.no_grad():
outputs = model(**inputs)
# 获取 CLS Token 对应的隐藏状态
# last_hidden_state 的形状为(batch_size, sequence_length, hidden_size)
last_hidden_state = outputs.last_hidden_state
cls_embedding = last_hidden_state[:, 0, :] # 取第一个 token([CLS])的输出
print(f"CLS Token embedding shape: {cls_embedding.shape}")
print(f"Sample values: {cls_embedding[0, :5]}")
# 在实际分类任务中,我们通常会接一个分类头
class BertClassifier(torch.nn.Module):
def __init__(self, num_classes):
super().__init__()
self.bert = BertModel.from_pretrained(model_name)
self.classifier = torch.nn.Linear(768, num_classes) # 假设是 base 模型,hidden_size=768
def forward(self, input_ids, attention_mask):
outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask)
cls_embedding = outputs.last_hidden_state[:, 0, :]
return self.classifier(cls_embedding)
# 示例使用
classifier = BertClassifier(num_classes=2)
logits = classifier(inputs['input_ids'], inputs['attention_mask'])
print(f"Classification logits: {logits}")
实际应用中的常见问题
- 长文本处理:BERT 的最大序列长度通常为 512。对于更长的文本:
- 可以截断中间部分,保留开头和结尾
- 使用滑动窗口方法,然后聚合多个 CLS 表示
-
考虑使用 Longformer 等支持更长序列的模型
-
微调策略:
- 在领域特定任务上微调时,CLS Token 的表示会适应新任务
- 可以尝试在 CLS 位置添加特殊的任务相关符号
-
对于多任务学习,可以为不同任务使用不同的分类头
-
常见误用:
- 错误地认为 CLS Token 只是第一个词的表示(它实际上聚合了整个序列信息)
- 在非分类任务中盲目使用 CLS Token(如序列标注任务更适合使用每个 token 的输出)
- 忽略注意力掩码对 CLS Token 的影响
性能优化建议
- 批处理:合理设置 batch size,充分利用 GPU 并行计算
- 混合精度训练:使用 torch.cuda.amp 进行自动混合精度训练
- 梯度检查点:对于大模型,可以激活梯度检查点以减少内存占用
- 缓存机制:对于重复输入的静态文本,可以缓存其 CLS 表示
开放性问题
- 在多语言任务中,CLS Token 是否能很好地跨语言传递语义信息?
- 对于极度不平衡的分类任务,CLS Token 是否仍然是最佳选择?
- 在少样本学习场景下,如何更好地利用 CLS Token?
- 对比学习框架下,CLS Token 应该如何设计和优化?
CLS Token 作为 BERT 等 Transformer 模型的核心设计之一,其简单性和有效性在实践中得到了充分验证。希望通过本文的讲解,开发者们能够更深入地理解其工作原理,并在实际项目中做出更合适的技术选型。
