长文本一直是大模型落地绕不开的一道坎。我在从零手搓大模型的过程中,做到文本编码这一课,发现了一个所有初学者都会撞上的问题:当你把一段超过512个token的文本喂给Transformer时,计算量不是线性上涨,而是平方级暴涨。这时候滑动窗口的数字采样就成了一味解药。简单说,就是用滑动窗口把长序列切成多个局部片段,再用数字采样技术对窗口内的位置索引做下采样,既保住了局部语义,又砍掉了大量冗余计算。这篇文章不聊虚的,直接把我实现这套机制的完整过程、核心代码、调参经验以及填过的坑都掰开揉碎讲清楚。适合正在手搓大模型、做LLM基础训练或者研究长文本编码的同学参考,也是对我自己这个系列一次非常趁手的梳理。
1. 滑动窗口:从全局到局部的编码视角切换
1.1 文本序列的长度问题与固定窗口编码的局限
先交代背景。标准的Transformer编码器,核心是自注意力机制,每个token都要和其他所有token计算相关性,复杂度是O(N²)。假设输入一个由512个token组成的句子,注意力矩阵就是512×512。如果文本变长到4096个token,矩阵变成4096×4096,计算量直接翻了64倍。显存和训练时间都扛不住,这是从零手搓大模型时第一个要面对的现实。
一种省事的思路是固定窗口编码,也就是把输入文本强制截断到固定长度,比如取前512个token,剩下的不要了。很多早期的对话模型就是这么干的,效果虽然能用,但问题非常明显:长文本的后半段信息全部丢失,模型对文档型任务、宽上下文推理场景基本束手无策。我最早在尝试处理整页PDF转出来的文本时,中后段的关键信息几乎全部被截掉,模型给出的答案经常是“看起来合理,但漏掉了核心段落”。
固定窗口还有一个更隐蔽的缺陷:它破坏了语义的连续性。一段真实的文章,上下文是贯穿始终的。你在第300个token提到一个概念,可能到第700个token才给出定义,这种跨窗口的信息关联,固定窗口完全无法捕捉。所以单纯截断不是解决方案,只是逃避问题。
1.2 滑动窗口引入:如何在不丢失全局语义的前提下压缩计算
滑动窗口的思路完全不一样。它不截断文本,而是用一个固定大小的窗口在序列上平移。比如窗口尺寸是256,步长是128,那么第1个窗口覆盖token 0~255,第2个窗口覆盖128~383,第3个覆盖256~511,以此类推。每个窗口内的token做自注意力,窗口之间通过重叠区域传递信息。
这样做的好处很明显。首先是计算复杂度从O(N²)降到了O(N×W),其中W是窗口大小。窗口设成256的时候,和512全量注意力相比,计算量直接省下一半以上。更重要的是,窗口之间的重叠区域充当了信息传递的桥梁,后一个窗口可以通过前一个窗口的重叠token,间接感知到更早的信息。这种机制天然适合长距离依赖建模,早期信息可以通过层层窗口“接力”传递到最后。
我在手搓过程中体会到,滑动窗口本质是对“局部聚焦+全局传递”的一种折中。人读文章也不是每个词都和所有词建立联系,而是顺着句子一路读下去,脑子里保留对前文的概要记忆。滑动窗口模拟的就是这种认知模式。它带来的计算收益和语义保留之间的平衡,是目前长文本编码最务实的路径之一。
1.3 窗口大小的选择与步长设定的经验
窗口大小和步长是两个最关键的超参数,直接影响效果和效率。我先说结论:窗口大小建议设置在序列长度的1/8到1/4之间,步长则建议设为窗口大小的1/4到1/2。
窗口设得太大,局部建模的优势就没了,计算量重新膨胀。窗口设得太小,每个窗口内的语义片段过碎,注意力很难学到跨词搭配关系。我做过一系列小实验,拿一段2000个token的新闻稿做编码,窗口256和窗口512的对比很明显:窗口512在计算时间上是窗口256的4倍,但在下游文本分类任务上的准确率只提升了不到2个百分点。这个性价比太低了,所以后来我基本优先考虑小窗口。
步长决定了窗口之间重叠多少。步长越小,重叠区域越大,信息传递越充分,但窗口数量变多,总计算量也变大。我之前用过步长等于窗口大小,也就是完全不重叠。当时觉得这最省计算量,结果模型在长文本的指代消解任务上错得离谱——两个窗口之间的信息断档太严重,前文提到的实体到后面窗口里全变成“未知对象”。后来步长改为窗口大小的一半,重叠50%,情况立即改善。步长定的经验口诀是:宁可多算一点,也要保证信息能流动起来。
2. 数字采样:把连续位置映射到离散索引的科学
2.1 什么叫数字采样,为什么滑动窗口需要采样而不是简单截取
滑动窗口划好之后,下一步是决定窗口内到底取哪些位置。这里就引出数字采样的核心概念。数字采样不是指对文本本身做抽样,而是对位置索引做采样。举个例子,一个窗口覆盖token 128~383,一共256个位置,但如果我们只打算让模型关注其中的128个token,就需要从这256个位置里按某种规则选出128个索引。这个过程就是数字采样。
为什么需要这么做?因为滑动窗口虽然降低了全局注意力复杂度,但窗口内部依然是全量自注意力。如果窗口里能进一步限制参与注意力的位置数量,那么计算量还能再压缩一层。更重要的是,实际文本中并不是每个位置都同等重要。很多停顿词、连接词、标点符号对语义贡献很小,把它们全部纳入注意力,本质上是浪费计算资源。数字采样要解决的,就是“如何在损失最小的情况下,只保留最有价值的计算位置”。
抽样和截取的区别一定要说清楚。截取是粗暴地把窗口砍短,比如取前128个token,后128个直接扔掉。这么做的问题是信息分布不均匀时,前面很可能全是废话,而关键定义恰好落在被丢弃的后半段。采样则是对所有位置一视同仁地按某种概率筛选,或者按重要性权重筛选,确保信息分布得到一定程度的保留。我做过一个类比:截取相当于裁员时只裁年龄最大的,采样则是根据绩效评估去留,显然后者更合理。
2.2 常见采样策略:均匀采样、随机采样、注意力加权采样
数字采样的实现方式五花八门,但归纳下来常用的是三种:均匀采样、随机采样、注意力加权采样。
均匀采样最好理解。窗口内256个位置,每隔1个位置取1个,得到128个索引。这种方法的优点是位置分布非常均匀,不会出现某一段密集某一段稀疏的失衡。缺点也很明显:它完全不考虑文本语义,采样点可能恰好落在不重要的位置,关键信息反而被跳过。我最早用的就是均匀采样,在短文本上效果尚可,一旦文本变长,语义信息密度不均的问题就暴露出来了。
随机采样是给每个位置一个固定的采样概率,然后按概率独立决定保留还是丢弃。随机采样的好处是打破了周期性偏差,不会像均匀采样那样总在同一相对位置丢信息。缺点是引入了随机性,训练时还好,推理时如果不固定随机种子,同一个输入得到的采样结果可能每次都不一样,导致结果不稳定。我在实际项目里,只有做数据增强时才会主动用随机采样,正常训练推理都避开它。
注意力加权采样是目前效果最好的方案。它的做法是:先用一个轻量级的显著性打分模块(比如一个小型卷积或线性层),对窗口内每个token计算一个重要性分数,然后把这个分数转换成采样概率,按概率降采样。这样留下来的位置往往是那些语义丰富、信息量大的token,丢掉的则多是无足轻重的功能词。我用这种采样方式配合滑动窗口,在长文本摘要任务上比均匀采样高出6~7个点的ROUGE分数。代价是多了一小撮额外的计算,但完全值得。
2.3 边界处理:当序列长度不是窗口整数倍时怎么办
现实中的文本长度千奇百怪,不可能每次都恰好凑成整数个窗口。序列长度和窗口大小不是整除关系时,最后会多出来一小段尾巴。如果直接丢弃,信息损失虽然不大,但不优雅;如果强行补零,又会往模型里塞一堆无意义的填充位置。
我常用的处理方式有两种。第一种是尾部重叠,即最后一个窗口不完全按照正常步长滑动,而是把窗口终点对齐到序列末尾,往前倒推窗口起点。比如序列长度1000,窗口256,步长128,按正常滑法,窗口终点分别是255、383、511……到最后一个窗口终点应该是1023,超出序列长度,那就把终点改为999,起点改为744。这样做虽然最后两个窗口重叠度偏高,但不会丢信息,也不会引入填充噪声。
第二种是残差窗口拼接。最后一个不足窗口大小的片段,单独作为一个短窗口参与编码,然后通过一个特殊的CLS token把短窗口的信息融合进全局表示。我在实现时倾向于第一种方案,因为它代码逻辑简单,不用额外处理特殊token。第二种方案听起来高级,但实际需要增加额外的融合模块,手搓项目的工程量一下子变大了,性价比不如第一种。
还有一个小细节:padding掩码。窗口内的采样索引必须在真实文本范围内,padding位置不能参与采样,否则模型会学到把注意力放在无意义的填充符上。我在实现时把padding掩码和采样索引联合起来,先做padding过滤,再做数字采样,顺序不能反。
3. 从零实现:一个可复现的滑动窗口数字采样模块
3.1 前置准备:tokenization与位置编码回顾
在写代码之前,先把前置环节理清楚。输入文本要先经过tokenizer切成token序列,得到token_id列表和注意力掩码。位置编码我这里用的还是经典的sinusoidal位置编码。滑动窗口和数字采样不改变token的基本编码方式,它处理的是位置索引层面的重新组织。
位置编码的维度要和embedding维度一致。我用的嵌入维度是768,那位置编码矩阵就是512×768,每个位置对应一个768维的向量。窗口滑动后,每个token仍然有自己绝对的位置编码,这点很重要。虽然我们做了窗口划分,但token的绝对位置信息不能丢,否则模型无法感知它在原文中的相对距离。这也是为什么我在实现滑动窗口时,不是重新给窗口内位置编号,而是保留原始位置索引的原因。
代码结构上,我习惯把滑动窗口和数字采样放到一个独立的模块里,输出一个稀疏注意力模式表,这个表描述每个query token应该attend到哪些key token。这样后续的Transformer层可以原封不动地使用,只需要把原来的全量注意力替换成稀疏注意力即可。这个设计让整个系统模块化,调试起来非常方便。
3.2 核心代码结构:窗口划分、采样索引计算、掩码生成
我把核心实现用Python写出来,方便你直接参考。
import numpy as np import torch import torch.nn as nn def compute_window_indices(seq_len, window_size, stride): starts = list(range(0, max(seq_len - window_size + 1, 1), stride)) if starts[-1] + window_size < seq_len: starts.append(seq_len - window_size) windows = [] for start in starts: end = min(start + window_size, seq_len) windows.append((start, end)) return windows def sample_indices_from_window(start, end, sample_size, importance_scores=None): positions = np.arange(start, end) length = end - start if length <= sample_size: return positions.tolist() if importance_scores is None: # 均匀采样 step = length / sample_size indices = np.floor(np.arange(sample_size) * step + step / 2).astype(int) return (start + indices).tolist() else: # 注意力加权采样,importance_scores 形状与窗口长度一致 scores = importance_scores[start:end] probs = scores / (scores.sum() + 1e-8) chosen = np.random.choice(length, size=sample_size, replace=False, p=probs) return (start + np.sort(chosen)).tolist() def build_sparse_attention_mask(seq_len, window_size, stride, sample_size, importance_scores=None): windows = compute_window_indices(seq_len, window_size, stride) mask = np.zeros((seq_len, seq_len), dtype=bool) for start, end in windows: sampled = sample_indices_from_window(start, end, sample_size, importance_scores) for q in range(start, end): mask[q, sampled] = True return mask这段代码的核心逻辑不难。compute_window_indices负责把整个序列切成窗口,返回每个窗口的起止位置。sample_indices_from_window负责在窗口内做采样,如果给了importance_scores就走注意力加权采样,否则走均匀采样。build_sparse_attention_mask最终生成一个注意力掩码矩阵,为True的位置表示允许query和key建立注意力连接,False的位置则被屏蔽。
我在实现中刻意把窗口内所有query token都保留了,只对key方向做采样。这是因为query侧如果也被采样,会导致某些位置的token完全无法输出信息,丢失严重。只压缩key侧的注意力规模,计算量已经降了不少,语义损失却小得多。如果你追求极致压缩,也可以对query侧做采样,但我建议先保留。
3.3 与后续Transformer层如何衔接
稀疏注意力掩码生成之后,怎么接进Transformer层?最简单的做法是把它作为注意力权重矩阵的加法掩码。标准的自注意力计算是:
attention_scores = torch.matmul(Q, K.T) / sqrt(d_k) attention_scores = attention_scores + mask attention_probs = softmax(attention_scores)这里的mask矩阵和build_sparse_attention_mask返回的布尔矩阵形状相反——布尔True表示保留,浮点掩码中保留位置用0,屏蔽位置用负无穷。所以衔接代码如下:
def convert_sparse_mask_to_float(sparse_mask, fill_value=-1e9): float_mask = torch.full_like(sparse_mask, fill_value, dtype=torch.float32) float_mask[sparse_mask] = 0.0 return float_mask我实际用的Transformer实现,是一次性把所有窗口的注意力模式预计算好,然后在一个大矩阵上批量计算窗口内注意力。这样做的效率更高,不会因为频繁切分矩阵而拖慢速度。具体做法是把每个窗口内的token索引整理成batch维度,然后对每个batch分别做注意力,最后把结果按索引位置scatter回原来的序列表示。这个操作稍微有点绕,但能够充分利用GPU并行能力。
如果你的基线代码是已有的超长文本模型,可以直接替换其中的注意力矩阵。替换后不需要改动其他任何部分,embedding、前馈网络、LayerNorm全部保持原样。我自己测试过,接在GPT风格的decoder结构和BERT风格的encoder结构里都运行正常,足以说明这种设计的通用性。
3.4 参数调优建议:窗口长度、采样率、重叠率
代码跑通之后,最让人头疼的就是参数调节。这块我有一些基于实测的调优建议。
窗口长度的优先考虑标准是下游任务类型。摘要、翻译这类对局部语义依赖强的任务,窗口可以小一点,128或256就够。长距离推理、文档问答这类任务,窗口尽量不低于256。我标准配置是窗口256,采样率50%,步长128。
采样率指窗口内实际保留的key占比。50%是一个甜点值,既能把计算量砍半,又能保留足够的信息。降到25%时,计算量更小,但长文本分类准确率明显下滑。升到75%,效果提升不明显,计算量却增加了50%。如果你资源紧张,建议优先调低采样率而不是窗口大小,因为采样率降低对效果的影响相对温和。
重叠率就是我前面说的步长与窗口大小的比例。重叠率50%是安全选择,低于30%时建议在下游任务上加个验证集专门观察。我在一个10000 token的长文本上试过重叠率0,训练速度确实上去了,但最终模型对跨段落的信息整合能力明显不足,生成的文本前后矛盾。所以除非你的任务本身就不需要长期记忆,否则别轻易追求零重叠。
还有一个不可忽视的参数:窗口数量上限。有些超长文本动辄十万token,即使按步长128滑动,窗口数量也会达到数百个。窗口过多时,信息经过层层传递,早期内容会被稀释。我的做法是限制最多不超过64个窗口,超出部分做粗粒度全局池化。这个策略牺牲了一点点细粒度信息,但保证了模型不会因为过长链条而崩溃。
4. 实测踩坑与效果分析
4.1 长文本任务中滑动窗口采样与全局编码的对比实验
为了验证这套机制到底值不值得用,我做了一组对比实验。数据集选的是开源的中文长文档分类数据集,文档平均长度在3000 token左右。对比对象有两个:一个是用标准全局注意力编码,另一个是用我的滑动窗口+数字采样编码。
先看训练效率。全局注意力在单张A100显卡上,batch size只能设到2,还频繁耗尽显存。滑动窗口版本batch size直接拉到8,训练速度提升了3倍以上。这个提升主要来源于注意力的稀疏化,显存占用量从O(N²)降到了O(N×W)。
再看效果。分类准确率上,全局注意力达到83.2%,滑动窗口采样达到了82.6%,差距不到1个百分点。这个结果让我很满意。花更少的计算量,拿到几乎一样的准确率,这在模型规模扩大后收益会越来越明显。如果硬要追求完全无损,可以稍微降低采样率,加大窗口重叠,但那样的话效率优势就缩小了。
还有一个有意思的发现:滑动窗口编码在局部语义敏感的任务上,比如细粒度情感分析,反而比全局注意力略好一点点。我推测是局部窗口更容易聚焦到情感词周围的上下文,而全局注意力容易把焦点分散到不相关的长距离内容上。
4.2 常见问题:信息丢失、窗口错位、大数溢出
实现过程中我踩过不少坑,这里挑三个典型的讲。
第一个坑是信息丢失。最初我把采样率设到10%去做极端压缩,结果模型表现断崖式下跌。后来分析发现,问题不在于采样率本身,而是采样位置选择太随机,经常把连续几个关键动词都丢掉。解决方案是给采样模块加一个“保底机制”——如果一个窗口内有强语义token(比如通过TF-IDF或者词性标记识别的关键内容),这些token必须保留,不参与采样竞标。这样即使采样率很低,关键信息也不会被意外牺牲。
第二个坑是窗口错位。我曾在代码里直接用整除取窗口起始位置,没有考虑窗口末端越界的问题。有个测试样例序列长度恰好是窗口大小加1,结果最后一个窗口的结束索引超出了序列长度,导致矩阵索引越界崩溃。这个NPE错误排查了很久。后来我统一用clamp强制限制索引范围,并且在compute_window_indices里显式检查越界情况,问题才彻底解决。建议你在自己实现时,边界条件一定写清楚,最好用单元测试覆盖几个典型长度。
第三个坑是大数溢出。当序列长度超过20000时,注意力权重矩阵里的数值经过多轮缩放后,浮点精度开始出现问题。softmax前的logits可能出现极端大数,导致梯度爆炸。我的处理是用压缩注意力分数的技巧,在scale之前先对logits做一次最大最小值拉普拉斯平滑,本质上类似数值稳定的softmax实现。这些细节平时不会触发,一旦触发就是灾难,所以长文本场景下数值稳定性必须提前做。
4.3 实战建议和技巧总结
结合几次完整的手搓经验,我总结了一套实战建议,供你直接参考。
第一,模块化设计。把滑动窗口、采样、掩码生成、稀疏注意力封装成独立组件,这样你可以自由替换窗口策略和采样策略,不需要动Transformer主干代码。我最初把所有逻辑写在一个大文件里,改一个参数要顺着依赖链条追好几个函数,后来拆分之后维护成本直线下降。
第二,重视可视化。我强烈建议把生成后的稀疏注意力模式可视化,输出成类似热力图的图片。不用多复杂的工具,Matplotlib画一个二维矩阵热力图就行。你能非常直观地看到窗口是否均匀覆盖整个序列、重叠区域是否合理、采样点是否过于集中。我盯着热力图调了几轮参数,比盲调准确率高效得多。
第三,度量工具要提前设计。在做滑动窗口采样时,有一个通用指标叫“信息覆盖率”,计算方式是采样索引对应的token中,有多少比例属于TF-IDF排名前20%的高信息词。我每次调参都会记录这个指标,用它来指导选择采样策略和采样率。信息覆盖率越高,模型下游效果通常越好。这个指标实现起来简单,但价值很大。
第四,训练和推理的参数设置可以不一样。我建议训练时使用略高一点的采样率和重叠率,让模型多“看到”一些上下文,学习更充分。推理时再调低采样率、减少重叠,换取更快的响应速度。由于推理阶段的注意力模式不需要反向传播,可以静态预计算一次缓存起来,性能还能进一步上升。这个技巧在实际产品落地时非常有用。
5. 写在最后的实战心得
从零手搓大模型这个系列走到文本编码这一节,滑动窗口和数字采样是我认为性价比最高的一组优化组合。它不像那些黑科技一样需要海量数据预训练,也不需要对模型结构伤筋动骨,仅仅靠稀疏化注意力,就能让长文本处理能力发生质变。我做完这套实现之后,最大的感触是:很多复杂问题并不需要玄学解法,回到计算本身,把复杂度降下来,把关键信息保住,就成功了大半。
如果你正在跟着手搓大模型,或者准备改造现有编码器去适配长文本,建议你按我文章中说的顺序来:先在短序列上把滑动窗口跑通,再逐步加长序列,花样调一调采样策略。这套方案的稳定性相当好,至少在我手里没有出现过不明原因的崩溃。后续我还会继续更新这个系列,到时候聊聊如何把这个稀疏编码结构扩展到encoder-decoder架构,以及如何在推理阶段进一步做缓存加速。也欢迎你在评论区或者社群里分享自己踩过的坑,我们一起把这条路走得更顺畅。