跳到主要内容

数据截至 (上游 commit cd6a3406572e)

01 · Arrow 列式存储与内存映射

这一章讲什么: 整个库的地基——为什么一个 100 GB 的数据集可以"打开"在 8 GB 内存的机器上,还能 ds[12345] 随机访问。读完你会明白 Dataset._data 到底是什么、pickle 一个 Dataset 为什么几乎不耗内存、以及拼接/切片为什么大多是零拷贝的。


1. 它要解决的小问题

训练数据常常比内存大。朴素做法有两条死路:

  • 全读进内存(pandas 式):500 GB 语料直接爆掉。
  • 自己写偏移索引逐行 seek:能做,但每种格式都要写一遍,而且读完还要把字节解析成对象,慢。

真正想要的是:数据待在磁盘上,访问起来却像在内存里,而且读出来的直接就是类型化数组,不需要逐行反序列化。


2. 思路/直觉

两个成熟技术拼起来就是答案:

  • 内存映射(mmap):系统调用把文件映射到进程的虚拟地址空间。读某个地址时,如果对应页不在物理内存,OS 触发缺页中断把那一页从磁盘换进来;内存紧张时 OS 自动把不用的页踢掉。换页策略是操作系统几十年打磨好的,不用自己写。
  • Apache Arrow 列式格式:同列数据在文件里是连续、定长对齐的二进制缓冲区。Arrow 的 IPC 流格式(pa.ipc.open_stream)支持直接从 memory-mapped 字节上"长出"数组,不做拷贝、不做逐行解析——这就是 zero-copy 的字面意思。

两件套合一:磁盘文件 ≈ 内存数组。RAM 只是 OS 页缓存,大小与数据集无关。


3. 图示:Dataset 的存储结构

怎么读这张图: 左边是用户面对的 Dataset;右边是它内部的两个表——主表 _data 必有,索引映射 _indices 可选(第 2 章讲)。主表通常是 MemoryMappedTable,指向磁盘上的 .arrow 文件。

Dataset (arrow_dataset.py:718)

├── _data: MemoryMappedTable ──────────► ~/.cache/huggingface/datasets/.../*.arrow
│ (mmap 视图,pickle 只存 path+replays) (Arrow IPC 流格式,数据本体)

├── _indices: MemoryMappedTable ───────► cache-<指纹>.arrow
│ (可选,单列 uint64 索引映射) (只有一列 "indices")

└── _fingerprint: "9f3a2c..." ◄── 同时写进文件 schema 的 metadata

三种表实现构成一个小的类层级:

Table (组合包装 pa.Table)
└── TableBlock
├── InMemoryTable 数据在 RAM(pickle = 复制全部数据)
└── MemoryMappedTable 数据在磁盘(pickle = 只存路径+replays)
ConcatenationTable 若干 TableBlock 的逻辑拼接,不复制数据

Table组合而非继承包住 pa.Table(src/datasets/table.py:210-224),子类的差别集中在 slice/filter/cast 等变换方法的行为上。


4. 原理演示

下面这段把这个机制的骨架演出来(真实实现见 §5):

# 示意,非源码
import pyarrow as pa

class MemoryMappedTable:
def __init__(self, path):
self.path = path
self.replays = [] # 变换历史,先空着
self.table = self._load()

def _load(self):
stream = pa.memory_map(self.path) # ① mmap:文件变成虚拟内存
reader = pa.ipc.open_stream(stream) # ② IPC 流:直接从映射字节读
table = reader.read_all() # ③ 零拷贝长出 pa.Table
for name, args, kwargs in self.replays:
table = getattr(table, name)(*args, **kwargs) # ④ 重放变换
return table

def slice(self, offset, length): # 变换不修改文件
new = MemoryMappedTable.__new__(MemoryMappedTable) # 只记录一条 replay
new.path, new.replays = self.path, self.replays + [("slice", (offset, length), {})]
new.table = self.table.slice(offset, length)
return new

def __getstate__(self): # ⑤ pickle 时:数据不进流
return {"path": self.path, "replays": self.replays}

重点看 ⑤:序列化出去的不是数据,而是"怎么重新得到这份数据"的配方——路径加一串操作。这就是整个库多进程传递 Dataset 不爆内存的原因。


5. 真实实现

5.1 入口:pa.memory_map 的三行封装

所有魔法从这三行开始(_memory_mapped_record_batch_reader_from_file,src/datasets/table.py:47-49):

def _memory_mapped_record_batch_reader_from_file(filename: str) -> pa.RecordBatchStreamReader:
memory_mapped_stream = pa.memory_map(filename)
return pa.ipc.open_stream(memory_mapped_stream)

上面包一层 read_all() 就是 _memory_mapped_arrow_table_from_file(src/datasets/table.py:119-122)。读表时走 ArrowReader.read_table(src/datasets/arrow_reader.py:317-329):in_memory=False(默认)返回 MemoryMappedTable.from_file,True 才返回 InMemoryTable.from_file

5.2 MemoryMappedTable:pickle 只存"配方"

类docstring 自己把设计说透了(src/datasets/table.py:1055-1073):pickle 时不复制数据,只存文件路径 + 一个变换列表(replays),重新加载时把变换重放一遍。

核心三个方法(src/datasets/table.py:1087-1107):

def __getstate__(self):
return {"path": self.path, "replays": self.replays}

def __setstate__(self, state):
...
table = _memory_mapped_arrow_table_from_file(path)
table = self._apply_replays(table, replays) # 重放 slice/cast/filter/...

每个变换方法都遵循同一个模式——先追加 replay,再返回新对象。以 slice 为例(src/datasets/table.py:1114-1131):

replay = ("slice", (offset, length), {})
replays = self._append_replay(replay)
return MemoryMappedTable(self.fast_slice(offset=offset, length=length), self.path, replays)

注意两点:一是立即生效(返回的新表已经是切完的),replay 只是给"未来反序列化"用的;二是 slice 用的是自家 fast_slice(见 §5.4),不是 pyarrow 原生切片。

5.3 InMemoryTable:对照组

InMemoryTable(src/datasets/table.py:695)的数据在 RAM 里,pickle 走标准流程复制全部数据。它服务于两种场景:数据本来就小(Dataset.from_dict 之类),或者 keep_in_memory=True 明确要求进内存。load_dataset 有个自动判断:数据集小于 IN_MEMORY_MAX_SIZE(默认 0,即不自动)才进内存(src/datasets/load.py:1730-1732 + is_small_dataset,src/datasets/utils/info_utils.py:93)。

5.4 IndexedTableMixin:跨 record batch 的快速随机访问

Arrow 文件内部是多个 record batch 串起来的。原生 pa.Table.slice 对长表偏慢,于是库自己维护了一个轻量索引(src/datasets/table.py:161-167):

  • _batches:所有非空 record batch 的列表;
  • _offsets:各 batch 起始行号的前缀和(np.cumsum)。

两个加速操作:

  • fast_slice(src/datasets/table.py:186-207):用插值搜索(_interpolation_search,src/datasets/table.py:135-158)直接定位起止 batch,只切首尾两个 batch,中间的整块保留。
  • fast_gather(src/datasets/table.py:169-184):对任意索引列表,用 np.searchsorted 批量算出每个索引落在哪个 batch,再逐行 slice 拼起来——注释里明说这比 pa.concat_tables(table.fast_slice(i,1) for i in indices) 快,因为二分查找被 NumPy 向量化了。

Dataset.__getitem__ 的读路径最终就是走到这里:_getitem(src/datasets/arrow_dataset.py:3125-3143)→ query_table(src/datasets/formatting/formatting.py:594-636)→ 有 _indices 时先经索引表重映射,再查主表。

5.5 ConcatenationTable:拼接不复制

concatenate_datasets 和 num_proc map 后的分片合并都靠它(src/datasets/table.py:1339)。设计要点:

  • blocks二维列表:第一维按行拼接(axis 0),第二维按列拼接(axis 1)(src/datasets/table.py:1332-1356 的注释)。
  • 拼接是逻辑的:blocks 里每个块仍指向各自的磁盘文件;只在需要物化一张 pa.Table 时才 concat(_concat_blocks_horizontally_and_vertically,src/datasets/table.py:1409-1417)。
  • pickle 时每个块各自序列化——内存块复制数据,mmap 块只存路径(src/datasets/table.py:1347-1351)。
  • 列方向拼接要求各 row_block 行数对齐;对不齐就互相切片对齐,_split_both_like(src/datasets/table.py:1485-1520)的 docstring 里画了一张很形象的 ASCII 对齐图。

list_table_cache_files(src/datasets/table.py:1832-1852)递归收集一张表背后所有被 mmap 的文件路径——Dataset.cache_files 属性和指纹计算(把缓存文件的 mtime 也哈希进去,src/datasets/fingerprint.py:243-245)都靠它。

5.6 元数据藏在 schema 里

Arrow 文件的 schema metadata 里塞了一个 b"huggingface" 键,JSON 序列化的 {"info": {"features": ...}, "fingerprint": ...}(ArrowWriter._build_metadata,src/datasets/arrow_writer.py:614-622)。Dataset.__init__ 会把它读出来:指纹不在参数里就从文件里恢复(src/datasets/arrow_dataset.py:745-750),features 也会和文件 schema 对齐(:752-767)。

一个推论: 缓存文件是自描述的。你把它拷给同事,对方 Dataset.from_file 打开,列类型和指纹都在,不需要任何伴随文件。

5.7 取数据的最后一公里:惰性格式化

_getitem 查回 pa_subtable 后交给 format_table 按当前格式(torch/numpy/默认 Python)转换。默认格式下有个容易被忽略的优化:LazyDict(src/datasets/formatting/formatting.py:285-296)——取一行时先返回一个 dict 壳,某列只在被访问的那一刻才从 Arrow 转 Python 对象_map_single 在非批处理且未指定 input_columns 时会打开这个惰性开关(src/datasets/arrow_dataset.py:3757-3765)。


6. 关键细节/坑

  • deepcopy 不复制数据。 Table.__deepcopy__self.tableself._batches 塞进 memo(src/datasets/table.py:226-233),因为 Arrow 表不可变,复制纯属浪费。注释还提到一个 pyarrow 的怪现象:对 pa.Table 做 deepcopy 反而会让 pa.total_allocated_bytes() 下降。
  • replay 链会累积。 对 mmap 表连续做 N 次 slice/cast,pickle 时 N 条变换全在 replays 里,反序列化时逐条重放。库没有做任何 replay 压缩/物化(inferred:代码里看不到压缩逻辑,_append_replay 只是 append,src/datasets/table.py:1109-1112)。长期演化的管道如果频繁切片又频繁 pickle,值得留意。
  • Parquet 读路径总是进内存。 ParquetReader._get_table_from_filename 虽然传了 memory_map=True,但注释写明 "Parquet read_table always loads data in memory"(src/datasets/arrow_reader.py:355-356)——Parquet 是压缩列存,mmap 只能映射压缩后的字节,解压这一跳省不掉。真正的零拷贝只属于 Arrow IPC 格式。
  • 空表切片会 segfault,所以到处有防御。 ArrowReader._get_table_from_filename 特意不切空表(src/datasets/arrow_reader.py:311-313),map 对空数据集直接短路(src/datasets/arrow_dataset.py:3368-3380)。

7. 代码地图

主题文件路径符号名
mmap 读表入口src/datasets/table.py_memory_mapped_record_batch_reader_from_file_memory_mapped_arrow_table_from_file
内存映射表src/datasets/table.pyMemoryMappedTable_apply_replays_append_replay
内存表src/datasets/table.pyInMemoryTable
逻辑拼接表src/datasets/table.pyConcatenationTable_split_both_likeconcat_tables
快速行访问src/datasets/table.pyIndexedTableMixinfast_slicefast_gather_interpolation_search
表背后的缓存文件src/datasets/table.pylist_table_cache_files
按指令读文件src/datasets/arrow_reader.pyArrowReader.read_table_get_table_from_filename
Dataset 初始化与元数据src/datasets/arrow_dataset.pyDataset.__init__
读路径src/datasets/formatting/formatting.pyquery_tableformat_tableLazyDict
文件内元数据src/datasets/arrow_writer.pyArrowWriter._build_metadata