Rerun 实验版 PyTorch DataLoader 缺失样本跳过机制:max_consecutive_skipped_samples语义与源码解析
【免费下载链接】rerunVisualize, query, and stream to train on multimodal robotics data.项目地址: https://gitcode.com/GitHub_Trending/re/rerun
本篇基于 Rerun 0.37 版本说明中“实验版 iterable dataloader 跳过缺失样本”这一条目展开(原草稿路径changelog/upcoming/dataloader-skip-missing-samples.md已并入正式版本说明,见 重定向配置 与 0.37 版本说明),系统讲解RerunIterableDataset在实时(live)训练数据流中如何处理解码失败与缺失字段:哪些情况会被判定为“缺失”、max_consecutive_skipped_samples的计数与上限语义、跳过在数据管道中的精确位置,以及 DDP 分布式训练下必须注意的DistributedDataParallel.join()约束,并给出对应源码与单元测试的可验证依据。
背景:为什么 live 数据流会出现“缺失样本”
Rerun 的实验版 PyTorch DataLoader(文档)面向真实机器人数据采集场景:一个 segment 中的某个字段可能在某些时刻根本没有数据,或者某个字段的压缩视频在某个 GOP 内解码失败。在 0.37 之前,这类情况会直接中断迭代;现在RerunIterableDataset的 live 路径改为跳过这些样本并继续迭代,而不是让整个训练挂掉。
核心设计原则只有一条:只有None表示“缺失”,合法的空张量(zero-sized tensor)必须原样保留。这一区分在 解码结果类型定义 中直接体现:
DecodedResult: TypeAlias = DecodedValue | None解码器返回None表示“这个样本该字段解不出来”,而不是“这个字段是空值”。
哪些情况会被判定为缺失并跳过
按 0.37 版本说明 的条目,live 路径下会解析为None的情况包括:
- 数值窗口行宽不一致:
NumericDecoder对窗口内行宽(width)不一致的样本返回None,使受影响的 live 样本被跳过而不是停止迭代。在 数值解码器源码 中,当 Arrow 列的 list 布局无法展平为统一宽度的数组(flat is None or offsets is None)时即返回None;而 Arrow 解码器 在任意 list 层级存在 null 行时同样返回(None, None)表示布局不支持映射。 - 源端 null 行:字段在源数据中本就是 null 的行。
- 预期的编解码失败:编码图像或压缩视频解码失败会解析为
None,其中视频解码失败被隔离在受影响的 GOP 范围内——同一个 GOP 内其他可正常解码的帧不受牵连。图像侧的实现见 图像解码器,_decode_request对每一处失败分支都显式return None。
_skip_incomplete:跳过逻辑的源码实现
跳过逻辑集中在 可迭代数据集实现 的_skip_incomplete生成器中,其行为与 changelog 条目逐条对应:
def _skip_incomplete( samples: Generator[DecodedSample, None, None], *, max_consecutive_skipped_samples: int | None = _DEFAULT_MAX_CONSECUTIVE_SKIPPED_SAMPLES, ) -> Generator[DecodedSample, None, None]: """Drop live samples with missing fields up to the configured limit.""" ... for sample in samples: missing = [key for key, value in sample.items() if value is None] if not missing: consecutive_skipped = 0 # 有效样本重置连续计数 yield sample continue skipped += 1 consecutive_skipped += 1 for key in missing: skipped_by_field[key] = skipped_by_field.get(key, 0) + 1 set_current_span_attributes({"rerun.dataloader.iter.num_samples_skipped": skipped}) for key in sorted(set(missing) - warned): warnings.warn( f"Skipping samples where field {key!r} has no value. " "Batches stay at full size; the epoch yields fewer samples.", RuntimeWarning, stacklevel=2, ) ...可以逐点核对四个关键行为:
- 按字段判缺失:遍历样本字典,
value is None的字段名进入missing列表;空张量不是None,因此不会被跳过。 - 每个缺失字段只告警一次:
warned集合保证同一字段的生命周期内最多发出一次RuntimeWarning(提示“批次保持满大小,epoch 产出的样本变少”)。 - OpenTelemetry 埋点:每次跳过都会把累计计数写入当前 span 的
rerun.dataloader.iter.num_samples_skipped属性,与 0.37 中“RerunIterableDataset的 OpenTelemetry tracing 覆盖两条 fetch 路径”的改动配套,方便在训练监控中直接看到跳过规模。 - 生成器资源安全:
finally: samples.close()确保消费者提前退出(例如break)时底层 fetch executor 也能正常关闭——这一点有专门的单元测试覆盖(见后文)。
max_consecutive_skipped_samples:连续跳过上限的语义
构造参数文档(源码 docstring)说明:每个 rank 和每个DataLoaderworker 独立计数,有效样本会把连续计数清零;超限后下一个缺失样本会抛出带有总跳过数与按字段拆分的计数的RuntimeError:
raise RuntimeError( f"Exceeded max_consecutive_skipped_samples={max_consecutive_skipped_samples} after " f"encountering {consecutive_skipped} consecutive incomplete samples " f"({skipped} total; missing fields: {field_counts})" )单元测试 test_dataloader_skip_incomplete.py 验证了错误消息的格式,例如max_consecutive_skipped_samples=1时连丢两个样本(一个缺image、两个都缺action与image)会报2 total; missing fields: action=1, image=2。
关于默认值需要注意版本差异:0.37 版本说明 中该条目写的是“defaults to 100”,而当前源码树的常量是:
_DEFAULT_MAX_CONSECUTIVE_SKIPPED_SAMPLES = 1000(见 _iterable_dataset.py),且单元测试test_default_consecutive_skip_limit_is_1000也断言默认上限为 1000。可以推断该默认值在 0.37 发布后的主干中从 100 上调到了 1000,实际取值请以你所安装版本的 dataloader 文档 与源码为准。传None表示不设上限;传负数会直接抛出ValueError(max_consecutive_skipped_samples must be non-negative)。
一个典型的用法示例(结合 dataloader 示例 的场景):
dataset = RerunIterableDataset( source, index="frame_nr", fields=fields, fetch_block_size=128, # 每个 rank / worker 连续跳过超过 50 个不完整样本即抛错 max_consecutive_skipped_samples=50, )跳过发生在管道的哪个位置
changelog 强调“过滤发生在可选的 emission shuffle 之前,不完整样本不会占用其缓冲区”。这一点在 live 路径 _iter_catalog 的管道装配顺序中可以直接验证:
samples = _skip_incomplete( _pipeline_blocks(blocks, fetch=fetch_block, process=process), max_consecutive_skipped_samples=self._max_consecutive_skipped_samples, ) if self._shuffle_buffer is not None: ... samples = self._shuffle_buffer.shuffle(samples, rng=rng) yield from _count_yields(samples)即:fetch → decode → 跳过不完整样本 → 可选的 shuffle 缓冲 → 计数与 yield。把跳过放在 shuffle 缓冲之前意味着不完整样本根本不会进入缓冲池,避免它们在缓冲窗口内“稀释”随机性并浪费内存。管道末尾的_count_yields(源码)还会记录num_samples_yielded与下游拉取间隔(pull gap),与 changelog 中“decode span 在样本 yield 之前结束,训练循环行为不再虚增解码时长”的改动相呼应。
DDP 分布式训练:跳过在分片之后发生
因为跳过发生在rank 分片之后,各 rank 的有效样本数可能不一致——某个 rank 的 shard 里恰好集中了较多缺失样本,它会比别的 rank 更早耗尽。因此 changelog 给出的硬性建议是:有限数据集的 DDP 训练循环必须用DistributedDataParallel.join()包裹,避免先完成的 rank 卡住其他 rank 的集合通信:
with model.join(): # 防止先耗尽的 rank 阻塞其他 rank for batch in loader: ...数据集类 docstring 对此有同样的说明:“rank shards can have different lengths, especially when incomplete samples are skipped after sharding”。
Manifest 回放保持严格模式:不跳过,直接报错
跳过机制只适用于 live 路径。manifest 回放(RerunIterableDataset.from_manifest)走的是另一条校验路径 _raise_if_incomplete:如果某个被 manifest 记录为 required 的字段解码为None,迭代直接抛出错误,提示用户重新生成 manifest,而不是悄悄改变 manifest 冻结下来的采样顺序:
raise RuntimeError( f"Required fields decoded to nothing: {', '.join(sorted(missing))}. The manifest was built " "as against different data, so regenerate it.\n" f"Segment: {target.segment.segment_id} at {target.index_value}" )这个设计保持了 manifest 作为“可复现、可断点续训采样顺序”的契约:live 路径允许丢样本换取鲁棒性,回放路径则用失败换可复现性。单元测试test_manifest_replay_raises_when_a_required_field_is_missing与test_manifest_replay_allows_an_optional_field_to_be_missing分别验证了 required 字段缺失抛错、optional 字段缺失放行两种行为(测试文件)。
另外注意一个边界:RerunMapDataset不能跳缺失项——map 风格数据集的索引映射必须稳定,跳过一个 item 会破坏索引与样本的一一对应,因此它的样本字典中对应字段可能为None(见 dataloader 文档 中关于 rank shard 与跳过的说明)。
单元测试覆盖的行为清单
rerun_py/tests/unit/test_dataloader_skip_incomplete.py 是对该特性最直接的证据,逐条对应 changelog 声明:
| 测试 | 验证的行为 |
|---|---|
test_drops_none_but_keeps_valid_empty_tensors | 丢弃None,但保留合法的空张量样本 |
test_warns_once_per_missing_field | 5 个同字段缺失样本只产生 1 条RuntimeWarning |
test_default_consecutive_skip_limit_is_1000 | 默认连续上限(当前源码为 1000) |
test_allows_exactly_the_configured_number_of_consecutive_skipped_samples | 上限是“恰好允许 N 个”而非 N-1 |
test_valid_sample_resets_consecutive_skip_count | 有效样本重置连续计数 |
test_raises_after_consecutive_skip_limit_with_per_field_counts | 抛错消息含总数与按字段计数 |
test_rejects_a_negative_skip_budget | 负数配置抛ValueError |
test_closes_the_source_when_the_consumer_stops_early | 消费者提前退出时关闭源生成器 |
test_manifest_replay_* | manifest 回放严格模式:required 缺失抛错、optional 放行 |
小结与实践建议
- live 训练:默认开启跳过,字段缺失会按字段各告警一次并计入 span;数据质量明显恶化(例如相机掉线导致连续数百样本缺
image)时,调低max_consecutive_skipped_samples可以让问题尽早以带按字段统计的RuntimeError暴露,而不是悄悄丢光一批数据。 - DDP:有限数据集务必
model.join();各 rank 因跳过导致的长度差是预期行为。 - 可复现训练:优先使用
RerunIterableDataset.from_manifest回放路径——它不做跳过,required 字段缺失会明确要求重新生成 manifest,保证采样顺序与冻结清单一致。 - 空张量不是缺失:如果你的解码逻辑会产生
torch.empty(0)之类的合法空张量,它们会被保留进批次;只有None会触发跳过,这一点在写自定义ColumnDecoder时尤其要注意返回值约定(见 解码器基类 的decode返回语义:result[i]为第i个请求的解码值,None表示该请求无结果)。
【免费下载链接】rerunVisualize, query, and stream to train on multimodal robotics data.项目地址: https://gitcode.com/GitHub_Trending/re/rerun
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考