1. 面试官为什么总爱考“手撕”注意力?

如果你最近在准备算法工程师或者深度学习相关的面试,我敢打赌,你肯定在刷题网站上见过“手撕注意力机制”这道题。它几乎成了面试的“保留节目”,从大厂到初创公司,面试官们都乐此不疲。很多朋友第一次看到这个要求时,心里可能会犯嘀咕:“现在框架这么成熟,PyTorch里nn.MultiheadAttention一行代码就搞定,为什么还要我手写?这不是为难人吗?”

我刚开始面试别人的时候,也问过这个问题。后来自己带团队、做项目,才真正明白面试官的良苦用心。他们想看的,绝不仅仅是你“知道”注意力是什么,而是想通过这几十行代码,考察你三个核心能力

  1. 对模型底层原理的透彻理解:你能不能用代码把公式softmax(QK^T/√dk)V清晰地表达出来?你知道QKV是怎么从输入数据变过来的吗?自注意力和交叉注意力在输入上到底有什么区别?这些细节,调用现成API是永远体会不到的。
  2. 扎实的矩阵操作基本功:注意力机制本质上就是一连串的矩阵乘法、转置和缩放。手写代码能立刻暴露你对张量形状变化的掌握程度。比如,Q的形状是(batch, seq_len_q, dim)K转置后是(batch, dim, seq_len_k),它们相乘得到的注意力分数矩阵形状应该是(batch, seq_len_q, seq_len_k)。这个推导过程,面试官会盯着看。
  3. 解决实际问题的工程思维:比如,你怎么处理序列长度不一致的情况(交叉注意力常见)?你怎么实现mask来屏蔽掉填充的padding或者未来的信息(在解码器中)?这些都是在真实模型中必须处理的工程细节。

所以,把“手撕注意力”看成一道单纯的代码题,就把它想简单了。它其实是一扇门,推开它,面试官能看到你对深度学习模型构建的整体认知水平。今天,我就把自己当年准备面试时总结的方法,以及后来在工作中反复使用的经验,掰开揉碎了讲给你听。我们从最简单的单头注意力开始,搞定Self-AttentionCross-Attention,让你下次面试时,能胸有成竹地在白板或编辑器里把它流畅地写出来。

2. 拆解单头注意力:从生活比喻到数学公式

在直接看代码之前,我们先把注意力机制用最“人话”的方式理解一遍。想象一个场景:你在一个嘈杂的咖啡馆里和朋友聊天。虽然周围有音乐声、别人的谈话声、咖啡机的噪音,但你的大脑能自动把注意力“聚焦”在你朋友的声音上,忽略其他杂音。这个过程,就是最朴素的“注意力”。

在深度学习中,我们要让机器学会这种“聚焦”能力。对于一段输入信息(比如一个句子),我们想知道其中每一个部分(比如每一个词)与其他所有部分的关联程度,然后根据这个关联程度,重新组合信息。

2.1 核心五步:把比喻变成计算步骤

我们把上面这个“聚焦”过程,分解成五个可计算的步骤。我会用一个非常简单的例子贯穿始终:假设我们的输入是三个词的嵌入向量,[我, 爱, 编程]

第一步:准备原料——输入与嵌入 输入就是我们的原始数据。对于自注意力,我们只有一个来源。比如,处理句子“我爱编程”时,这三个词的向量就是我们的输入X,形状假设是(3, 4),表示3个词,每个词用4维向量表示。 对于交叉注意力,我们有两个来源。比如在机器翻译中,Q可能来自目标语言(英文)的某个词,而KV来自源语言(中文)的整个句子。这是自注意力和交叉注意力最根本的区别:数据源是一个还是两个

第二步:提出三个关键问题——线性变换得到Q, K, V 这是注意力机制最精妙的设计。我们通过三个不同的可学习权重矩阵W_Q, W_K, W_V,把同样的输入(或不同的输入)投影到三个不同的空间。

  • Query:可以理解为“当前正在关注的词”发出的“询问”。比如,当处理“爱”这个词时,它的Query向量就在问:“我和周围的词(‘我’和‘编程’)有什么关系?”
  • Key:可以理解为所有词(包括自己)提供的“标签”或“身份标识”。它用来回应Query的询问。
  • Value:可以理解为每个词所携带的“实际信息”或“内容”。最终我们根据注意力权重聚合的是Value

用一个不太严谨但形象的类比:你在图书馆(输入序列)找关于“深度学习”(Query)的书。图书馆每本书都有一个索引标签(Key)。你通过比对“深度学习”和每本书的标签,计算出一个匹配分数(注意力分数),然后根据这个分数,去组合这些书里的具体内容(Value),最终得到你需要的知识。

第三步:计算匹配度——注意力分数 这一步计算Query和每一个Key的匹配程度。最常用的方法是点积。Q的每一行(一个Query)会和K的每一行(一个Key)做点积,得到一个分数矩阵。分数越高,表示匹配度越高。 这里有一个关键操作:缩放。我们会把点积结果除以√dk,其中dkKey向量的维度。为什么要缩放?因为当dk很大时,点积的结果可能会非常大,导致经过softmax后梯度变得极小(饱和),不利于模型训练。缩放是为了稳定梯度。

第四步:归一化成权重——Softmax 上一步得到的分数可能有正有负,数值范围也不固定。我们用softmax函数沿着最后一个维度(Key的维度)进行归一化,使得所有Key对于当前这个Query的权重之和为1,变成一个概率分布。这样,权重就代表了“注意力”的分配比例。

第五步:聚合信息——加权求和 最后,我们用第四步得到的注意力权重,对所有的Value向量进行加权求和。对于当前这个Query,权重大的Value对最终结果的贡献就大。这个加权求和的结果,就是自注意力交叉注意力在这个位置上的输出。它已经融合了整个序列中所有相关信息。

把这五步用公式串起来,就是那个经典的注意力公式: Output = softmax(Q * K^T / √dk) * V

理解了这个流程,代码其实就是对这个公式的忠实翻译。下面我们就进入实战环节。

3. 手把手实现:从零搭建单头注意力类

理论懂了,不写代码等于零。我们直接用PyTorch来实现一个灵活的单头注意力类,让它同时支持自注意力和交叉注意力。我会一行一行解释,确保你知其然也知其所以然。

3.1 搭建类骨架与初始化

首先,我们把必要的工具包引进来。torchnn是基础,math用来计算开方。

import torch
from torch import nn
import math

接下来,我们定义OneHeadAttention类。在初始化函数__init__里,我们需要确定几个关键尺寸:

  • emb_size: 输入词向量的维度。
  • qk_size: QueryKey投影后的维度。在标准Transformer中,这个值通常是emb_size,但也可以不同,我们的实现保持灵活。
  • v_size: Value投影后的维度。同样,它可以和qk_size不同。
class OneHeadAttention(nn.Module):
    def __init__(self, emb_size, qk_size, v_size):
        super(OneHeadAttention, self).__init__()
        # 定义三个线性变换层,分别生成Q, K, V
        self.Wq = nn.Linear(emb_size, qk_size) # 将输入投影到Query空间
        self.Wk = nn.Linear(emb_size, qk_size) # 将输入投影到Key空间
        self.Wv = nn.Linear(emb_size, v_size)  # 将输入投影到Value空间
        # 定义softmax,在最后一个维度上做归一化
        self.softmax = nn.Softmax(dim=-1)

这里有个细节值得注意:我们使用了nn.Linear层。它内部包含了可学习的权重矩阵和偏置项。在注意力机制中,偏置项有时会被省略以简化计算,但使用nn.Linear是更通用和标准的做法,模型会自己学习是否需要一个偏置。

3.2 核心前向传播逻辑

forward函数是魔法发生的地方。它的设计要能同时处理自注意力(一个输入)和交叉注意力(两个输入)。一个优雅的做法是:让函数接收两个参数x_qx_kv。当它们是同一个张量时,就是自注意力;当它们是不同张量时,就是交叉注意力。

    def forward(self, x_q, x_kv, mask=None):
        """
        参数:
            x_q:   Query的来源,形状为 (batch_size, seq_len_q, emb_size)
            x_kv:  Key和Value的共同来源,形状为 (batch_size, seq_len_kv, emb_size)
            mask: 可选的注意力掩码,形状为 (batch_size, seq_len_q, seq_len_kv)。
                  在需要屏蔽的位置为True(或1),例如padding位置或解码器的未来信息。
        返回:
            注意力输出,形状为 (batch_size, seq_len_q, v_size)
        """
        # 1. 线性投影,得到Q, K, V
        Q = self.Wq(x_q)   # (batch, seq_len_q, qk_size)
        K = self.Wk(x_kv)  # (batch, seq_len_kv, qk_size)
        V = self.Wv(x_kv)  # (batch, seq_len_kv, v_size)

        # 2. 计算注意力分数: Q * K^T / sqrt(dk)
        # 首先将K转置,使它的最后两维从 (seq_len_kv, qk_size) 变为 (qk_size, seq_len_kv)
        # 这样matmul(Q, K.transpose(1,2)) 得到 (batch, seq_len_q, seq_len_kv)
        dk = Q.size(-1)  # 获取qk_size,即dk
        scores = torch.matmul(Q, K.transpose(1, 2)) / math.sqrt(dk) # (batch, seq_len_q, seq_len_kv)

        # 3. 应用掩码(如果提供)
        # 掩码通常在解码器中使用,用于防止看到“未来”的信息。
        # 掩码为True的位置,我们希望其注意力权重为极小值(如-1e9),这样经过softmax后权重接近0。
        if mask is not None:
            # 使用masked_fill,将mask中为True的位置替换为一个很大的负值
            scores = scores.masked_fill(mask, -1e9)

        # 4. Softmax归一化得到注意力权重
        attn_weights = self.softmax(scores) # (batch, seq_len_q, seq_len_kv)

        # 5. 加权求和:注意力权重 * Value
        # attn_weights: (batch, seq_len_q, seq_len_kv)
        # V: (batch, seq_len_kv, v_size)
        # 输出: (batch, seq_len_q, v_size)
        output = torch.matmul(attn_weights, V)

        return output

关于mask的深入解释:这是面试中常被追问的点。mask主要有两种用途:

  1. Padding Mask:在批次训练中,句子长度不一,我们会用0填充到统一长度。在计算注意力时,这些填充位置不应该参与。我们会生成一个布尔掩码,在padding的位置标记为True,然后在scores上用masked_fill将其替换为一个极大的负值(如-1e9)。这样,在softmax之后,这些位置的权重就几乎为0。
  2. Causal Mask (Sequence Mask):在Transformer的解码器中,为了保证自回归特性(生成当前词时只能看到它之前的词),我们需要一个下三角掩码矩阵。这个矩阵的主对角线及以上(未来位置)都为True,同样用masked_fill处理。

我们的代码通过一个通用的mask参数支持了这两种场景,这是工程完备性的体现。

4. 自注意力 vs. 交叉注意力:测试与深度对比

类写好了,不跑通测试心里总不踏实。我们来写一个测试脚本,直观地看看自注意力和交叉注意力到底怎么用,以及它们的输出有什么不同。

4.1 自注意力测试:句子理解自己

自注意力就像是让句子里的每个词“反省”自己,并观察与句子中其他所有词的关系。

if __name__ == '__main__':
    print("=== 测试1:单头自注意力 ===")
    # 定义参数
    batch_size = 2  # 两个句子
    seq_len = 5     # 每个句子5个词
    emb_size = 8    # 词向量维度8(为了演示方便设小一点)
    qk_size = 6     # Q和K的投影维度
    v_size = 4      # V的投影维度

    # 创建模拟输入数据:两个句子,每个句子5个词,每个词8维向量
    x = torch.randn(batch_size, seq_len, emb_size)
    print(f"输入x的形状: {x.shape}") # [2, 5, 8]

    # 初始化注意力层
    self_attn = OneHeadAttention(emb_size, qk_size, v_size)

    # 模拟一个padding mask:假设第一个句子的最后2个词是padding,第二个句子的最后1个词是padding
    mask = torch.zeros(batch_size, seq_len, seq_len, dtype=torch.bool)
    mask[0, :, -2:] = True  # 第一个句子,所有Query都不能关注最后两个Key(padding)
    mask[1, :, -1:] = True  # 第二个句子,所有Query都不能关注最后一个Key(padding)
    print(f"掩码形状: {mask.shape}")

    # 前向传播:自注意力就是x既作为Query源,也作为Key/Value源
    output_self = self_attn(x, x, mask)
    print(f"自注意力输出形状: {output_self.shape}") # 应为 [2, 5, 4]
    print(f"输出示例(第一个句子的第一个词): {output_self[0, 0]}")

运行结果分析: 输入x的形状是(2, 5, 8)。经过自注意力层后,输出的形状变成了(2, 5, 4)。这里的变化很有意思:

  • batch_sizeseq_len(5)保持不变,因为我们对序列中的每个位置都计算了一个输出。
  • 最后一个维度从输入的emb_size=8变成了v_size=4。这是因为我们聚合的是投影后的Value向量,其维度是v_size自注意力的输出,其序列长度与Query源的长度一致,特征维度与Value的维度一致。

4.2 交叉注意力测试:连接两个世界

交叉注意力是连接两个不同序列或信息源的桥梁。最典型的应用是Transformer解码器层:Query来自已生成的目标语言序列,而KeyValue来自编码器输出的源语言序列表示。

    print("\n=== 测试2:单头交叉注意力 ===")
    # 定义参数(可以和自注意力不同)
    batch_size = 2
    q_seq_len = 3      # 目标序列长度(例如已生成的翻译词)
    kv_seq_len = 6     # 源序列长度(例如待翻译的原文)
    emb_size = 8
    qk_size = 6
    v_size = 4

    # 创建模拟输入:Query源和Key/Value源不同
    x_q = torch.randn(batch_size, q_seq_len, emb_size)   # 目标序列
    x_kv = torch.randn(batch_size, kv_seq_len, emb_size) # 源序列
    print(f"Query源x_q形状: {x_q.shape}")   # [2, 3, 8]
    print(f"Key/Value源x_kv形状: {x_kv.shape}") # [2, 6, 8]

    # 初始化注意力层(可以和自注意力是同一个实例,因为结构一样)
    cross_attn = OneHeadAttention(emb_size, qk_size, v_size)

    # 交叉注意力的mask通常用于屏蔽源序列中的padding
    mask_cross = torch.zeros(batch_size, q_seq_len, kv_seq_len, dtype=torch.bool)
    mask_cross[0, :, -2:] = True  # 假设第一个样本的源序列最后两个位置是padding
    # 第二个样本没有padding

    # 前向传播:输入两个不同的张量
    output_cross = cross_attn(x_q, x_kv, mask_cross)
    print(f"交叉注意力输出形状: {output_cross.shape}") # 应为 [2, 3, 4]
    print(f"输出示例(第一个样本的第一个目标词): {output_cross[0, 0]}")

运行结果分析: 这次,Queryx_q的形状是(2, 3, 8)Key/Valuex_kv的形状是(2, 6, 8)。输出的形状是(2, 3, 4)。 关键点来了:输出的序列长度(3)等于Query源的长度(q_seq_len),而不是Key/Value源的长度(6)。这完美体现了交叉注意力的工作模式:对于目标序列中的每一个位置(Query),我们去源序列(Key/Value)中寻找相关信息,并聚合回来。输出的每个位置,都融合了整个源序列的信息。

4.3 对比总结与面试要点

为了让你在面试时能清晰表达,我把两者的核心区别总结成下面这个表格:

特性自注意力交叉注意力
输入源单一源 (x_q == x_kv)双源 (x_qx_kv 通常不同)
物理意义序列内部元素间的相互关联两个不同序列或模态间的信息对齐与检索
Query来源序列自身序列A (如目标序列)
Key/Value来源序列自身序列B (如源序列)
输出序列长度等于输入序列长度等于Query源序列A的长度
典型应用Transformer编码器、BERT等Transformer解码器、视觉问答、多模态融合

面试时,如果被要求手写,写完代码后,面试官很可能会追问:“如果让你改成多头注意力,你会怎么做?” 这时你可以从容回答:核心思想是并行。我们可以让W_Q, W_K, W_V的投影维度变成num_heads * head_dim,然后在forward中,通过reshapetranspose(batch, seq_len, num_heads*head_dim)的张量变为(batch, num_heads, seq_len, head_dim),接着在num_heads这个维度上并行计算多个单头注意力,最后再把结果合并起来。这其实就是把我们今天实现的单头模块复制多份,并行计算。

5. 面试实战技巧与常见坑点

理论懂了,代码会写了,最后我们聊聊面试现场怎么发挥,以及哪些地方容易踩坑。这些是我自己面试别人和被面试时,总结出的血泪经验。

技巧一:先讲流程,再写代码 不要一上来就闷头写class。先在白板或共享屏幕上画出注意力计算的流程图,或者写出那五个步骤。向面试官清晰地阐述Q, K, V的含义、缩放因子的作用、softmax的目的。这能展示你的沟通能力和结构化思维。面试官确认你思路正确后,再动笔写代码,会顺畅很多。

技巧二:重视形状注释 在代码的关键步骤后,用注释标明张量的形状变化。比如:

Q = self.Wq(x_q)   # (batch, seq_len_q, qk_size)
K = self.Wk(x_kv)  # (batch, seq_len_kv, qk_size)
scores = torch.matmul(Q, K.transpose(1, 2)) # (batch, seq_len_q, seq_len_kv)

这能极大减少你在矩阵乘法时犯维度错误的风险,同时也让面试官一眼看到你对数据流的把握。

技巧三:主动提及Mask和性能 写完基础代码后,可以主动说:“在实际应用中,我们通常还需要支持注意力掩码,比如处理变长序列时的padding mask,或者在解码器中防止信息泄露的causal mask。” 然后简要说明实现思路。如果时间允许,甚至可以提一句:“对于超长序列,点积注意力的计算复杂度是O(n²),在实际工业级模型中可能会用到稀疏注意力、线性注意力等优化方法。” 这能展现你的知识广度。

常见坑点与检查清单

  1. 维度错误QK^T相乘时,务必确保qk_size维度对齐。最安全的做法就是像我们代码里那样,先转置K的最后两个维度。
  2. 缩放因子忘记开方dkqk_size,缩放时是除以√dk,不是dk。这个细节面试官一定会检查。
  3. Softmax维度nn.Softmax(dim=-1)中的dim=-1意味着在最后一个维度(即seq_len_kv,所有Key)上进行归一化,为每个Query生成一个权重分布。这是最容易写错的地方之一。
  4. Mask的值:使用masked_fill时,填充的值应该是一个极大的负数(如-1e9),这样经过softmax后权重才趋近于0。如果填0softmax后还会有一个基础概率,就起不到屏蔽作用了。
  5. 忘记处理Batch维度:所有线性层和矩阵乘法都要考虑batch维度在最前面。我们的实现从一开始就考虑了batch

最后,保持冷静。手写代码时有点小紧张很正常,如果一时卡壳,可以边写边向面试官解释你的思考过程。面试官考察的不仅是最终代码的正确性,更是你解决问题的逻辑和调试能力。把今天的内容消化好,单头注意力这道题,你一定能稳稳拿下。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐