反馈记忆网络FMN:让搜索系统拥有持续学习型对话能力

发布时间:2026/8/5 1:17:57
反馈记忆网络FMN:让搜索系统拥有持续学习型对话能力 1. 项目概述从“猜你喜欢”到“懂你所想”的进化在信息检索和推荐系统的日常工作中我们经常遇到一个经典难题用户的初始查询Query往往简短、模糊甚至词不达意。比如用户搜索“苹果”他到底是想买水果了解科技公司还是想看电影《苹果》传统的查询建议Query Suggestion系统大多基于搜索日志的共现频率或语义相似度给出“苹果手机”、“苹果电脑”、“苹果价格”等选项。这就像是一个记忆力超群的图书管理员能立刻报出历史上所有问过“苹果”的人后来都看了什么书。但这个管理员有个缺点他记性虽好却不会“察言观色”。他无法根据当前与你对话的上下文尤其是你对他之前推荐的反应来动态调整后续的建议。2018年国际万维网大会WWW上的一篇论文《Query Suggestion with Feedback Memory Network》所提出的“反馈记忆网络”Feedback Memory Network, FMN就是为了赋予这个“图书管理员”实时学习和记忆对话上下文的能力。它不再是一次性的、静态的推荐而是构建了一个持续的、有记忆的交互过程。核心思想很直观将用户与搜索系统的每一次交互提交查询、点击建议、忽略结果都视为一次“反馈”系统将这些反馈存储在一个结构化的“记忆”中。当用户输入下一个查询时系统不仅看查询本身还会“回忆”之前相关的交互历史从而生成更个性化、更精准的查询建议。这个模型的价值在于它模拟了人类对话中常见的澄清与递进过程。例如你第一次问“深度学习”系统推荐了“深度学习 教程”你点击了它这被记作一次正反馈。接着你问“有什么框架”系统会结合你刚刚对“教程”感兴趣的记忆优先推荐“TensorFlow 教程”或“PyTorch 入门”而不是泛泛的“机器学习框架”。这背后是FMN对短期会话上下文和长期用户兴趣的协同建模能力。对于从事搜索、广告、内容推荐等领域的技术人来说理解FMN不仅是为了复现一个模型更是掌握一种“让系统拥有持续学习型对话能力”的设计范式。接下来我将拆解其核心架构、实现细节并分享在复现与应用中积累的实战经验。2. 核心架构与原理深度拆解FMN的整体设计可以看作一个精巧的“记忆增强型”神经网络。它主要由四个核心模块构成查询编码器、反馈记忆模块、记忆读取与推理模块以及建议生成器。其工作流程是对当前查询和历史的反馈序列进行编码将历史反馈存入一个可外部寻址的记忆矩阵根据当前查询从记忆矩阵中读取最相关的历史信息融合当前查询与读取的记忆生成最终的查询建议。2.1 反馈的表示与记忆存储反馈Feedback是FMN的基石。论文中主要考虑了两种显式反馈点击用户点击了某个查询建议和忽略用户看到了但未点击。每一种反馈都被表示为一个三元组(q_t, a_t, f_t)其中q_t是用户在第t轮输入的原始查询a_t是系统当时给出的建议或用户实际执行的动作如点击的链接对应的查询f_t是反馈类型如点击为1忽略为-1。注意在实际工程中反馈的定义可以更丰富。例如可以将“停留时长”、“后续搜索行为”甚至“滚动深度”作为软性反馈信号但初期复现建议严格遵循论文的简化定义以降低复杂度。这些反馈三元组不会直接以原始文本形式存储。首先q_t和a_t会分别通过一个查询编码器通常是一个RNN如GRU或BERT等预训练模型转换为固定维度的向量表示u_t和v_t。然后这个反馈三元组被融合成一个统一的记忆向量m_t。论文采用的方法是一个简单的神经网络层m_t tanh(W_m * [u_t; v_t; f_t] b_m)这里[;]表示向量拼接W_m和b_m是可学习参数f_t作为标量嵌入。这个m_t就被写入到一个固定大小的外部记忆矩阵 M中。M的每一行都是一个记忆槽memory slot存储着一个历史反馈m_t。2.2 记忆的读取与注意力机制当新的查询q_{now}到来时系统首先将其编码为向量u_{now}。关键的一步是系统需要从记忆矩阵M中读取与当前查询最相关的历史信息。这不是简单的全量召回而是通过一种注意力机制Attention Mechanism进行加权读取。具体来说计算u_{now}与记忆矩阵M中每一个记忆向量m_i的相关性权重α_ie_i v^T * tanh(W_h * u_{now} W_m * m_i b_a)α_i softmax(e_i)其中W_h,W_m,v,b_a都是可学习参数。权重α_i反映了第i个历史反馈对当前查询的重要程度。最终读取的记忆向量o是所有权重下的加权和o Σ (α_i * m_i)这个设计非常巧妙。它允许模型动态地决定哪些历史反馈是相关的。例如如果当前查询是“Python”那么过去关于“Java安装”的反馈权重可能很低而关于“编程语言对比”或“Anaconda”的反馈权重会很高。这种能力使得FMN能够实现会话内Session内的连贯性推荐。2.3 建议生成从融合向量到候选排序读取的记忆向量o与当前查询向量u_{now}进行融合形成一个增强的上下文表示cc tanh(W_c * [u_{now}; o] b_c)这个向量c蕴含了“用户当前想问什么”以及“他之前对什么感兴趣”的双重信息。接下来FMN需要生成具体的查询建议。论文将其建模为一个排序问题。系统维护一个庞大的候选查询集合通常从搜索日志中挖掘得到。对于每一个候选查询q_cand同样将其编码为向量v_cand。然后计算该候选与上下文向量c的匹配分数score(q_cand) c^T * v_cand这个点积分数越高说明候选查询与当前上下文越相关。最后对所有候选查询按分数降序排列取Top-K作为查询建议输出。实操心得这里的候选集构建是影响效果的关键。不能直接用全量查询日志那样计算开销太大。通常的做法是先用传统的协同过滤或语义匹配模型如Sentence-BERT进行一次粗排从百万级候选集中筛选出千级别的候选再由FMN进行精排。这属于经典的“召回-排序”两阶段架构。3. 模型实现与训练细节理解了原理我们来看如何动手实现。我将基于PyTorch框架分享关键代码片段和训练技巧。3.1 数据准备与预处理首先需要构建训练数据。每条训练样本是一个会话序列[ (q1, a1, f1), (q2, a2, f2), ..., (q_T, a_T, f_T) ]。其中q_T是会话中最后一个查询作为当前查询(q1...q_{T-1})及其对应的反馈构成历史记忆而a_T用户实际在T时刻采取的动作对应的查询则作为正样本label。数据清洗从搜索日志中提取会话通常以30分钟用户无活动为界划分。清洗掉过短如长度3或过长如长度50的会话以及包含无效字符的查询。构建词典对所有查询进行分词中英文不同处理建立词到索引的映射。建议保留高频词如出现次数5低频词统一用UNK表示。生成训练对对于会话中每一个位置tt1将q_t作为当前查询其之前的所有(q_i, a_i, f_i)(it) 作为历史记忆a_t作为正例候选。同时需要为每个正例采样若干负例候选如随机从全局候选池中抽取未在本次会话中出现的查询。# 示例一个训练样本的数据结构 sample { ‘current_query‘: ‘python tutorial‘, # 编码后的索引序列 ‘history_feedbacks‘: [ # 历史反馈列表 {‘query‘: ‘how to code‘, ‘action‘: ‘programming basics‘, ‘feedback‘: 1}, # 点击 {‘query‘: ‘software‘, ‘action‘: ‘install python‘, ‘feedback‘: -1}, # 忽略 ], ‘positive_candidate‘: ‘python for beginners‘, # 正样本 ‘negative_candidates‘: [‘java download‘, ‘c book‘, ‘ai news‘] # 负样本采样得到 }3.2 模型组件的PyTorch实现以下是核心组件的简化实现import torch import torch.nn as nn import torch.nn.functional as F class QueryEncoder(nn.Module): 查询编码器使用GRU def __init__(self, vocab_size, embed_dim, hidden_dim): super().__init__() self.embedding nn.Embedding(vocab_size, embed_dim) self.gru nn.GRU(embed_dim, hidden_dim, batch_firstTrue) def forward(self, query_seq): # query_seq: [batch_size, seq_len] embedded self.embedding(query_seq) # [batch, seq_len, embed_dim] _, hidden self.gru(embedded) # hidden: [1, batch, hidden_dim] return hidden.squeeze(0) # [batch, hidden_dim] class FeedbackMemoryNetwork(nn.Module): def __init__(self, query_hidden_dim, mem_dim, candidate_size): super().__init__() self.query_encoder QueryEncoder(...) # 假设已定义 # 记忆融合层 self.feedback_fusion nn.Linear(query_hidden_dim*2 1, mem_dim) # 拼接q, a向量和反馈标量 # 注意力相关参数 self.attn_W_h nn.Linear(query_hidden_dim, mem_dim) self.attn_W_m nn.Linear(mem_dim, mem_dim) self.attn_v nn.Linear(mem_dim, 1) # 上下文融合层 self.context_fusion nn.Linear(query_hidden_dim mem_dim, mem_dim) # 候选评分层 (点积相似度无需额外参数) def forward(self, current_query, history_feedbacks, candidate_vecs): current_query: 当前查询索引序列 history_feedbacks: 列表每个元素是(q_seq, a_seq, feedback) candidate_vecs: 所有候选查询的预编码向量 [candidate_size, hidden_dim] # 1. 编码当前查询 u_now self.query_encoder(current_query) # [batch, hidden_dim] # 2. 构建记忆矩阵M memory_slots [] for (q_seq, a_seq, fb) in history_feedbacks: u_t self.query_encoder(q_seq) v_t self.query_encoder(a_seq) # 将反馈标量fb转换为与batch匹配的形状并拼接 fb_tensor torch.tensor(fb, deviceu_t.device).view(-1, 1).float() m_t torch.tanh(self.feedback_fusion(torch.cat([u_t, v_t, fb_tensor], dim1))) memory_slots.append(m_t) M torch.stack(memory_slots, dim1) # [batch, mem_len, mem_dim] # 3. 注意力读取记忆 # 计算注意力权重 u_expanded self.attn_W_h(u_now).unsqueeze(1) # [batch, 1, mem_dim] M_transformed self.attn_W_m(M) # [batch, mem_len, mem_dim] attn_scores self.attn_v(torch.tanh(u_expanded M_transformed)).squeeze(-1) # [batch, mem_len] attn_weights F.softmax(attn_scores, dim1) # [batch, mem_len] # 加权求和得到读取向量o o torch.sum(attn_weights.unsqueeze(-1) * M, dim1) # [batch, mem_dim] # 4. 融合上下文 c torch.tanh(self.context_fusion(torch.cat([u_now, o], dim1))) # [batch, mem_dim] # 5. 计算候选分数 (点积) # candidate_vecs: [candidate_size, hidden_dim] - 需要投影到mem_dim或者确保维度一致 # 这里假设candidate_vecs已经是mem_dim维度 scores torch.matmul(c, candidate_vecs.t()) # [batch, candidate_size] return scores3.3 训练目标与技巧FMN采用对比学习Contrastive Learning的思路进行训练使用最大间隔损失Max-margin Loss也称为铰链损失Hinge Loss。对于一批训练数据模型会为当前查询计算正例候选的分数s_pos和一系列负例候选的分数s_neg_i。损失函数鼓励正例分数比所有负例分数至少高出一个边界值marginLoss Σ_i max(0, margin - s_pos s_neg_i)def max_margin_loss(positive_scores, negative_scores, margin0.1): positive_scores: [batch_size, 1] negative_scores: [batch_size, num_neg] # positive_scores: [batch, 1], negative_scores: [batch, num_neg] # 计算正例与每个负例的差值 loss_per_pair F.relu(margin - positive_scores negative_scores) # [batch, num_neg] loss loss_per_pair.mean() # 对所有负例和batch求平均 return loss训练技巧负例采样负例的质量至关重要。除了随机负例可以加入“困难负例”Hard Negatives例如与正例语义相似但用户未点击的查询。这能显著提升模型的分辨能力。记忆容量与衰减外部记忆矩阵M的大小是超参数。太小记不住长历史太大会增加计算负担并引入噪声。可以为记忆槽设计衰减机制让久远的反馈权重自然降低。梯度裁剪RNN和注意力机制叠加训练时梯度可能不稳定。使用梯度裁剪torch.nn.utils.clip_grad_norm_能有效防止梯度爆炸。验证指标使用Mean Reciprocal Rank (MRR)和PrecisionK如P1, P5作为验证集指标而不是只看损失。这更贴近线上效果。4. 实战部署与优化经验将FMN从论文搬到生产环境会遇到许多在实验室里不曾有的挑战。以下是几个关键的实战环节。4.1 线上服务架构设计FMN模型不能直接部署为实时服务因为每次都要对当前查询编码、读取记忆、与海量候选计算点积延迟无法接受。必须采用预计算与索引的策略。候选查询索引离线预编码所有候选查询构建向量索引如使用FAISS。线上服务时只需计算上下文向量c然后用c去FAISS索引中进行最近邻搜索即最大点积搜索瞬间返回Top-K结果。用户记忆存储每个用户的记忆矩阵M需要持久化存储如Redis。当用户有新会话时从Redis读取其历史记忆向量用于计算。记忆的更新写入新的反馈可以异步进行以降低请求延迟。服务化将模型封装为gRPC或HTTP API服务。输入是用户ID和当前查询服务端拉取用户记忆执行模型前向传播查询FAISS返回建议列表。[用户请求] -- [API网关] -- [FMN服务] |--- [Redis] 读取用户记忆 |--- [模型计算] 生成上下文向量c |--- [FAISS] 检索Top-K候选 |--- 返回查询建议列表4.2 冷启动与记忆初始化问题新用户没有历史反馈记忆矩阵M为空FMN会退化为一个普通的语义匹配模型。为了解决冷启动利用用户画像如果能有用户的基本属性如地域、设备可以用一个轻量级网络将这些属性映射为一个“伪记忆向量”作为初始记忆。会话内快速启动即使同一会话内用户的前几次搜索也能快速积累记忆。确保模型对短记忆序列也能稳健工作。4.3 记忆的噪声与遗忘不是所有历史反馈都有用。用户可能误点或者兴趣发生了转移。长期存储所有反馈会导致记忆被噪声污染。反馈权重衰减为每个记忆向量m_t附加一个时间衰减因子例如weight exp(-λ * Δt)Δt是距离当前的时间间隔。在计算注意力权重前先对记忆向量进行衰减加权。记忆压缩与淘汰定期对用户的记忆矩阵进行“清理”。可以聚类相似的记忆向量或用一个小型自编码器学习其压缩表示只保留最具代表性的记忆。也可以设定一个最大容量采用LRU最近最少使用策略淘汰旧记忆。4.4 效果评估与A/B测试离线指标MRR, PK好不代表线上业务指标如点击率CTR、转化率一定提升。定义核心指标明确优化目标是提升建议的点击率还是提升点击后结果页的满意度如停留时长、后续转化。设计A/B测试将用户流量随机分为实验组使用FMN和对照组使用旧版建议系统。对比两组在核心指标上的差异。测试周期要足够长以覆盖不同用户的不同生命周期阶段。分析日志深入分析实验组日志看FMN在哪些场景下表现更好如长会话、复杂信息需求在哪些场景下不如基线如简单的事实型查询。这能为后续迭代提供方向。5. 常见陷阱与排查指南在实际开发和调优FMN的过程中我踩过不少坑。这里总结一份问题排查清单希望能帮你节省时间。问题现象可能原因排查步骤与解决方案离线指标很高线上CTR无提升甚至下降1. 训练/测试数据分布与线上真实分布不一致数据泄漏或偏差。2. 负例采样过于简单模型过拟合。3. 线上服务延迟高导致用户感知差放弃点击。1. 检查训练数据时间窗口是否包含了未来数据。确保训练/验证/测试集按时间划分。2. 在验证集中加入“困难负例”检查模型表现是否骤降。若是需改进负例采样策略。3. 监控线上服务P99延迟优化FAISS索引参数如nprobe或对候选集进行更激进的粗排过滤。模型对所有查询都给出相似的建议个性化不强1. 记忆读取注意力机制失效权重趋于均匀。2. 反馈信号f_t点击/忽略的区分度不够或融合方式有问题。3. 记忆矩阵维度太小无法有效存储信息。1. 可视化注意力权重α_i看是否集中。如果均匀尝试增大注意力层的维度或加入LayerNorm。2. 尝试不同的反馈编码方式例如将点击和忽略分别用不同的向量表示而非一个标量。3. 逐步增加mem_dim观察验证集指标变化找到一个饱和点。训练损失震荡不收敛1. 学习率设置过大。2. 梯度爆炸。3. Batch内样本差异过大如会话长度悬殊。1. 使用学习率预热Warmup和衰减Decay策略。2. 加入梯度裁剪clip_grad_norm_。3. 按会话长度进行分层采样或使用动态Padding确保一个Batch内长度相近。新用户无记忆效果比老用户差很多冷启动问题突出。模型过度依赖记忆自身语义匹配能力弱。1. 在损失函数中增加一项鼓励模型在无记忆时也能做好例如同时优化一个不依赖记忆的基准模型。2. 引入用户侧特征如人口属性作为记忆的补充输入缓解冷启动。内存或显存占用过高1. 记忆矩阵M存储了过长的历史序列。2. 候选集向量全部加载到内存。1. 为记忆长度设置上限或实现记忆淘汰策略。2. FAISS索引支持从磁盘分页加载。对于超大候选集考虑使用量化索引如IVFPQ在精度和内存间取得平衡。一个关键的调试技巧构建一个小的、可解释的测试集。例如人工构造几个有明确逻辑的会话序列比如[(音乐, 流行音乐, 1), (推荐一些, ?, ?)]。然后打印模型内部的注意力权重和最终的建议排名看是否符合你的直觉预期。这能帮你快速定位是数据问题、特征问题还是模型结构问题。FMN模型为我们打开了一扇门让搜索系统从“被动应答”走向“主动对话”。它的思想并不局限于查询建议可以扩展到对话系统、序列推荐等任何需要利用历史交互信息的场景。实现它的过程是对记忆机制、注意力模型和在线学习的一次深刻实践。最开始复现时你可能纠结于每一个超参数但当你看到系统能根据用户之前的行为精准地猜出他下一个可能想问的问题时那种成就感是实实在在的。记住模型是骨架而高质量、高信噪比的反馈数据才是其灵魂在工程中投入精力做好数据清洗和反馈定义往往比调参带来的收益更大。