ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

PyTorch变长序列处理:从DataLoader报错到LSTM pack全攻略

PyTorch变长序列处理:从DataLoader报错到LSTM pack全攻略 前阵子一个做水文预报的朋友半夜给我发来一段报错代码在单条数据上明明跑得好好的一放进DataLoader就炸。我把报错往下一拖看到一行熟悉的话RuntimeError: stack expects each tensor to be equal size, but got [142, 6] at entry 0 and [118, 6] at entry 1。这属于PyTorch里最经典的“数据长度不一致”问题——DataLoader默认按stack方式拼batchstack的硬性要求就是所有样本形状相等序列长度对不上训练直接崩。这篇文章就围绕这个坑展开DataLoader的默认行为为什么不容忍变长、LSTM处理变长序列的标准姿势是什么、怎样通过自定义collate_fn把数据管线彻底改造成变长友好以及我在实际项目里踩过的和序列长度有关的各种边角坑。适合用LSTM做时间序列预测水文径流、气象、金融都算、做NLP分类或序列标注的读者只要你需要把不定长序列喂给PyTorch这篇东西应该能帮你省点时间。1. 先搞懂报错DataLoader默认batch策略为什么不容忍变长1.1 runtime error背后发生了什么很多人看到“stack expects each tensor to be equal size”第一反应是去检查模型输入维度但问题压根没到模型那一步。DataLoader在产出每个batch之前会调用collate函数把当前batch里的若干样本拼成一个tensor。默认的collate逻辑是torch.utils.data._utils.collate.default_collate它内部对Tensor类型字段做的就是torch.stack。stack要求所有tensor除新增批次维度外其余维度必须完全一致。举个例子假设batch里有两条样本一条形状是[142, 6]另一条是[118, 6]。它们代表的是“142个时间步、每个时间步6个特征”和“118个时间步、每个时间步6个特征”。强行stack成三维tensor时第二维一个是142一个是118对不上于是抛出上面那个错误。这个过程发生在数据加载阶段和你的LSTM层、优化器、loss函数没有任何关系。这个问题的本质是序列数据天然不是矩形。NLP里一句话10个词另一句话50个词水文站的径流序列上游站建站早可能积累了几千个时段的观测下游站可能只有几百个时段的记录。只要你的模型输入是变长序列就绕不开这个矩形化矛盾。1.2 默认collate_fn的实现逻辑为了把问题看得更透拆一下default_collate的简化版逻辑def default_collate(batch): # batch是一个list每个元素是Dataset.__getitem__返回的一个tuple # 按字段位置合并假设每个样本是 (x, y)这里先对x字段处理 elem batch[0] if torch.is_tensor(elem): return torch.stack(batch, 0) ...它把batch里所有样本的同一个字段收集起来然后stack。同理如果你的Dataset返回的是(x, length, y)这样的三元组它会对x、length、y三个字段分别执行stack。注意这里第二个字段length如果每个样本返回的是Python intdefault_collate会尝试把它转成tensor然后stack。这个倒是能成功因为标量形状一致但x是变长tensor的时候第一关就过不去。1.3 为什么很多人前期没踩到这个坑很多示例代码和课程项目里数据在进入DataLoader之前就已经被窗口化切分成了固定长度。比如取“过去30天每天的气温、降雨、流量”作为输入预测“未来一天流量”那么所有样本天然就是[30, F]stack毫无压力。这种预处理掩盖了变长问题也让很多人第一次遇到真实变长数据时毫无头绪。实际生产环境里变长几乎躲不掉。文本分类里句子长度差异巨大语音特征按帧数变化水文站记录长度各异事件序列天生不齐。一旦你从“玩具固定长度”切到“真实变长数据”第一个跳出来的就是DataLoader这个报错。所以这不是一个冷门边角问题而是每个做序列模型的人都要过的坎。2. padding之后为什么还不能直接喂LSTM信息污染问题2.1 最朴素的解决方案pad到相同长度既然stack要求形状一致最直接的办法就是把短的序列后面补0补到batch内最长序列的长度。这样一个batch内的样本都是[max_len, F]stack顺利通过。具体实现也简单可以在Dataset的__getitem__里预处理也可以放在collate_fn里动态处理。很多初学者也是这么干的但随后会发现模型指标不对劲或者训练不收敛。问题出在哪2.2 LSTM会把pad位置当真实数据算RNN系模型的递推是强制按时间步执行的。假定一条真实长度是118的样本被pad到142那么在t119到t142这24个时间步上输入是0向量但LSTM的隐状态依然在更新而且更新受前一步隐状态和偏置影响输出不会是0。更麻烦的是信息的“传染”。LSTM在t118这个真实最后一步算出的隐状态按理说已经包含了整条序列的完整信息。但如果你直接把pad后的序列喂进去模型在t119看到输入0它并不知道这个时间步不存在于是基于t118的隐状态和当前输入0再算一次更新。这个被污染的隐状态又继续向后传。你后续无论是取最后一个时间步输出还是用hn[-1]做分类拿到的都是被pad噪声污染过的向量。有人会想那我用mask把loss屏蔽掉pad位置不就行了这是一个常见的误解。mask确实能让loss不计算pad位置的误差但LSTM内部的前向递推已经跑过了那些时间步隐状态已经被“脏数据”改写。mask只能管损失管不了递推污染。所以正确做法是在进入LSTM之前就把pad位置物理上拿掉。2.3 pack_padded_sequence到底做了什么PyTorch提供的答案就是torch.nn.utils.rnn.pack_padded_sequence。它的核心思想是把变长数据按“有效时间步”压缩成一维序列让LSTM在每一步只处理“还没结束的样本”。我用一个直观例子说明。假设batch里有三个样本真实长度分别是3、2、1pad后的形状是[3, 4, F]时间步1三个样本都有效计算batch_sizes[0]3时间步2前两个样本有效计算batch_sizes[1]2时间步3只有第一个样本有效计算batch_sizes[2]1时间步4没人有效不需要计算pack之后内部data的形状是[sum(lengths), F]而不是[batch, max_len, F]也就是[321, F][6, F]。PyTorch的LSTM在底层会按照batch_sizes逐个时间步处理每步只处理batch_sizes[i]个样本长度短的样本跑完就退出计算。这种设计的收益体现在两方面。第一是语义正确pad位置不会再参与任何递推隐状态不会被污染第二是计算效率LSTM的计算量从batch*max_len降到了sum(lengths)。当序列长短差异悬殊时这个省下来的算力非常可观——比如平均长度只有最大长度的一半直接省一半计算量。2.4 enforce_sorted参数和排序要求pack_padded_sequence有一个关键参数enforce_sorted。当它为True默认值时PyTorch要求传入的lengths必须已经按降序排列也就是说batch里的样本必须按长度从长到短排好。这是为了底层实现能直接按连续内存切分数据不需要额外索引操作效率最高。如果数据没排序你有两个选择手动按lengths降序对batch内样本重新排列然后把排序后的lengths传给pack设置enforce_sortedFalse让PyTorch内部自动排序我强烈建议选第一种。手动sort只多几行代码但能省掉PyTorch内部的重复排序开销而且你自己清楚batch内样本的排列状态。选了enforce_sortedFalse之后返回的PackedSequence里会带一个unsorted_indices字段想还原原始顺序很容易搞混。与其和这个索引纠缠不如一开始就自己排好。3. 完整改造自定义collate_fn 排序 pack3.1 Dataset端就要把长度信息带出来整个过程的第一步是在Dataset的__getitem__里返回序列长度的信息。这一步很多人会漏掉因为以前固定长度时压根不需要length。现在必须让DataLoader知道每条序列的真实长度否则collate_fn里无从得知该pad到多少、pack时该传什么。下面是一个水文径流预测场景的最小Dataset实现sequences是一个列表每个元素是形状[L_i, F]的tensorL_i各不相同import torch from torch.utils.data import Dataset class StreamflowDataset(Dataset): def __init__(self, sequences, targets): # sequences: List[Tensor]每个Tensor形状为 [L_i, F] # targets: List[float] 或 List[int] self.sequences sequences self.targets targets def __len__(self): return len(self.sequences) def __getitem__(self, idx): x self.sequences[idx] y self.targets[idx] return x, torch.tensor(len(x)), torch.tensor(y)注意length这里我直接用torch.tensor(len(x))包了一下返回的是0维tensor。这样default_collate在stack标量的时候不会出问题。其实Python int也行但统一成tensor在后续处理时更顺手。3.2 collate_fn的完整实现接下来是核心的collate_fn。它负责三件事把batch内所有序列pad到当前batch最大长度把lengths和targets转成tensor按lengths降序对padded、lengths、targets统一排序。def collate_fn(batch): seqs, lengths, targets zip(*batch) max_len max(lengths) feat_dim seqs[0].size(-1) batch_size len(seqs) # 动态pad到当前batch的最大长度 padded torch.zeros(batch_size, max_len, feat_dim) for i, (seq, length) in enumerate(zip(seqs, lengths)): padded[i, :length] seq lengths torch.tensor(lengths) targets torch.tensor(targets) # 按长度降序排序pack_padded_sequence的强制要求 sorted_lengths, sort_idx torch.sort(lengths, descendingTrue) padded padded[sort_idx] targets targets[sort_idx] return padded, sorted_lengths, targets这里有一个极其容易踩的坑排序时padded、targets必须跟着lengths一起按sort_idx重排否则输入和标签就对不上了。我就见过有人只对padded做了索引targets没动结果训练loss忽高忽低模型怎么调都不收敛查了一整天才发现是标签乱序。另外pad操作我直接用切片赋值padded[i, :length] seq这个写法在seq是tensor时是in-place拷贝效率没问题。不要用cat或者append去拼那样又会引入新的形状问题。3.3 最小可运行的LSTM分类示例有了dataset和collate_fn整个训练循环可以这样写。我特意把batch_size设为2方便你观察变长数据在里面的流转import torch import torch.nn as nn from torch.utils.data import DataLoader from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence # 伪造变长序列数据5条长度不齐特征维度1 raw_x [torch.randn(l, 1) for l in [142, 118, 200, 96, 160]] raw_y [1, 0, 1, 0, 1] # 二分类标签 class LSTMClassifier(nn.Module): def __init__(self, input_size, hidden_size, num_layers1): super().__init__() self.lstm nn.LSTM(input_size, hidden_size, num_layers, batch_firstTrue) self.fc nn.Linear(hidden_size, 2) def forward(self, packed_input): packed_output, (hn, cn) self.lstm(packed_input) # hn的形状是 [num_layers, batch, hidden_size] # 取最后一层最后一个有效时间步的隐状态作为整条序列的表示 last_hidden hn[-1] return self.fc(last_hidden) dataset StreamflowDataset(raw_x, raw_y) dataloader DataLoader(dataset, batch_size2, shuffleTrue, collate_fncollate_fn) model LSTMClassifier(input_size1, hidden_size16) optimizer torch.optim.Adam(model.parameters(), lr1e-3) criterion nn.CrossEntropyLoss() for epoch in range(3): for padded, lengths, targets in dataloader: # lengths留在CPUpack_padded_sequence不支持GPU上的lengths packed_input pack_padded_sequence( padded, lengths, batch_firstTrue, enforce_sortedTrue ) logits model(packed_input) loss criterion(logits, targets) optimizer.zero_grad() loss.backward() optimizer.step() print(fepoch {epoch}, loss {loss.item():.4f})这段代码可以直接跑通。核心就两行collate_fn里利用动态pad排序统一形状训练时用pack_padded_sequence把pad位置压缩掉LSTM前向只计算有效时间步。3.4 device相关的坑lengths必须留在CPU我第一次写这个流程时习惯性地在collate_fn的最后把所有tensor都.to(device)结果pack_padded_sequence直接报出一堆关于lengths的设备不匹配的错误。后来查文档才知道pack_padded_sequence要求lengths是CPU上的tensor而padded可以放在GPU上。所以正确做法是collate_fn里不要动device把padded、lengths、targets原样返回在训练循环里padded和targets再显式.to(device)lengths保持CPU不动。这一点如果你用多卡DataParallel/DDP还要留意同样的问题lengths在打包前别挪设备。另外pad_packed_sequence返回的output会回到和输入padded相同的设备上所以后续接全连接层、算loss都不需要额外操心设备位置。唯一要注意的是如果你要手动构造mask就得把lengths转回目标设备因为mask通常是在GPU上做的。4. 输出、loss与预测阶段变长数据的三处隐藏bug4.1 不要用output[:, -1, :]提取最后一步把pack_padded_sequence的输出经过pad_packed_sequence还原后output形状是[batch, max_len, hidden_size]。在固定长度数据的时代取最后一步通常写output[:, -1, :]。但在变长数据里这个操作取到的是max_len位置对短序列来说那是padding位值可能是0也可能是一堆无意义的隐状态。正确做法有两个# 方法一直接用pack后的hn[-1]语义是最后一个有效时间步的隐状态 last_hidden hn[-1] # 方法二手动按lengths索引还原后的output batch_indices torch.arange(output.size(0), deviceoutput.device) last_hidden output[batch_indices, lengths - 1, :]注意lengths在这里必须在output的设备上如果output在GPUlengths还是CPU的tensor就需要先lengths.to(device)。方法二多一步设备转换方法一省事很多我平时基本都用hn[-1]。需要提醒的是hn[-1]只在单层LSTM时等价于“最后一个有效时间步的隐状态”。多层LSTM时hn[-1]取的是最后一层每个样本最后一个有效时间步的隐状态这正是分类/回归任务最常用的语义。如果你要的是每一层各自的结果那就按层索引。4.2 mask losspad位置不能参与损失计算如果你的任务不是序列级分类而是对每个时间步做预测比如逐日径流预报、序列标注、语音帧分类那么loss计算时pad位置的预测结果必须屏蔽掉。否则模型会拼了命去把padding位置预测成某个默认值把整个优化方向带偏。mask的构造很简单核心是一个形状比较# seq_len: 当前batch的max_len # lengths: 每个样本真实长度注意先转到和设备一致 mask torch.arange(seq_len, deviceoutput.device).unsqueeze(0) lengths.unsqueeze(1).to(output.device) # mask形状 [batch, seq_len]True表示该位置有效以逐时间步回归任务为例masked MSE loss可以这样写# pred_seq: [batch, seq_len, 1] # target_seq: [batch, seq_len, 1] mask mask.unsqueeze(-1) # [batch, seq_len, 1] loss ((pred_seq - target_seq) ** 2 * mask).sum() / mask.sum()这行代码里有一个细节除以的是mask.sum()而不是整个batch的样本数或所有元素数。含义是只对有效时间步的误差求平均这样不管batch里序列长短差异多大loss的尺度不会忽大忽小。如果是分类问题更简单的做法是直接用CrossEntropyLoss的ignore_index参数criterion nn.CrossEntropyLoss(ignore_index0) # 0是pad token id它内部会自动忽略target中等于ignore_index的位置效果和手动mask一样但省去写mask矩阵的麻烦。缺点是你得保证pad位置的标签恰好是你指定的ignore_index如果标签本身就是0就需要另选一个特殊值或者手动mask。4.3 预测阶段的处理方式和训练不一样上线推理的时候如果一次只预测一条样本完全不需要pack。直接把输入形状组织成[1, L, F]喂给LSTM即可PyTorch不会强迫你pad到某个统一长度。这一点很多人训练时用了pack预测时却把流程写复杂了。如果是批量预测那就和训练时保持一致动态pad、排序、pack。唯一可以省略的是loss相关的mask因为预测阶段没有标签只有前向过程。另外一个实践建议训练好模型导出时尽量固定一个最大输入长度。不是因为模型限制而是方便服务端做batch预测和显存规划。把超长序列截断、短序列pad到固定长度在线推理的吞吐往往比每次都动态pad要高因为省去了排序和变长逻辑的调度开销。4.4 标签也是变长encoder-decoder场景标题里说“lstm等”实际碰到的还有一类是输入变长输出也变长典型如seq2seq、语音识别、文本生成。这种情况下decoder的target序列同样需要pad到相同长度且loss计算时要把“结束符之后的pad位置”全部屏蔽。做法和上面基本一致对target做pad时pad位置填入一个特殊token id通常是0然后用CrossEntropyLoss(ignore_index0)直接忽略pad位或手动构造target mask乘到loss上这里额外提醒一个坑decoder端做teacher forcing时输入是target序列整体右移一位pad位置会跟着移动mask也要相应地在时间维度上平移。很多人只记得给target加mask忘了给decoder输入也加对应的mask导致模型在训练时“偷看”了pad位置的信息。5. 性能优化和工程建议变长数据的高级玩法5.1 三种padding方案对比聊完正确性再聊性能和工程。同样是处理变长方案不同内存占用和训练速度差别很大。方案做法内存占用实现复杂度适用场景全局pad预处理时pad到整个数据集最大长度最浪费可能80%以上是0最低数据长度差异不大或固定窗口场景batch内动态padcollate_fn里pad到当前batch最大长度合理只浪费当前batch内部中等大多数变长序列任务推荐默认分桶动态pad先按长度分桶再在桶内组batch最优padding比例极低较高序列长短差异悬殊训练速度优先全局pad的问题在第一节已经提过如果最长序列有1000个时间步平均只有200你每一步训练都在算80%的无效数据显存和算力双重浪费。batch内动态pad是我最推荐的做法它不需要任何预处理成本只在collate_fn里多几行代码已经在第三节实现了。5.2 一个简单的LengthBucketBatchSampler当序列长短差异极大时batch内动态pad也会出现浪费。比如一个batch里最长序列1000步其他几条只有50步那么这个batch的padding比例高达90%。分桶的思路是先把所有样本按长度排序切成若干个长度区间然后在每个区间内随机组batch。这样每个batch内部的长度差异很小自然就减少了padding。一个可以直接抄的版本import numpy as np from torch.utils.data import Sampler class LengthBucketBatchSampler(Sampler): def __init__(self, lengths, batch_size, shuffleTrue): self.batch_size batch_size self.shuffle shuffle sorted_idx sorted(range(len(lengths)), keylambda i: lengths[i]) # 桶的数量取sqrt(n)经验值可根据数据分布调整 n_buckets max(1, int(np.sqrt(len(lengths)))) self.buckets np.array_split(sorted_idx, n_buckets) def __iter__(self): for bucket in self.buckets: bucket list(bucket) if self.shuffle: np.random.shuffle(bucket) for i in range(0, len(bucket), self.batch_size): yield bucket[i:i self.batch_size] def __len__(self): return sum( (len(bucket) self.batch_size - 1) // self.batch_size for bucket in self.buckets )用法上要注意用了batch_sampler之后DataLoader里就不能再传batch_size和shufflelengths [x.size(0) for x in raw_x] dataloader DataLoader( dataset, batch_samplerLengthBucketBatchSampler(lengths, batch_size2), collate_fncollate_fn, )我在气象降水数据上做过一次对比序列长度从几十到上千不等不分桶训练一个epoch耗时10分钟分桶后不到6分钟提速接近40%。代价是每个epoch内的样本顺序不再完全随机如果比较在意随机性可以每隔几个epoch重新按长度排序再分桶让长短样本有机会在同一个batch里出现。5.3 监控padding比例无论用哪种方案我都建议在训练代码里加一行padding比例的统计用来判断当前的数据组织方式是否合理。def compute_pad_ratio(padded, lengths): total_elements padded.numel() valid_elements lengths.sum().item() * padded.size(-1) return 1.0 - valid_elements / total_elements如果这个值长期超过50%说明你的batch里有一半以上的计算都在处理pad位置这时候就应该考虑分桶或者调整batch_size让序列长度相近的样本更多出现在同一个batch里。这个指标也适合用来横向对比不同采样策略的效果。5.4 另一种“长度不一致”多特征时间轴不对齐最后聊一个标题里“等”字涵盖的延伸场景。除了样本之间的长度不一致还有一种更隐蔽的不一致同一个样本内部多个特征在时间轴上对不齐。做水文径流预报时很典型降雨量数据可能是逐小时网格产品上游流量站是逐日观测遥感蒸散发产品又是8天合成。它们的时间网格不同直接拼成一个特征矩阵会错位。这种问题不在DataLoader层解决而是在预处理阶段就要统一时间基准把所有特征重采样到同一个时间网格或者对缺失时段做插值。千万不要试图直接在collate_fn里对不同时间步的特征做对齐那会引入未来信息泄漏。另一种思路是分多路encoder每个分支处理不同时间分辨率的输入最后把各自的隐状态拼接起来再交给后续层。这条路实现成本高一些但避免了插值带来的信息损失适合特征本身时间尺度差异过大、且数据量充足的场景。5.5 我建议你写进代码注释里的几个结论这里把我平时实际项目里沉淀下来的几条规则列一下都是吃过亏之后总结的只要数据是变长的就必须在Dataset里返回lengths不要试图在模型内部猜长度pack之前必须排序排序时必须同时重排padded、lengths、targets三者永远保持一致pack_padded_sequence的lengths必须留在CPUpadded可以在GPU不要统一to(device)序列级任务取表示用hn[-1]不要用output[:, -1, :]时间步级任务要显式构造mask或者在loss里用ignore_index在线推理单条样本不需要pack保持输入维度[1, L, F]即可训练前先算一算padding比例超过50%优先考虑分桶batch_sampler第四节的mask问题和第五节的排序问题是两个最容易出“软bug”的地方。软bug是指代码能跑、loss能降、但最终结果就是不如预期排查起来非常恶心。尤其排序不同步训练曲线看起来正常可一旦跑验证集就露馅因为验证时样本顺序变了标签对应关系混乱。我自己现在写任何带变长序列的DataLoader第一件事就是在纸上把所有tensor的shape演算一遍pack前是什么形状、pack后是什么形状、pad_packed_sequence还原后是什么形状、hn[-1]是什么形状、mask是什么形状。这个习惯看着笨但真的帮我少踩了很多次坑。变长数据的处理链路其实不长坑就藏在那些“看起来显然是对的”地方。
返回列表