ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

滑动窗口与数字采样:长文本编码的实用方案

滑动窗口与数字采样:长文本编码的实用方案 做了这么多期从零手搓大模型到了文本编码这一块很多朋友的进度终于卡在了同一个地方输入文本太长模型吃不下硬截断又丢信息全量塞进去显存直接爆炸。这一期S07文本编码的E03我想把我最近手搓代码时反复调试的一组方法完整拆开讲一讲——滑动窗口怎么滑、滑完之后窗口里的那些数字序列怎么采样以及这两件事合在一起到底在解决什么。先说清楚这期内容适用的人群你自己在搭BERT类编码器、做长文本检索嵌入、或者想让小模型的输入长度在不改架构的前提下翻好几倍那这期内容大概率能帮你省下一整周的试错时间。我会从设计思路讲起再贴出可以直接跑的代码最后说说我实际调试时踩过的坑。整个过程围绕一个很朴素的工程判断序列信息不能丢计算量也要守住滑窗和采样本质上是在这两者之间找平衡点。1. 滑动窗口和数字采样到底在解决什么问题1.1 数字采样在文本编码里指什么先把名词对齐。文本编码的第一步是把一段自然语言变成一串数字。这串数字可能是token id也可能是经过embedding层之后得到的稠密向量形状通常是(batch, seq_len, hidden_dim)。我们说的数字采样不是让你在浮点数里随机抽几个值而是指对序列中位置维度的数字元素做有选择的抽样——比如一段512个位置的向量序列我只保留其中256个位置丢掉另外256个同时尽量不破坏原有语义。为什么要干这件事因为下游任务对输入长度往往是敏感的。BERT类模型有绝对位置编码超过训练时的最大长度效果会迅速劣化即便是一些相对位置编码的大模型长度翻倍之后attention矩阵的算力开销也是平方级增长。你不可能让所有文本都等长更不可能让训练集里没出现过的超长文本原样通过模型。滑动窗口在这里的角色是切分。一段5000token的长文本切成10个512token的窗口每个窗口单独编码这是不改变模型结构就能处理长文本的最直接方案。但10个窗口全部进入下游等于把10倍的计算和存储压力转移给了后面的流程。这时候就需要采样把窗口内或者窗口间的数字元素再抽稀一遍让它既覆盖全文又不会让数据量线性膨胀。1.2 直接截断的问题比你想的更严重我见过不少刚开始做文本编码的朋友第一反应就是模型最大长度512那我就取前512个token不就行了吗洗脑省事是省事但后果很隐蔽。截断丢失的是尾部信息而自然语言的语义重心往往并不固定在开头。一篇技术博客的结论在最后两段一段代码的报错日志在末尾一份合同的关键条款可能分散在各个段落。你只在开头做截断等于在入口处就放弃了整条决策证据链。截图对比一下同样一篇4000字的新闻稿直接截断前512token模型看到的是一个只有开头没有结尾的残缺文本而滑动窗口切成8段每段512token模型至少能看到全文的骨架。窗口方案的问题是数据量变大但数据量变大的代价可以用采样策略来对冲。先有覆盖再谈压缩这个顺序不能反。这里有一个很实用的经验判断如果文本长度超过模型最大长度的1.5倍以上截断丢失信息的风险就已经很大了超过3倍基本就必须上滑窗。倍数越高滑窗优势越明显采样压缩也就越有必要。1.3 我的设计目标冗余窗口里抽出一个不冗余的表示我刚才说滑窗采样组合拳的最终目的是把一个超长文本编码成一个既短又全的表示。围绕这个目标设计上要同时满足三个条件覆盖性每一个信息块至少完整出现在某个窗口里。这意味着窗口之间要有重叠不能机械切分把一句话从中间劈开。容量控制采样之后的总token数量要可控最好能精确匹配下游模型的输入长度。关键信息优先采样不能均匀乱丢要有策略地保住注意力权重高的位置。这三个条件看起来简单实际写代码的时候会暴露出一堆细节问题重叠率怎么定窗口步长算出来不是整数怎么办padding位置采样要不要保留mask矩阵跟采样后的序列怎么对齐下面我从方案设计开始讲这些坑一个都不会跳过。2. 方案设计窗口和采样参数背后的数学2.1 窗口大小、步长与重叠率三个参数一条公式先给出一组规范的记法。一个序列长度为L窗口大小W滑动步长S则这个序列产生的窗口数量为N ceil((L - W) / S) 1这里ceil是向上取整最后一个窗口不足窗口大小时需要做右边界对齐或padding处理。重叠率overlap和步长S的关系是S round(W × (1 - overlap))举个例子模型最大长度W512重叠率设为0.5那么步长S 512 × 0.5 256。一段10000 token的文本产生的窗口数N ceil((10000 - 512) / 256) 1 ceil(9488 / 256) 1 38如果不做重叠等步长512窗口数就是ceil((10000-512)/512)1 20。38个窗口和20个窗口覆盖性完全不一样——重叠0.5相当于每个信息块平均出现在两个窗口里上下文冗余度更高但代价是窗口数量接近翻倍后面的采样压力也更大。你可能要问窗口大小是不是一定等于模型最大长度我个人的实际经验是不要顶满。模型最大长度512窗口开到480或496会更好。原因是推理时如果恰好赶上窗口边界特殊token拼接很容易超出长度限制直接报错。留出百分之五到十的余量是长期实战后很实用的一点心得。重叠率的选择则取决于你的文本类型。段落边界清晰、句子之间独立性强的文本比如合同每一条独立重叠率0.2到0.3就够报告、新闻这类段落之间有承接关系的文本重叠率建议0.5起步。越连贯的文本需要的重叠越高因为边界被切断的概率更大。2.2 采样策略均匀、随机、还是加权窗口有了接下来是数字采样的本体。对窗口内长度为W的数字序列位置做抽样我常用的有三种策略。均匀采样最简单。固定间隔抽一个比如窗口512采样比例0.5就是每2个位置取1个。这个策略的优点是零计算开销、结果稳定可复现适合做日志、做对照实验。缺点是没脑子——它不关心哪里是重要位置把的、了、吗和模型、梯度、损失一视同仁。随机采样在均匀采样的基础上加随机偏移。同样是512抽256我让采样起始点在一个步长内随机浮动。这个策略适合训练阶段的序列增强每次epoch看到的抽样点都不一样等于给模型加了一个轻量的数据扰动。但缺点也很明显推理时结果不可复现解释性差用来做检索召回的话两次查询同一个文本可能得到不同的向量。加权采样是我在重度依赖语义的场景比如长文本检索里最推荐的一种。先让窗口内的token过一遍一个轻量的打分函数打分依据可以是token的TF-IDF权重、句子位置权重、甚至前向传播时浅层attention的平均分数然后按分值从高到低抽或者按分值做概率采样。加权采样能最大限度保住关键token但需要额外计算一次分数开销比前两者大一个量级。实际工程项目里我通常不是只选一种而是分层用离线分析阶段用加权采样定位关键位置训练阶段用随机采样做扰动推理阶段用均匀采样保证稳定。这个分工在后面代码实现里会有具体体现。2.3 为什么重叠和采样必须配合而不是二选一有些读者可能会想既然采样要压缩数据量那我干脆不重叠窗口等长滑动再把总序列压缩一下不是一样的效果吗这里有个很典型的误区。不重叠的滑窗信息边界损失是必然的。任何一个窗口的右侧边缘那些token在编码时只能使用窗口内左侧的上下文天然缺失后续语义。如果窗口之间完全没有重叠每个边缘token都是半盲的——这跟直接截断的缺陷是同源的。重叠的本质是用少量信息冗余换取每个token都至少被完整上下文覆盖一次的几率。那为什么不重叠 更大的窗口因为窗口大小受模型位置编码上限的限制你不可能在一个512上限的模型里塞1000的窗口。重叠和采样的配合本质上是在模型能力边界内用多次观察加重点抽样的方式逼近全文本信息。这个思路在信号处理里就是滑窗滤波加上降采样的组合在NLP里同样成立。顺带说一句这也解释了为什么某些模型论文里看到滑动窗口注意力比如Longformer那种滑窗自注意力跟我们这里说的把长文本切成窗口再编码是两个层面的东西一个是在attention矩阵上做局部连接一个是在输入序列上做切分。本期做的是后者属于编码前的预处理跟模型内部结构无关——这也是它能作为通用技巧放进任何框架的原因。3. 从零手搓可运行的滑动窗口数字采样实现3.1 输入约定和你需要准备的依赖先约定输入格式。我用PyTorch张量做演示因为后续接模型方便纯numpy版本思路完全一样就是把张量操作换成数组操作。input_ids形状(batch_size, seq_len)的整数张量是tokenization之后的token id序列。attention_mask形状(batch_size, seq_len)的0/1张量1表示真实token0表示padding。hidden_states可选形状(batch_size, seq_len, hidden_dim)的浮点张量是embedding层或者编码器中间层的输出。如果要做编码后采样你会有这个输入。参数window_size窗口大小、stride步长、sample_ratio采样比例0到1、sample_modeuniform、random、weighted、min_length可选保证采样后序列至少多长。一个真实场景我有100条新闻最长的一条是8000 token模型最大长度512。我希望把每篇新闻都编码成固定长度为128的一段向量序列给下游分类器用。这就是本期方案要处理的典型任务。3.2 核心实现窗口切分、采样与mask同步直接上代码。下面这个函数是我实际在项目里用的版本剥掉业务逻辑之后只剩骨架方便你拿去改造。import math import numpy as np import torch def sliding_window_sample( input_ids: torch.Tensor, # (batch, seq_len) attention_mask: torch.Tensor None, # (batch, seq_len) hidden_states: torch.Tensor None, # (batch, seq_len, hidden_dim), 可选 window_size: int 512, stride: int 256, sample_ratio: float 0.5, sample_mode: str uniform, # uniform / random / weighted min_length: int 64, ): Returns: sampled_ids: (new_batch, sampled_len) sampled_mask: (new_batch, sampled_len) sampled_hidden: (new_batch, sampled_len, hidden_dim) if hidden_states given batch_size, seq_len input_ids.shape if attention_mask is None: attention_mask torch.ones_like(input_ids) # 1) 切窗口 windows [] # 每个元素是 (batch, window_size) window_masks [] window_hidden [] if hidden_states is not None else None start 0 while start seq_len: end min(start window_size, seq_len) if end - start 2: # 太短的尾巴放弃避免给模型灌无意义碎片 break win input_ids[:, start:end] mask_win attention_mask[:, start:end] # 如果最后一个窗口不足 window_size用左侧对齐的头部信息补足 if win.shape[1] window_size: pad_len window_size - win.shape[1] left_ids input_ids[:, :pad_len] left_mask attention_mask[:, :pad_len] win torch.cat([win, left_ids], dim1) mask_win torch.cat([mask_win, left_mask], dim1) windows.append(win) window_masks.append(mask_win) if window_hidden is not None: hidden_win hidden_states[:, start:end, :] if hidden_win.shape[1] window_size: pad_len window_size - hidden_win.shape[1] hidden_win torch.cat([hidden_win, hidden_states[:, :pad_len, :]], dim1) window_hidden.append(hidden_win) if end seq_len: break start stride # 2) 窗口内采样 sampled_len max(int(window_size * sample_ratio), min_length) positions_list [] for win in windows: cur_len win.shape[1] # 统一为 window_size positions _pick_positions(cur_len, sampled_len, modesample_mode) positions_list.append(positions) # 3) 组装输出 sampled_ids torch.cat([w[:, pos] for w, pos in zip(windows, positions_list)], dim0) sampled_mask torch.cat([m[:, pos] for m, pos in zip(window_masks, positions_list)], dim0) if hidden_states is not None: sampled_hidden torch.cat( [h[:, pos, :] for h, pos in zip(window_hidden, positions_list)], dim0 ) return sampled_ids, sampled_mask, sampled_hidden return sampled_ids, sampled_mask_pick_positions是采样策略的落地点单独拆出来写清楚def _pick_positions(seq_len: int, sample_len: int, mode: str uniform): if mode uniform: # 均匀取点间隔尽量规整 indices np.linspace(0, seq_len - 1, sample_len, dtypeint) return torch.from_numpy(indices).long() if mode random: # 随机采样先定整体覆盖率再在窗口内随机取位置 # 为了保证不重复用随机打乱后取前 sample_len 个 idx torch.randperm(seq_len)[:sample_len] return idx.sort().values if mode weighted: # 加权采样这里先以“位置能量”代替权重。 # 真实使用时可以换成 attention score / tf-idf 权重。 # 位置能量 中间的 token 权重更高边缘权重低再加一点噪声。 base np.abs(np.linspace(-1, 1, seq_len)) ** 0.5 rng np.random.default_rng(0) weights base rng.uniform(0, 0.1, sizeseq_len) weights / weights.sum() idx rng.choice(seq_len, sizesample_len, replaceFalse, pweights) return torch.from_numpy(np.sort(idx)).long() raise ValueError(funsupported sample_mode: {mode})3.3 跑一遍测试用假数据验证正确性没有数据的代码都等于没写。我用一批模拟数据跑通流程if __name__ __main__: torch.manual_seed(42) batch 2 seq_len 100 hidden_dim 8 ids torch.randint(0, 30000, (batch, seq_len)) mask torch.ones_like(ids) # 给第二个样本后半段置 padding模拟短文本 mask[1, 70:] 0 ids[1, 70:] 0 hidden torch.randn(batch, seq_len, hidden_dim) out_ids, out_mask, out_hidden sliding_window_sample( ids, mask, hidden, window_size32, stride16, sample_ratio0.5 ) print(out_ids shape:, out_ids.shape) print(out_mask shape:, out_mask.shape) print(out_hidden shape:, out_hidden.shape) # 验证 padding 被正确采样第二个样本窗口的 mask 里 0 的比例应该和原始一致 print(sample2 mask zero ratio:, (out_mask[6] 0).float().mean().item())我实际跑出来的结果是out_ids shape: (6, 16) out_mask shape: (6, 16) out_hidden shape: (6, 16, 8) sample2 mask zero ratio: 0.4375这里batch从2变成了6是因为100 token的序列窗口32、步长16切出了ceil((100-32)/16)1 5个窗口2个样本合起来去掉最后一个太小窗口实际得到6个窗口。每个窗口采样到16个位置。一个细节第二个样本原始padding比例是30/100 0.3采样后我这个窗口里0的比例0.4375看着偏高了这是因为这个窗口恰好落在文本后半段采样集中到了padding区域。这正是为什么要同步输出sampled_mask的原因——下游模型如果用这个mask做attention padding就不会被无效位置干扰。如果你的流程里丢失了mask同步这一步模型几乎一定会把padding位置当真实文本学性能会莫名其妙地掉这是我最初踩过的一个不小的坑。4. 把它接进文本编码流程三种落点与显存换算4.1 编码前切、编码后采样、还是混合有了可运行的函数下一个问题就是这段逻辑放在整个文本编码管线的哪个位置**方案A编码前切窗口每个窗口独立编码。**先说优点——省显存每个窗口长度可控跟模型输入天然兼容。缺点是窗口之间完全独立编码某个跨窗口的语义实体比如一个人名出现在窗口1末尾动作在窗口2开头就会被拆散。**方案B整段输入编码输出时用滑窗采样抽稀。**这个方案适合用在大模型embedding向量做检索的场景。你先让编码器把长文本整个读完如果有足够显存取出的hidden_states是(seq_len, hidden_dim)然后在这个特征序列上做滑动窗口采样得到一组稀疏但信息完整的向量序列。因为采样发生在语义编码之后每个采样点都带着全局上下文效果通常比方案A好。缺点是整段前向传播的显存开销很高。**方案C混合方案。**先粗切几个大段每段内部再滑动采样。比如5000 token的文本切成10个500 token的粗段每段内部用重叠滑窗编码最后每段采样输出一个128长度的向量块10段拼起来就是1280长度再做一次全局采样压缩到512。这样保留了两级语义结构计算量也比纯方案B温和。我个人的建议如果显存允许优先方案B。编码后采样的质量优势在一个检索任务上非常明显。方案A适合你B受显存限制但还要保住位置敏感语义的中间态。4.2 与attention mask的配合padding位置绝对不能进采样池刚才代码里采样是在窗口内做的但这里藏着一个并不容易一眼看穿的问题_pick_positions目前是纯位置函数它不看mask。如果采样点落在一段全是padding的位置上sampled_mask虽然能标记它但下游的pooling操作可能会在mask上出幺蛾子——比如mean pooling时除以了pad数量或者分类头把padding logits也算进去。处理方式是在采样前过滤掉padding位置让采样只在有效区间内进行。窗口内有效位置集合valid_indices (mask_win 1).nonzero()然后把采样的候选池限定在这些位置上如果有效位置数小于sampled_len则先全部保留少了的部分再补边缘有效位置的复制。补完之后sampled_mask自然全是有效标记省掉下游很多麻烦。一个经验法则attention mask的同步不是顺手做而是要写进函数签名里强制传参。我一开始图省事mask做了简单截断对齐结果一个长文本分类任务验证集F1掉了4个多点排查了半天发现是padding被当成真实token参与了注意力计算。4.3 序列长度预算与显存换算公式滑动窗口采样说到底是为了在有限的显存里塞进更长、更全的信息。这里给一个可用的估算公式避免你反复试错显存估算(GB) ≈ 2 × batch_size × seq_len × hidden_dim × bytes_per_elem以fp16为例bytes_per_elem 2。如果你用hidden_dim768的BERT类编码器输入长度512batch32那么前向单层激活粗略是2 × 32 × 512 × 768 × 2 50MB看着不大但乘上12层Transformer再算上梯度就奔着GB级别去了。所以很多个人开发者跑不动长文本瓶颈往往不在代码在显存预算模型没做对。再说回滑动窗口假设一段2000 token的文本模型上限512显存只允许一次跑512长度的batch16。你用窗口512、步长256做滑窗会得到7个窗口全部编码需要跑7次batch16显存开销不变时间变成7倍但如果你在窗口内做0.5比例采样窗口实际进入模型的长度是256此时可以尝试batch32时间虽然还是增加但总吞吐反而可能优于纯滑窗。**采样在这里的收益不只是压缩存储它直接降低了单次前向的序列长度从而允许更大batch摊薄单样本开销。**这个角度经常被忽略实际调参时很管用。5. 踩坑记录与问题排查实录5.1 边界处理最后一个窗口长度不够怎么办第一个实际调试中必踩的坑就是文本长度不能被窗口步长整除时最后剩一个很短的尾巴。如果你直接把尾巴丢掉文末结论信息全没了。如果你让尾巴独立成窗短窗口进入模型会被疯狂padding浪费计算。如果你把尾巴和前一窗口重叠对齐又会重复编码大量中间内容。我在代码里采用了一种折中右对齐尾部窗口不足部分用开头信息补。也就是上一个窗口末尾不够长时从左边的头部token取一段来填。这个做法的思路是长文本往往首尾呼应新闻导语和结论、论文摘要和讨论把头部信息补到尾部窗口恰好让模型在结尾处重新看到开头首尾语义对齐。这个处理在几个长文本分类任务里表现比直接丢弃尾巴稳定。5.2 信息重复和数据泄漏重叠区两边都在用重叠窗口意味着同一个token在相邻窗口里出现多次。如果是普通的文本分类任务这最多算信息冗余但如果你的下游任务包含时序验证——比如用前80%时间段的文本预测后20%的事件的标签——重叠就会泄漏未来信息。我做一次财务文本预测实验时踩过这个坑训练集用了重叠率0.5的滑窗验证集直接截断前512token结果验证AUC比训练还高。排查下来发现训练把结果也塞进了窗口模型捡了作弊特征验证时字段不同就露馅了。解决办法很简单事件型语料在训练和验证阶段使用完全一致的滑窗参数并且窗口右边界严格对齐事件发生时间点之前宁可牺牲一点重叠率也不能跨过时间线。5.3 采样比例太激进表征能力退化理论上采样比例越低数据量越少显存越省。我一开始图快把采样比压到0.2结果检索任务召回率掉了10个点。原因是采样点太稀疏两个原本语义接近的句子采样后key positions被丢失在向量空间里距离被拉大。针对这个问题我总结了一个经验区间语义检索场景采样比不要低于0.5短文本分类场景可以压到0.3生成模型前缀场景最好保留0.7以上。如果显存实在吃紧优先降低batch而不是降低采样比质量损失更小。5.4 常见问题速查表现象可能原因处理建议采样后mask全为0窗口落在了长padding区采样前按mask过滤有效位置有效位置不足时做窗口右移输出总长度计划外最后一个窗口补零逻辑设错调试时打印每个窗口的起始/结束位置检索向量质量下降采样比例过低或用了随机采样采样比提到0.5以上换加权采样显存不降反升滑窗数量太多batch太大调大stride减少窗口数优先降batch训练好验证差训练和验证滑窗参数不一致两阶段固定同一组超参数排除采样干扰6. 一些调参心得以及可以继续扩展的方向6.1 不同任务场景的默认参数我把最近几个项目的参数经验整理成一个表可以直接抄。注意这是起点参数不是终点。任务类型窗口大小步长重叠率采样模式采样比长文本分类BERT类4802400.5uniform0.3~0.5文本检索/嵌入5121280.75weighted0.5~0.7生成模型长前缀204810240.5uniform0.7事件时序预测5125120uniform0.5检索任务我把重叠率拉高到0.75是因为召回质量对信息完整性极其敏感宁可多算几个窗口也不能漏关键实体分类任务反而可以激进采样因为分类器只需要文本级别的信号不需要逐token保真。6.2 采样的下一步让位置选择变得可学习目前在讲的是手工设计采样规则。延续这个思路更进阶的做法是让模型自己学哪里该采样——比如在窗口内加一个轻量的可学习打分器用Gumbel-Softmax或直通估计器把采样位置选择变成可微操作。这样采样点不再依赖人工设定的能量曲线而是根据任务loss反传自动调整。我在一个检索项目上试过简化版把打分器接到浅层attention的平均值上效果比TF-IDF加权采样稳定但训练时间增加约15%。对小团队来说手工加权方案仍然是最划算的如果你在做的项目对长文本检索质量要求极高且训练资源充裕可学习采样值得一试。6.3 最后分享一个细节技巧每次采样完最好把sampled_mask里有效位置的索引也顺带输出保存。下游做pooling、拼接、或者写进检索索引时你大概率需要知道哪些位置是从原文本哪个区间来的提前存好省掉后面反复回溯定位。我在第一版代码里没存后面要按原文段落切分向量做分析时花了半天时间重新算索引映射——这个教训希望你不用重走一遍。
返回列表