案例 13:从零实现注意力与小型 Transformer 垃圾短信分类#
学习目标与记号#
本案例不用预训练模型,而是从张量运算开始实现缩放点积注意力、多头注意力、位置编码和 Transformer 编码器,并把它用于短信二分类。本案例的学习目标包括:
说明查询 \(\boldsymbol{Q}\)、键 \(\boldsymbol{K}\) 与值 \(\boldsymbol{V}\) 各自承担的角色;
写出并核验注意力分数、Softmax 权重和加权聚合的维度;
正确构造填充掩码,保证填充词元不参与键和值的聚合,也不进入池化分母;
理解多头注意力为什么先拆分通道、独立计算,再拼接并投影;
说明位置编码、残差连接和 LayerNorm 在编码器中的作用;
在完全相同的数据划分上比较词频-逆文档频率(term frequency-inverse document frequency,TF-IDF)逻辑回归与小型 Transformer;
谨慎解释注意力权重:它描述模型内部的一种加权关系,不天然等于因果贡献或人类理由。
垃圾短信识别涉及误拦截正常通信的代价。本例只展示算法流程,不应直接部署。
数据来源、许可与隐私边界#
使用 UCI SMS Spam Collection,共约 5,574 条英文短信,标签为 ham 或 spam。
说明页:https://archive.ics.uci.edu/dataset/228/sms%2Bspam%2Bcollection
匿名 HTTPS 直链:https://archive.ics.uci.edu/static/public/228/sms+spam+collection.zip
许可:CC BY 4.0,发布改编结果时需要署名并说明修改。
数据可能包含电话号码、网址、金额诱导或令人不适的垃圾内容。探索性数据分析(exploratory data analysis,EDA)只显示经过遮蔽的片段:数字替换为 <NUM>、网址替换为 <URL>,避免在课程网页重复传播联系方式。原始文本只存在用户缓存和内存中。
划分数据后才拟合 TF-IDF 词表和神经网络词表;验证集和测试集中的短信不会决定词元集合或逆文档频率(inverse document frequency,IDF)。短信分类允许读取整条消息,因此编码器采用双向注意力而非因果掩码。若任务改成逐字生成,就必须增加因果掩码,禁止位置 \(t\) 读取未来位置。
from __future__ import annotations
import os
import random
import re
import urllib.request
import zipfile
from collections import Counter
from pathlib import Path
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import accuracy_score, average_precision_score, f1_score
from sklearn.model_selection import train_test_split
SEED = 20260812
FAST_MODE = os.getenv("FAST_MODE", "1") != "0"
random.seed(SEED)
np.random.seed(SEED)
try:
import torch
from torch import nn
from torch.utils.data import DataLoader, Dataset
except ImportError as exc:
raise ImportError("本案例需要 PyTorch 和 scikit-learn。") from exc
torch.manual_seed(SEED)
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CACHE_ROOT = Path(os.getenv("AI_COURSE_DATA_DIR", Path.home() / ".cache" / "ai-course"))
CACHE_DIR = CACHE_ROOT / "uci_sms_spam"
CACHE_DIR.mkdir(parents=True, exist_ok=True)
print({"快速模式": FAST_MODE, "计算设备": str(DEVICE), "缓存目录": str(CACHE_DIR)})
{'快速模式': True, '计算设备': 'cpu', '缓存目录': '/private/tmp/ai-course-case-data/uci_sms_spam'}
URL = "https://archive.ics.uci.edu/static/public/228/sms+spam+collection.zip"
ZIP_PATH = CACHE_DIR / "sms_spam_collection.zip"
if not ZIP_PATH.exists():
urllib.request.urlretrieve(URL, ZIP_PATH)
assert ZIP_PATH.stat().st_size > 100_000
with zipfile.ZipFile(ZIP_PATH) as zf:
members = [n for n in zf.namelist() if Path(n).name == "SMSSpamCollection"]
assert len(members) == 1
text = zf.read(members[0]).decode("utf-8", errors="replace")
rows = []
for line_number, line in enumerate(text.splitlines(), start=1):
if not line.strip():
continue
parts = line.split("\t", 1)
assert len(parts) == 2, f"第 {line_number} 行缺少制表符"
label, message = parts
rows.append((label, message))
df = pd.DataFrame(rows, columns=["label_text", "text"])
df["label"] = df["label_text"].map({"ham": 0, "spam": 1})
assert df["label"].notna().all() and df["text"].str.len().gt(0).all()
assert len(df) >= 5_500
print({"数据维度": df.shape, "正常短信数": int((df["label"] == 0).sum()), "垃圾短信数": int((df["label"] == 1).sum())})
{'数据维度': (5574, 3), '正常短信数': 4827, '垃圾短信数': 747}
def redact(message):
message = re.sub(r"https?://\S+|www\.\S+", "<URL>", message, flags=re.I)
message = re.sub(r"\b\d[\d\s().+-]{4,}\d\b", "<NUM>", message)
return message[:120]
eda = df.assign(
characters=df["text"].str.len(),
words=df["text"].str.findall(r"\b\w+\b").str.len(),
)
eda_summary = eda.groupby("label_text")[["characters", "words"]].agg(["median", "mean", "max"])
eda_summary = eda_summary.rename(index={"ham": "正常短信", "spam": "垃圾短信"})
eda_summary.index.name = "类别"
eda_summary.columns = pd.MultiIndex.from_tuples([
({"characters": "字符数", "words": "词数"}[feature],
{"median": "中位数", "mean": "均值", "max": "最大值"}[stat])
for feature, stat in eda_summary.columns
])
display(eda_summary)
for label in ("ham", "spam"):
example = eda.loc[eda.label_text == label, "text"].iloc[0]
print({"ham": "正常短信", "spam": "垃圾短信"}[label], ":", redact(example))
ax = eda.boxplot(column="characters", by="label_text", showfliers=False)
ax.set_title("消息长度分布(离群点隐藏,仅用于观察主体)")
ax.set_xlabel("类别")
ax.set_ylabel("字符数")
ax.set_xticklabels(["正常短信", "垃圾短信"])
plt.suptitle("")
plt.tight_layout()
| 字符数 | 词数 | |||||
|---|---|---|---|---|---|---|
| 中位数 | 均值 | 最大值 | 中位数 | 均值 | 最大值 | |
| 类别 | ||||||
| 正常短信 | 52.0 | 71.471929 | 910 | 11.0 | 14.780402 | 190 |
| 垃圾短信 | 149.0 | 138.676037 | 223 | 27.0 | 25.483266 | 39 |
正常短信 : Go until jurong point, crazy.. Available only in bugis n great world la e buffet... Cine there got amore wat...
垃圾短信 : Free entry in 2 a wkly comp to win FA Cup final tkts 21st May 2005. Text FA to 87121 to receive entry question(std txt r
train_df, temp_df = train_test_split(
df, test_size=0.30, random_state=SEED, stratify=df["label"]
)
valid_df, test_df = train_test_split(
temp_df, test_size=0.50, random_state=SEED, stratify=temp_df["label"]
)
train_df = train_df.reset_index(drop=True)
valid_df = valid_df.reset_index(drop=True)
test_df = test_df.reset_index(drop=True)
assert set(train_df.index) == set(range(len(train_df))) # 重置后的局部索引连续
assert len(train_df) + len(valid_df) + len(test_df) == len(df)
for part in (train_df, valid_df, test_df):
assert set(part["label"]) == {0, 1}
print({
"训练集": {"样本数": len(train_df), "垃圾短信比例": train_df.label.mean()},
"验证集": {"样本数": len(valid_df), "垃圾短信比例": valid_df.label.mean()},
"测试集": {"样本数": len(test_df), "垃圾短信比例": test_df.label.mean()},
})
print("所有文本变换都将在该单元之后仅用训练集拟合。")
{'训练集': {'样本数': 3901, '垃圾短信比例': np.float64(0.1340681876441938)}, '验证集': {'样本数': 836, '垃圾短信比例': np.float64(0.1339712918660287)}, '测试集': {'样本数': 837, '垃圾短信比例': np.float64(0.13381123058542413)}}
所有文本变换都将在该单元之后仅用训练集拟合。
# 强基线:词/字符 n-gram TF-IDF + 带类别权重的逻辑回归。
# fit 只接收训练文本;transform 才用于验证和测试。
tfidf = TfidfVectorizer(
lowercase=True,
ngram_range=(1, 2),
min_df=2,
max_features=20_000,
sublinear_tf=True,
)
X_train_tfidf = tfidf.fit_transform(train_df["text"])
X_valid_tfidf = tfidf.transform(valid_df["text"])
X_test_tfidf = tfidf.transform(test_df["text"])
baseline = LogisticRegression(
max_iter=1_000, class_weight="balanced", random_state=SEED
)
baseline.fit(X_train_tfidf, train_df["label"])
baseline_prob = baseline.predict_proba(X_test_tfidf)[:, 1]
baseline_pred = (baseline_prob >= 0.5).astype(int)
baseline_metrics = {
"accuracy": accuracy_score(test_df["label"], baseline_pred),
"macro_f1": f1_score(test_df["label"], baseline_pred, average="macro"),
"average_precision": average_precision_score(test_df["label"], baseline_prob),
}
print("TF-IDF 基线:", {"准确率": baseline_metrics["accuracy"], "宏平均 F1": baseline_metrics["macro_f1"], "平均精确率": baseline_metrics["average_precision"]})
assert X_train_tfidf.shape[1] == len(tfidf.vocabulary_)
TF-IDF 基线: {'准确率': 0.982078853046595, '宏平均 F1': 0.9611988639348772, '平均精确率': 0.9787955376403855}
TOKEN_PATTERN = re.compile(r"[A-Za-z]+(?:'[A-Za-z]+)?|\d+|[^\w\s]", re.UNICODE)
def tokenize(message):
return TOKEN_PATTERN.findall(message.lower())
counter = Counter(token for text in train_df["text"] for token in tokenize(text))
vocab_tokens = [token for token, count in counter.most_common(8_000) if count >= 2]
itos = ["<pad>", "<unk>"] + vocab_tokens
stoi = {token: i for i, token in enumerate(itos)}
PAD, UNK = stoi["<pad>"], stoi["<unk>"]
MAX_LEN = 64
def encode(message):
ids = [stoi.get(token, UNK) for token in tokenize(message)][:MAX_LEN]
valid = [1] * len(ids)
ids += [PAD] * (MAX_LEN - len(ids))
valid += [0] * (MAX_LEN - len(valid))
return np.asarray(ids, np.int64), np.asarray(valid, bool)
class SMSDataset(Dataset):
def __init__(self, frame):
encoded = [encode(text) for text in frame["text"]]
self.ids = np.stack([x[0] for x in encoded])
self.mask = np.stack([x[1] for x in encoded])
self.y = frame["label"].to_numpy(np.int64)
def __len__(self):
return len(self.y)
def __getitem__(self, idx):
return (
torch.from_numpy(self.ids[idx]),
torch.from_numpy(self.mask[idx]),
torch.tensor(self.y[idx], dtype=torch.long),
)
train_ds, valid_ds, test_ds = map(SMSDataset, (train_df, valid_df, test_df))
assert train_ds.ids.shape == (len(train_df), MAX_LEN)
assert np.all(train_ds.ids[~train_ds.mask] == PAD)
print("词表大小:", len(itos), ";训练未知词元率:",
np.mean(train_ds.ids[train_ds.mask] == UNK))
词表大小: 3486 ;训练未知词元率: 0.04705046197583511
缩放点积注意力:数值、维度和掩码#
一个注意力头先计算
若批量大小为 \(B\)、头数为 \(H\)、序列长度为 \(T\)、每个头的维度为 \(d_k\),则 \(\boldsymbol{Q},\boldsymbol{K},\boldsymbol{V}\) 的维度均为 \((B,H,T,d_k)\),分数矩阵的维度为 \((B,H,T,T)\)。除以 \(\sqrt{d_k}\) 可以避免点积的方差随 \(d_k\) 快速增大,从而减轻 Softmax 函数过早饱和的问题。
填充掩码 \(\boldsymbol{M}\) 把作为“键”的填充位置得分设为很小的数,使其 Softmax 权重接近 0。若查询本身也是填充位置,还应在后续把相应输出置为 0;最终池化只除以有效词元数。若注意力计算遮住了填充键,但平均池化仍把填充位置计入分母,短信表示就会受到填充长度影响。
多头机制让不同表示子空间分别计算上述关系,拼接后的维度仍为模型维度。注意力得分矩阵的元素数随 \(T^2\) 增长:序列长度从 64 增加到 128 后,元素数约变为原来的 4 倍。
def scaled_dot_product_attention(q, k, v, key_mask=None):
assert q.shape == k.shape == v.shape and q.ndim == 4
scores = q @ k.transpose(-2, -1) / (q.shape[-1] ** 0.5)
if key_mask is not None:
assert key_mask.shape == (q.shape[0], q.shape[2])
scores = scores.masked_fill(~key_mask[:, None, None, :], -1e4)
weights = torch.softmax(scores, dim=-1)
output = weights @ v
return output, weights
q = torch.randn(2, 3, 5, 4)
mask = torch.tensor([[1,1,1,0,0], [1,1,1,1,0]], dtype=torch.bool)
out, weights = scaled_dot_product_attention(q, q, q, mask)
assert out.shape == (2,3,5,4) and weights.shape == (2,3,5,5)
assert torch.allclose(weights.sum(-1), torch.ones(2,3,5), atol=1e-6)
assert weights[0, :, :, 3:].max() < 1e-6
print("注意力维度与填充键权重核验通过。")
注意力维度与填充键权重核验通过。
class MultiHeadSelfAttention(nn.Module):
def __init__(self, d_model=64, n_heads=4):
super().__init__()
assert d_model % n_heads == 0
self.n_heads = n_heads
self.head_dim = d_model // n_heads
self.qkv = nn.Linear(d_model, 3 * d_model)
self.out = nn.Linear(d_model, d_model)
def forward(self, x, valid):
b, t, d = x.shape
qkv = self.qkv(x).reshape(b, t, 3, self.n_heads, self.head_dim)
q, k, v = qkv.unbind(dim=2)
q, k, v = [z.transpose(1, 2) for z in (q, k, v)]
z, weights = scaled_dot_product_attention(q, k, v, valid)
z = z.transpose(1, 2).reshape(b, t, d)
z = self.out(z) * valid.unsqueeze(-1)
return z, weights
class EncoderBlock(nn.Module):
def __init__(self, d_model=64, n_heads=4, dropout=0.1):
super().__init__()
self.attn = MultiHeadSelfAttention(d_model, n_heads)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, 4*d_model), nn.GELU(),
nn.Dropout(dropout), nn.Linear(4*d_model, d_model)
)
self.dropout = nn.Dropout(dropout)
def forward(self, x, valid):
attn_out, weights = self.attn(self.norm1(x), valid)
x = x + self.dropout(attn_out)
x = x + self.dropout(self.ffn(self.norm2(x))) * valid.unsqueeze(-1)
return x, weights
class SMSMiniTransformer(nn.Module):
def __init__(self, vocab_size, max_len=64, d_model=64, n_heads=4):
super().__init__()
self.token = nn.Embedding(vocab_size, d_model, padding_idx=PAD)
self.position = nn.Embedding(max_len, d_model)
self.blocks = nn.ModuleList([
EncoderBlock(d_model, n_heads),
EncoderBlock(d_model, n_heads),
])
self.norm = nn.LayerNorm(d_model)
self.head = nn.Linear(d_model, 2)
def forward(self, ids, valid, return_attention=False):
b, t = ids.shape
assert valid.shape == ids.shape and t <= self.position.num_embeddings
positions = torch.arange(t, device=ids.device)[None, :]
x = (self.token(ids) + self.position(positions)) * valid.unsqueeze(-1)
attentions = []
for block in self.blocks:
x, weights = block(x, valid)
attentions.append(weights)
x = self.norm(x)
denom = valid.sum(1, keepdim=True).clamp_min(1)
pooled = (x * valid.unsqueeze(-1)).sum(1) / denom
logits = self.head(pooled)
return (logits, attentions) if return_attention else logits
probe_ids, probe_mask, _ = next(iter(DataLoader(train_ds, batch_size=4)))
probe_model = SMSMiniTransformer(len(itos))
probe_logits, probe_attention = probe_model(probe_ids, probe_mask, True)
assert probe_logits.shape == (4,2)
assert probe_attention[0].shape == (4,4,MAX_LEN,MAX_LEN)
generator = torch.Generator().manual_seed(SEED)
train_loader = DataLoader(train_ds, batch_size=64, shuffle=True, generator=generator)
valid_loader = DataLoader(valid_ds, batch_size=128, shuffle=False)
test_loader = DataLoader(test_ds, batch_size=128, shuffle=False)
def evaluate_loader(model, loader):
model.eval()
ys, probs = [], []
with torch.no_grad():
for ids, valid, y in loader:
logits = model(ids.to(DEVICE), valid.to(DEVICE))
ys.append(y.numpy())
probs.append(torch.softmax(logits, dim=1)[:,1].cpu().numpy())
return np.concatenate(ys), np.concatenate(probs)
torch.manual_seed(SEED)
model = SMSMiniTransformer(len(itos)).to(DEVICE)
optimizer = torch.optim.AdamW(model.parameters(), lr=2e-3, weight_decay=1e-4)
epochs = 3 if FAST_MODE else 12
history = []
for epoch in range(epochs):
model.train()
losses = []
for ids, valid, y in train_loader:
ids, valid, y = ids.to(DEVICE), valid.to(DEVICE), y.to(DEVICE)
optimizer.zero_grad()
logits = model(ids, valid)
loss = nn.functional.cross_entropy(logits, y)
assert torch.isfinite(loss)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
losses.append(loss.item())
vy, vp = evaluate_loader(model, valid_loader)
history.append((np.mean(losses), average_precision_score(vy, vp)))
print(f"训练轮次={epoch+1},训练损失={history[-1][0]:.4f},验证集平均精确率={history[-1][1]:.4f}")
训练轮次=1,训练损失=0.2258,验证集平均精确率=0.9162
训练轮次=2,训练损失=0.0762,验证集平均精确率=0.9573
训练轮次=3,训练损失=0.0431,验证集平均精确率=0.9576
test_y, transformer_prob = evaluate_loader(model, test_loader)
transformer_pred = (transformer_prob >= 0.5).astype(int)
transformer_metrics = {
"accuracy": accuracy_score(test_y, transformer_pred),
"macro_f1": f1_score(test_y, transformer_pred, average="macro"),
"average_precision": average_precision_score(test_y, transformer_prob),
}
results = pd.DataFrame([
{"model": "tfidf_logistic", **baseline_metrics},
{"model": "mini_transformer", **transformer_metrics},
]).set_index("model")
results_display = results.rename(
index={"tfidf_logistic": "TF-IDF 逻辑回归", "mini_transformer": "小型 Transformer"},
columns={"accuracy": "准确率", "macro_f1": "宏平均 F1", "average_precision": "平均精确率"},
)
results_display.index.name = "模型"
display(results_display)
assert np.isfinite(results.to_numpy()).all()
print("类别不平衡下,平均精确率与宏平均 F1 比只看准确率更有信息;阈值 0.5 也不是部署时唯一选择。")
| 准确率 | 宏平均 F1 | 平均精确率 | |
|---|---|---|---|
| 模型 | |||
| TF-IDF 逻辑回归 | 0.982079 | 0.961199 | 0.978796 |
| 小型 Transformer | 0.980884 | 0.957808 | 0.975508 |
类别不平衡下,平均精确率与宏平均 F1 比只看准确率更有信息;阈值 0.5 也不是部署时唯一选择。
# 定性查看一个已遮蔽样本的最后一层平均注意力。
model.eval()
sample_id = int(np.flatnonzero(test_ds.y == 1)[0])
ids, valid, y = test_ds[sample_id]
with torch.no_grad():
logits, attentions = model(
ids[None].to(DEVICE), valid[None].to(DEVICE), return_attention=True
)
mean_attention = attentions[-1][0].mean(0).cpu().numpy()
valid_count = int(valid.sum())
token_labels = [itos[i] for i in ids[:valid_count].tolist()]
received_attention = mean_attention[:valid_count, :valid_count].mean(axis=0)
top = np.argsort(received_attention)[-8:][::-1]
print("遮蔽文本:", redact(test_df.iloc[sample_id]["text"]))
print("平均收到注意力较高的词元:", [(token_labels[i], float(received_attention[i])) for i in top])
print("这些权重不是因果解释;删词、置换和反事实测试才能进一步检验模型依赖。")
遮蔽文本: Thanks for your ringtone order, reference number X29. Your mobile will be charged 4.50. Should your tone not arrive plea
平均收到注意力较高的词元: [('call', 0.09967180341482162), ('services', 0.0825686976313591), ('not', 0.061526160687208176), ('your', 0.06094611808657646), ('your', 0.05327988043427467), ('be', 0.04985040798783302), ('reference', 0.04409356415271759), ('your', 0.03621407225728035)]
这些权重不是因果解释;删词、置换和反事实测试才能进一步检验模型依赖。
逐步读懂本案例#
从隐私检查到无泄漏划分#
短信原文可能包含电话号码和不适内容,因此展示样本前先遮蔽号码;模型仍使用统一清洗后的文本,不能把测试样本手工改得更容易分类。分层划分保持正常短信与垃圾短信的比例,词表、词频和 TF-IDF 统计量都只能由训练集拟合。确定模型与训练方案后,测试集才用于最终评价。
词袋逻辑回归是重要基线:它计算量较小,常能抓住明显关键词。小型 Transformer 若不能超过基线,仍可用于演示注意力机制,但不能据此宣称复杂结构更有效。比较时应同时报告平均精确率、召回率、F1 分数和预测时间。
词元、填充和掩码的维度#
简单分词器把文本变成词元编号,并保留填充和未知词元等特殊编号。批量编号的维度为 \(B,L\),其中 \(B\) 是批量大小,\(L\) 是统一后的序列长度;填充掩码与其维度相同,真实词元位置为真,填充位置为假。嵌入后的表示为 \(\boldsymbol{X}\in\mathbb{R}^{B\times L\times d}\)。如果掩码方向写反,模型会关注填充位置并忽略正文,因此程序用小样本检查被遮蔽位置的注意力权重接近 0。
查询、键、值与缩放点积#
线性映射得到查询 \(\boldsymbol{Q}\)、键 \(\boldsymbol{K}\) 和值 \(\boldsymbol{V}\)。拆分多头后的维度为 \(B,H,L,d_k\),且 \(d=Hd_k\)。相关得分 \(\boldsymbol{Q}\cdot\boldsymbol{K}^{\mathsf T}\) 的维度是 \(B,H,L,L\):倒数第二个轴对应查询位置,最后一个轴对应键位置。除以 \(\sqrt{d_k}\) 可以避免维度增大时得分过大、Softmax 函数过早饱和。掩码在归一化前把无效键位置设为很小的数。
注意力输出是权重对 \(\boldsymbol{V}\) 的加权和,先得到 \(B,H,L,d_k\),再拼接为 \(B,L,d\)。残差连接要求输入与输出的维度完全相同;Layer Normalization 在特征轴上计算。文本分类可以对有效词元进行带掩码的平均,再送入输出层,从而避免填充长度改变短信表示。
训练和评价模式#
训练模式下,Dropout 会随机丢弃部分数值;评价模式必须关闭这一随机操作,否则同一短信每次得到的概率可能不同。损失函数直接接收输出层的线性运算结果,不应先手动计算 Softmax。训练曲线应分别记录训练损失和验证损失;根据验证集确定停止轮次,不能用测试集选择轮数。
类别不平衡时,准确率容易被多数类主导。平均精确率衡量整个概率排序,F1 分数则依赖具体阈值。若实际应用更重视少漏掉垃圾短信或少误拦正常短信,应在独立验证集上根据代价选择阈值,并在最终测试集上只评价一次。
注意力权重不是完整解释#
权重显示某一层、某一头在当前计算中的信息分配,但不等于某个词对最终预测的因果贡献。残差路径、其他注意力头、后续层和输出层都会改变结果。可以把经过遮蔽的注意力图作为检查模型的线索,但不能据此断言某个词“导致”预测,更不能用于自动惩罚发送者。
与 BERT 案例的关系#
这里从随机初始化开始训练小型编码器,目的是逐项核对查询、键、值、缩放、掩码、多头拼接、位置编码和 Layer Normalization。下一案例的 BERT 使用大规模预训练表示,重点讨论冻结编码器与完整微调。两者不能只按参数量或单次准确率简单比较。
复现与中文呈现#
固定随机种子,并记录词表、最大长度、划分索引、依赖版本和计算设备。快速模式只用于检查程序能否完整运行;正式结论应使用多个随机种子并报告结果的波动范围。所有步骤说明、代码注释、运行提示、图题、坐标轴和展示表头使用中文;官方数据集名称、短信原文、Python API、变量名和 Transformer 缩写保留原文。
结果解释、局限与常见陷阱#
为什么保留强基线。 短短信中,少量关键词、字符模式和网址往往已经具有很强的区分力。小型 Transformer 若不能稳定超过 TF-IDF,并不意味着代码失败;它说明当前数据规模与任务可能不需要更复杂的顺序交互。模型比较必须同时报告计算成本和随机种子波动。
掩码不是可选细节。 若填充词元进入注意力或池化,短消息会比长消息含更多人为的零位置,模型可能学到长度捷径。代码通过权重断言和池化分母明确排除填充。
注意力不等于解释。 权重受层、头、残差和后续非线性共同影响。高权重只说明某一内部计算的相对分配,不能证明该词元导致预测。更可信的诊断应结合删除、替换、梯度、反事实样本及错误类别分析。
数据和部署漂移。 数据来自较早时期,短信语言、短链服务和诈骗策略会变化。现实系统还需要时间外测试、误拦截代价、概率校准、人工复核和反馈回路监控。
隐私与安全。 文本可能携带号码和网址;课程展示应遮蔽,缓存应限制访问。模型输出不能作为执法或个人信誉判断依据。
比较计算量。 TF-IDF 和 Transformer 的建模方式、参数量与训练成本不同。相同的训练轮数不等于相同的参数更新次数或计算时间;正式比较应报告参数更新次数、参数量、固定批量的预测时间、峰值内存及多个随机种子的结果。
综合练习#
手算一个长度为 3、每个词元维度为 2 的单头注意力例子,逐项给出得分、缩放、Softmax 权重和输出,并与函数结果核对。
故意移除填充掩码,构造同一句话但不同填充长度的两个输入,证明错误实现会改变预测。
保持参数量和训练预算近似相同,比较 1、2、4、8 个头;解释“更多头”为什么不保证更好。
加入因果掩码并写断言:位置 (t) 对所有 (j>t) 的权重为零。说明分类任务为什么通常不需要它。
用验证集选择“精确率至少达到某个预设值时,召回率最大”的阈值,然后只在测试集上评价一次。
设计删除高注意力词元和删除随机词元的配对实验;至少运行三个随机种子,检验高注意力词元是否真的更影响预测。
报告注意力矩阵随序列长度的元素数、耗时和峰值内存,验证二次复杂度趋势。