词嵌入与序列模型


文档摘要

词嵌入与序列模型 词嵌入把稀疏的符号化文本压缩进稠密的向量空间,在那里,语义相似性就变成了几何上的邻近。本文件涵盖 Word2Vec(CBOW、Skip-gram)、GloVe、FastText、RNN、LSTM、GRU、带注意力的 seq2seq,以及编码器-解码器范式——从词袋模型走向上下文表示的整条进阶之路。 在第 1 节里,我们引入了分布式假设:出现在相似语境中的词,往往意义也相似。在第 2 节中,我们用 TF-IDF 向量这类手工构造的稀疏特征来表示文本。这些向量生活在非常高维的空间里(每个词表词一维),而且大多是零。词嵌入(word embedding)把这些信息压缩成稠密的低维向量,捕捉语义关系,并且是直接从数据中学出来的。

词嵌入与序列模型

词嵌入把稀疏的符号化文本压缩进稠密的向量空间,在那里,语义相似性就变成了几何上的邻近。本文件涵盖 Word2Vec(CBOW、Skip-gram)、GloVe、FastText、RNN、LSTM、GRU、带注意力的 seq2seq,以及编码器-解码器范式——从词袋模型走向上下文表示的整条进阶之路。

  • 在第 1 节里,我们引入了分布式假设:出现在相似语境中的词,往往意义也相似。在第 2 节中,我们用 TF-IDF 向量这类手工构造的稀疏特征来表示文本。这些向量生活在非常高维的空间里(每个词表词一维),而且大多是零。**词嵌入(word embedding)**把这些信息压缩成稠密的低维向量,捕捉语义关系,并且是直接从数据中学出来的。

  • Word2Vec(Mikolov 等,2013)通过在一个简单的预测任务上训练一个浅层神经网络来学习词嵌入。它有两种架构。

  • **连续词袋模型(Continuous Bag of Words,CBOW)**根据周围的上下文词来预测目标词。给定一窗上下文词(例如 "the cat ___ on the"),模型把它们嵌入向量平均一下,再过一个线性层,去预测缺失的那个词("sat")。训练目标是最大化:

P(w_t \mid w_{t-k}, \ldots, w_{t-1}, w_{t+1}, \ldots, w_{t+k})
  • Skip-gram 模型则反过来:给定目标词,预测周围的上下文词。对目标词 "sat",模型会分别尝试预测 "the"、"cat"、"on"、"the"。它的目标是最大化:
P(w_{t+j} \mid w_t) \quad \text{for each } j \in [-k, k], \; j \neq 0

Skip-gram 和 CBOW 架构并排:CBOW 把上下文嵌入平均来预测中心词,skip-gram 用中心词嵌入来预测每个上下文词

  • Skip-gram 往往对罕见词效果更好,因为每个词会生成多个训练样本(每个上下文位置一个)。CBOW 更快,对常见词略好,因为它在多个上下文信号上做了平均。

  • 在完整词表上训练代价很高,因为 softmax 的分母要在所有 V 个词上求和。**负采样(negative sampling)**通过把问题变成二分类来近似它:区分真正的上下文词(正样本)和随机采样的噪声词(负样本)。模型不再计算完整的 softmax,而是只更新目标词、真正的上下文词以及少数几个负样本的嵌入:

\mathcal{L} = \log \sigma(v_{w_O}^T v_{w_I}) + \sum_{i=1}^{k} \mathbb{E}_{w_i \sim P_n} [\log \sigma(-v_{w_i}^T v_{w_I})]
  • 这里 v_{w_I} 是输入词嵌入,v_{w_O} 是输出(上下文)词嵌入,P_n 是噪声分布,通常取 unigram 频率的 3/4 次方(这会下调像 "the" 这种超高频词的权重)。

  • 为什么这么简单的目标能产生有意义的嵌入?Levy 和 Goldberg(2014)证明,带负采样的 skip-gram 实际上是在隐式地分解一个平移过的点互信息(pointwise mutual information,PMI)矩阵。在收敛时,两个词向量的点积近似于:

v_w^T v_c \approx \text{PMI}(w, c) - \log k
  • 其中 \text{PMI}(w, c) = \log \frac{P(w, c)}{P(w) P(c)} 衡量的是词 wc 共现的频率比偶然期望高出多少(第 5 章信息论),k 是负样本数。共现远高于偶然的词,PMI 高,点积也就高(嵌入相似)。共现低于期望的词,PMI 为负,嵌入也就不相似。这说明 Word2Vec 干的事情和经典的分布式语义方法(比如对共现矩阵做 SVD 的潜在语义分析)是一样的,只是更可扩展、可以在线学习。

  • Word2Vec 嵌入最令人惊讶的性质是,它能通过向量运算来捕捉类比(analogies)。向量 v_{\text{king}} - v_{\text{man}} + v_{\text{woman}} 最接近 v_{\text{queen}}。这之所以行得通,是因为嵌入空间把语义关系编码成近似线性的方向:"royalty(王权)"这个方向大致是 v_{\text{king}} - v_{\text{man}},把它加到 v_{\text{woman}} 上就落在 v_{\text{queen}} 附近。这呼应了第 1 章的线性代数:语义关系就是向量平移。

  • GloVe(Global Vectors for Word Representation,Pennington 等,2014)走的是另一条路。它不再一次只从一个局部上下文窗口学习,而是构建一个全局词共现矩阵 X,其中 X_{ij} 统计在整个语料中词 j 出现在词 i 上下文里的次数。模型然后学习这样的嵌入:让它们的点积近似对数共现次数:

w_i^T \tilde{w}_j + b_i + \tilde{b}_j = \log X_{ij}
  • 损失函数用一个截顶函数 f(X_{ij}) 给每对词加权,防止非常频繁的共现主导训练:
\mathcal{L} = \sum_{i,j=1}^{V} f(X_{ij}) \left(w_i^T \tilde{w}_j + b_i + \tilde{b}_j - \log X_{ij}\right)^2
  • GloVe 把全局矩阵分解(如潜在语义分析)的好处与 Word2Vec 的局部上下文学习结合起来。在实践中,GloVe 和 Word2Vec 产生的嵌入质量相当。

  • FastText(Bojanowski 等,2017)通过把每个词表示成一袋字符 n-gram 来扩展 skip-gram。词 "where" 在 n = 3 时变成:"<wh"、"whe"、"her"、"ere"、"re>",再加上整词词元 ""。该词的嵌入是它所有 n-gram 嵌入之和。

  • 这带来了一个关键优势:FastText 能为训练时从未见过的词生成嵌入。词 "whereabouts" 与 "where" 共享 n-gram,所以即便 "whereabouts" 从未出现在训练数据里,它的嵌入也合情合理。这对于形态丰富的语言(第 1 节)尤其有用,因为这些语言的词有大量屈折变体。

  • **嵌入评测(embedding evaluation)**通常用两类基准。类比任务测试是否 v_a - v_b + v_c \approx v_d(例如 "Paris" - "France" + "Italy" \approx "Rome")。相似度基准把词对之间的余弦相似度(第 1 章)与人类判断做比较。常见的数据集包括 WordSim-353、SimLex-999 和 Google 类比测试集。一个实际提醒:在类比上表现出色的嵌入,未必在情感分类这类下游任务上最好。最好的评测往往是任务本身。

  • 在第 6 章中,我们把 RNN、LSTM、GRU 作为处理序列数据的架构介绍过。这里我们聚焦于它们具体如何应用到语言任务上。

  • 一个语言模型 RNN 一次读一个词元,并在每一步预测下一个词元。隐藏状态 h_t 把整段历史 w_1, \ldots, w_t 压缩进一个固定大小的向量,一个线性层加 softmax 把 h_t 映射到词表上的分布。训练用交叉熵损失对齐真正的下一个词元,这等同于最小化困惑度(第 2 节)。关键局限在于:这个固定大小的隐藏状态必须把历史的一切都编码进去,而早期词元的信息会被逐步覆盖掉。

  • **双向 RNN(bidirectional RNN)**从两个方向处理序列:一个 RNN 从左往右读,另一个从右往左读。在每个位置 t,前向隐藏状态 \overrightarrow{h}_t 和后向隐藏状态 \overleftarrow{h}_t 拼接起来,形成一个上下文感知的表示 h_t = [\overrightarrow{h}_t ; \overleftarrow{h}_t]。这让模型同时能拿到过去和未来的上下文,对 POS 标注和 NER(第 2 节)这类任务很强——在这些任务里,一个词的标签同时依赖于它前后的词。双向 RNN 不能用于语言建模,因为在预测未来词元时你不能偷看它们。

双向 RNN:前向 RNN 从左往右读产生隐藏状态,后向 RNN 从右往左读,每个位置上把两个方向的输出拼接起来

  • **深层堆叠 RNN(deep stacked RNN)**把多个 RNN 层叠在一起。第 l 层在所有时间步的隐藏状态成为第 l + 1 层的输入序列。堆叠 2-4 层通常能通过构建层次化表示来提升性能,就像更深的 CNN 构建特征层级那样(第 6 章)。超过 4 层之后,梯度消失和过拟合就会成为问题,除非在层间加上残差连接。

  • 序列到序列(sequence-to-sequence,seq2seq)架构(Sutskever 等,2014)把一个变长输入序列映射到一个变长输出序列。它由一个编码器(encoder) RNN 和一个解码器(decoder) RNN 组成:编码器读入输入并把它压缩成一个上下文向量(最终隐藏状态),解码器则以这个上下文向量为条件,一次一个词元地生成输出。

Seq2seq 编码器-解码器:编码器 RNN 从左往右读入输入词元,最终隐藏状态作为初始状态传给解码器 RNN,解码器自回归地生成输出词元

  • Seq2seq 是机器翻译的突破性架构。编码器读一句法语,解码器产出英语翻译。解码器从一个特殊的序列起始词元开始,自回归地生成词元,直到产出一个序列结束词元为止。一个实用的小技巧:把输入序列反过来(喂 "chat le" 而不是 "le chat")能改善结果,因为这样让第一个输入词在计算图中更靠近第一个输出词,缩短了梯度路径。

  • 瓶颈问题:整个输入必须被压缩进一个单一的、固定大小的向量。对长句子来说,这个向量无法承载所有信息,性能就退化。这催生了注意力机制(attention)

  • 第 6 章介绍了现代的 Q、K、V 形式的注意力。NLP 中最初的注意力机制形式不同,是作为编码器状态和解码器状态之间的对齐模型来构建的。

  • **Bahdanau 注意力(Bahdanau attention,加性注意力 additive attention,Bahdanau 等,2015)**用一个学习的前馈网络,在解码器隐藏状态 s_t 和每个编码器隐藏状态 h_i 之间计算一个对齐分数:

e_{ti} = v^T \tanh(W_s s_{t-1} + W_h h_i)
  • 分数经 softmax 归一化成注意力权重,上下文向量就是编码器状态的加权和:
\alpha_{ti} = \frac{\exp(e_{ti})}{\sum_j \exp(e_{tj})}, \quad c_t = \sum_i \alpha_{ti} h_i
  • 解码器随后同时使用 s_{t-1}c_t 来产出下一个输出。关键洞见在于:不再为整句用一个固定的上下文向量,解码器的每一步都得到编码器状态的一种不同的加权组合,让模型能够"回看"输入里相关的部分。

  • **Luong 注意力(Luong attention,乘性注意力 multiplicative attention,Luong 等,2015)**简化了分数的计算。**点积(dot)**变体用 e_{ti} = s_t^T h_i。**通用(general)**变体用 e_{ti} = s_t^T W h_i。这些都比 Bahdanau 的加性分数快,因为它们用的是矩阵乘法而不是前馈网络。Luong 注意力还用当前解码器状态 s_t(而不是 s_{t-1})来计算上下文向量,这让它能拿到更多信息,但计算上也略有不同。

源句子与其翻译之间的注意力对齐热力图,展示每个目标词关注的是哪些源词,颜色越亮表示注意力权重越高

  • 注意力权重经常被可视化为热力图,展示解码器在产出每个输出词元时关注的是哪些输入词元。在翻译中,这些热力图大致描画出源语言和目标语言之间的词对齐轨迹,对角线模式会被语序重排打破(比如法语和英语之间形容词-名词的语序差异)。

  • 在推理时,解码器必须在每一步选一个词元。**贪心解码(greedy decoding)**在每个位置都挑概率最高的词元,但这可能导致次优的序列:一个局部的好选择可能把模型逼进一个全局很差的句子里。**束搜索(beam search)**每一步都保留最优的 k 个(束宽)部分序列,把每一个用所有可能的下一个词元扩展,再保留总体最优的 k 个。

  • 当束宽 k = 1 时,束搜索退化为贪心解码。典型取值是 k = 4k = 10。束越大找到的序列越好,但也按比例更慢。束搜索还需要长度归一化,以避免偏向更短的序列——短序列因为相乘的项更少,天然总概率更高。归一化分数是:

\text{score}(y) = \frac{1}{|y|^\alpha} \sum_{t=1}^{|y|} \log P(y_t \mid y_{<t})
  • 其中 |y| 是序列长度,\alpha(通常 0.6-0.7)控制长度惩罚的强度。\alpha = 0 时没有长度归一化。\alpha = 1 时分数就是每个词元的对数概率(几何平均)。中间值在偏向简洁输出和不过早截断之间取得平衡。

  • RNN 是顺序处理文本的,而 1D CNN 则通过在词元序列上滑动卷积核来并行处理。每个卷积核检测一种局部模式(一个 n-gram 特征)。

  • TextCNN(Kim,2014)把多个不同宽度(比如 3、4、5 个词元)的 1D 卷积核作用到输入嵌入矩阵上。每个卷积核产生一张特征图,**按时间取最大池化(max-over-time pooling)**从每张特征图里取出唯一的最大值,捕捉这个模式是否在文本中任何位置被检测到,与位置无关。所有卷积核池化后的特征拼接起来,送进一个分类器。

TextCNN 架构:输入嵌入经过宽度为 3、4、5 的并行卷积核,每个后接按时间最大池化,然后拼接起来送入全连接分类器

  • TextCNN 快,在情感分析这类文本分类任务上出奇地有效。它捕捉局部 n-gram 模式,但建模不了长程依赖:宽度为 5 的卷积核只能看到 5 个连续词元。**膨胀因果卷积(dilated causal convolutions)**通过在卷积核元素之间插入间隔(膨胀)来缓解这一点。把膨胀率按指数增长(1、2、4、8……)堆叠起来,可以在不增加参数的情况下让感受野指数增长,让模型能捕捉跨越数百个词元的依赖。

  • 到目前为止讨论过的所有嵌入(Word2Vec、GloVe、FastText)都为每种词型产生一个单一向量,与上下文无关。"Bank" 无论是指金融机构还是河岸,都拿到同一个嵌入。这是一个根本性的局限,而**上下文嵌入(contextual embeddings)**正是为了解决它。

  • ELMo(Embeddings from Language Models,Peters 等,2018)通过在输入文本上运行一个深双向 LSTM 语言模型来产生上下文词表示。前向 LSTM 在每个位置预测下一个词;一个单独的后向 LSTM 预测上一个词。两者都作为语言模型在大语料上训练。

  • 在每个位置 k,ELMo 用与任务相关的学习权重把所有 L 层的隐藏状态组合起来:

\text{ELMo}_k = \gamma \sum_{j=0}^{L} s_j \, h_{k,j}
  • 这里 h_{k,j} 是位置 k、第 j 层的隐藏状态(第 0 层是原始词元嵌入),s_j 是经 softmax 归一化的标量权重,\gamma 是一个与任务相关的缩放因子。不同层捕捉不同的信息:低层捕捉句法(POS 标签、词的形态),高层捕捉语义(词义、语义角色)。通过用学习到的权重混合所有层,ELMo 嵌入能适配多种下游任务。

  • ELMo 标志着**先预训练再微调(pre-train then fine-tune)**范式的开端:在大量无标注文本上训练一个大语言模型,然后用它的表示来做下游任务。ELMo 具体是把预训练表示作为固定或轻度微调的特征,与任务相关的输入拼接在一起。BERT 和 GPT(第 4 节)则更进一步,端到端地微调整个模型,这被证明效果要好得多。

  • 从 Word2Vec 到 ELMo 的演进,刻画了 NLP 中一个反复出现的主题:从静态走向动态表示、从局部走向全局上下文、从浅层走向深层模型。每一步都是用计算开销换取更丰富的表示。Transformer(第 4 节)用注意力完全取代循环,完成了这一演进,使深层上下文化与并行计算两者兼得。

编程练习(使用 CoLab 或 notebook)

  1. 从零实现带负采样的 Word2Vec skip-gram。在一个小语料上训练,并用 PCA 可视化学到的嵌入。
import jax import jax.numpy as jnp import matplotlib.pyplot as plt # 小型语料 corpus = """the king ruled the kingdom . the queen ruled the kingdom . the prince is the son of the king . the princess is the daughter of the queen . a man worked in the castle . a woman worked in the castle . the king and queen lived in the castle . the prince and princess played outside .""".lower().split() vocab = sorted(set(corpus)) word2idx = {w: i for i, w in enumerate(vocab)} idx2word = {i: w for w, i in word2idx.items()} V = len(vocab) # 用窗口大小 2 生成 skip-gram 对 window = 2 pairs = [] for i, word in enumerate(corpus): for j in range(max(0, i - window), min(len(corpus), i + window + 1)): if i != j: pairs.append((word2idx[word], word2idx[corpus[j]])) pairs = jnp.array(pairs) print(f"Vocabulary: {V} words, Training pairs: {len(pairs)}") # 模型参数 embed_dim = 16 key = jax.random.PRNGKey(42) k1, k2 = jax.random.split(key) W_in = jax.random.normal(k1, (V, embed_dim)) * 0.1 # 输入嵌入 W_out = jax.random.normal(k2, (V, embed_dim)) * 0.1 # 输出嵌入 # 单个对的负采样损失 def neg_sampling_loss(W_in, W_out, target, context, neg_ids): v_in = W_in[target] # (embed_dim,) v_out = W_out[context] # (embed_dim,) v_neg = W_out[neg_ids] # (k, embed_dim) pos_loss = -jax.nn.log_sigmoid(jnp.dot(v_in, v_out)) neg_loss = -jnp.sum(jax.nn.log_sigmoid(-v_neg @ v_in)) return pos_loss + neg_loss # 训练循环 num_neg = 5 lr = 0.05 @jax.jit def train_step(W_in, W_out, target, context, neg_ids): loss, (g_in, g_out) = jax.value_and_grad(neg_sampling_loss, argnums=(0, 1))( W_in, W_out, target, context, neg_ids) return loss, W_in - lr * g_in, W_out - lr * g_out key = jax.random.PRNGKey(0) for epoch in range(50): total_loss = 0.0 for i in range(len(pairs)): key, subkey = jax.random.split(key) neg_ids = jax.random.randint(subkey, (num_neg,), 0, V) loss, W_in, W_out = train_step(W_in, W_out, pairs[i, 0], pairs[i, 1], neg_ids) total_loss += loss if (epoch + 1) % 10 == 0: print(f"Epoch {epoch+1}: avg loss = {total_loss / len(pairs):.4f}") # 用 PCA 可视化(第 1 章) embeddings = W_in mean = embeddings.mean(axis=0) centered = embeddings - mean U, S, Vt = jnp.linalg.svd(centered, full_matrices=False) coords = centered @ Vt[:2].T # 投影到前两个主成分 plt.figure(figsize=(10, 8)) for i, word in idx2word.items(): plt.scatter(coords[i, 0], coords[i, 1], c='#3498db', s=40) plt.annotate(word, (coords[i, 0] + 0.02, coords[i, 1] + 0.02), fontsize=9) plt.title("Word2Vec Skip-gram Embeddings (PCA projection)") plt.grid(alpha=0.3); plt.show()
  1. 构建一个字符级 RNN 语言模型,让它从一小段训练字符串学会生成文本。
import jax import jax.numpy as jnp # 极小的训练文本 text = "to be or not to be that is the question " chars = sorted(set(text)) char2idx = {c: i for i, c in enumerate(chars)} idx2char = {i: c for c, i in char2idx.items()} V = len(chars) data = jnp.array([char2idx[c] for c in text]) # RNN 参数 hidden_dim = 64 key = jax.random.PRNGKey(0) k1, k2, k3, k4, k5 = jax.random.split(key, 5) params = { 'Wx': jax.random.normal(k1, (V, hidden_dim)) * 0.1, 'Wh': jax.random.normal(k2, (hidden_dim, hidden_dim)) * 0.05, 'bh': jnp.zeros(hidden_dim), 'Wy': jax.random.normal(k3, (hidden_dim, V)) * 0.1, 'by': jnp.zeros(V), } def rnn_step(params, h, x_idx): x = jnp.eye(V)[x_idx] # one-hot h = jnp.tanh(x @ params['Wx'] + h @ params['Wh'] + params['bh']) logits = h @ params['Wy'] + params['by'] return h, logits def loss_fn(params, inputs, targets): h = jnp.zeros(hidden_dim) total_loss = 0.0 for t in range(len(inputs)): h, logits = rnn_step(params, h, inputs[t]) log_probs = jax.nn.log_softmax(logits) total_loss -= log_probs[targets[t]] return total_loss / len(inputs) grad_fn = jax.jit(jax.grad(loss_fn)) # 训练 inputs = data[:-1] targets = data[1:] lr = 0.01 for step in range(500): grads = grad_fn(params, inputs, targets) params = {k: params[k] - lr * grads[k] for k in params} if (step + 1) % 100 == 0: l = loss_fn(params, inputs, targets) print(f"Step {step+1}: loss = {l:.4f}") # 生成文本 def generate(params, seed_char, length=60): h = jnp.zeros(hidden_dim) idx = char2idx[seed_char] result = [seed_char] key = jax.random.PRNGKey(42) for _ in range(length): h, logits = rnn_step(params, h, idx) key, subkey = jax.random.split(key) idx = jax.random.categorical(subkey, logits) result.append(idx2char[int(idx)]) return ''.join(result) print(f"\nGenerated: {generate(params, 't')}")
  1. 实现一个带 Bahdanau 注意力的玩具 seq2seq 模型,做序列反转。可视化注意力对齐矩阵。
import jax import jax.numpy as jnp import matplotlib.pyplot as plt # 任务:反转一个数字序列(例如 [3, 1, 4] -> [4, 1, 3]) vocab_size = 10 # 数字 0-9 SOS, EOS = 10, 11 # 特殊词元 total_vocab = 12 embed_dim, hidden_dim = 16, 32 max_len = 5 key = jax.random.PRNGKey(42) keys = jax.random.split(key, 8) params = { 'embed': jax.random.normal(keys[0], (total_vocab, embed_dim)) * 0.1, 'enc_Wx': jax.random.normal(keys[1], (embed_dim, hidden_dim)) * 0.1, 'enc_Wh': jax.random.normal(keys[2], (hidden_dim, hidden_dim)) * 0.05, 'dec_Wx': jax.random.normal(keys[3], (embed_dim, hidden_dim)) * 0.1, 'dec_Wh': jax.random.normal(keys[4], (hidden_dim, hidden_dim)) * 0.05, # Bahdanau 注意力 'Ws': jax.random.normal(keys[5], (hidden_dim, hidden_dim)) * 0.1, 'Wh_att': jax.random.normal(keys[6], (hidden_dim, hidden_dim)) * 0.1, 'v_att': jax.random.normal(keys[7], (hidden_dim,)) * 0.1, # 输出投影(从 hidden + context 到词表) 'Wo': jax.random.normal(keys[0], (hidden_dim * 2, total_vocab)) * 0.1, } def encode(params, seq): """编码输入序列,返回所有隐藏状态。""" h = jnp.zeros(hidden_dim) states = [] for t in range(len(seq)): x = params['embed'][seq[t]] h = jnp.tanh(x @ params['enc_Wx'] + h @ params['enc_Wh']) states.append(h) return jnp.stack(states), h def bahdanau_attention(params, dec_state, enc_states): """计算 Bahdanau 注意力权重和上下文向量。""" scores = jnp.tanh(enc_states @ params['Wh_att'] + dec_state @ params['Ws']) e = scores @ params['v_att'] # (src_len,) alpha = jax.nn.softmax(e) context = alpha @ enc_states return context, alpha def decode_step(params, dec_h, prev_token, enc_states): x = params['embed'][prev_token] dec_h = jnp.tanh(x @ params['dec_Wx'] + dec_h @ params['dec_Wh']) context, alpha = bahdanau_attention(params, dec_h, enc_states) combined = jnp.concatenate([dec_h, context]) logits = combined @ params['Wo'] return dec_h, logits, alpha def seq2seq_loss(params, src, tgt): enc_states, enc_final = encode(params, src) dec_h = enc_final loss = 0.0 prev_token = SOS for t in range(len(tgt)): dec_h, logits, _ = decode_step(params, dec_h, prev_token, enc_states) log_probs = jax.nn.log_softmax(logits) loss -= log_probs[tgt[t]] prev_token = tgt[t] return loss / len(tgt) # 生成训练数据:反转序列 key = jax.random.PRNGKey(0) train_srcs, train_tgts = [], [] for _ in range(200): key, subkey = jax.random.split(key) length = jax.random.randint(subkey, (), 3, max_len + 1) key, subkey = jax.random.split(key) seq = jax.random.randint(subkey, (int(length),), 0, vocab_size) train_srcs.append(seq) train_tgts.append(seq[::-1]) # 反转 # 训练 grad_fn = jax.grad(seq2seq_loss) lr = 0.01 for epoch in range(100): total_loss = 0.0 for src, tgt in zip(train_srcs, train_tgts): grads = grad_fn(params, src, tgt) params = {k: params[k] - lr * grads[k] for k in params} total_loss += seq2seq_loss(params, src, tgt) if (epoch + 1) % 20 == 0: print(f"Epoch {epoch+1}: avg loss = {total_loss / len(train_srcs):.4f}") # 可视化某一个样本的注意力 test_src = jnp.array([3, 1, 4, 1, 5]) test_tgt = test_src[::-1] enc_states, enc_final = encode(params, test_src) dec_h = enc_final attentions = [] prev_token = SOS for t in range(len(test_tgt)): dec_h, logits, alpha = decode_step(params, dec_h, prev_token, enc_states) attentions.append(alpha) prev_token = test_tgt[t] att_matrix = jnp.stack(attentions) fig, ax = plt.subplots(figsize=(6, 5)) im = ax.imshow(att_matrix, cmap='Blues') ax.set_xlabel("Source position"); ax.set_ylabel("Target position") src_labels = [str(int(x)) for x in test_src] tgt_labels = [str(int(x)) for x in test_tgt] ax.set_xticks(range(len(src_labels))); ax.set_xticklabels(src_labels) ax.set_yticks(range(len(tgt_labels))); ax.set_yticklabels(tgt_labels) for i in range(len(tgt_labels)): for j in range(len(src_labels)): ax.text(j, i, f"{att_matrix[i,j]:.2f}", ha='center', va='center', fontsize=9) ax.set_title("Bahdanau Attention Alignment (sequence reversal)") plt.colorbar(im); plt.tight_layout(); plt.show()

作者与出处
原作者: HenryNdubuaku
来源:HenryNdubuaku
许可证:Apache-2.0
整理: 灏天文库整理
由灏天文库结构化整理,提供目录导航、全文检索与在线阅读,便于系统化学习
发布者: 作者: HenryNdubuaku 转发
评论区 (0)
U