跳到主要内容

数据截至 (上游 commit cd6a3406572e)

02 · map/filter 与缓存指纹机制

这一章讲什么: HF Datasets 最「用过就回不去」的特性——ds.map(fn) 跑过一次,第二次运行同一个脚本时零重算,直接从磁盘缓存打开结果。读完你会明白缓存键(指纹)是怎么算出来的、什么时候会失手不命中、多进程 map 怎么分片,以及 shuffle/filter 为什么不用搬数据。


1. 它要解决的小问题

预处理很贵:对 1000 万条语料跑分词可能要几十分钟。而开发是迭代的——改配置、改训练脚本、重跑,但预处理函数往往没改

传统做法是自己存中间结果、自己起文件名、自己记「这份文件是哪份代码产出的」。三周后没人记得 train_v3_final_FINAL.arrow 是什么。

想要的是:预处理结果像内容寻址存储一样自动管理——函数没变就复用,函数变了自动作废,全程不需要人管文件名。


2. 思路/直觉

关键洞察:一次预处理 = 一个纯函数。

输出数据 = f(输入数据, 变换函数, 变换参数)

如果「输入数据 + 函数 + 参数」三者都能哈希,那么哈希值就是输出数据的天然主键。第二次运行时主键相同 → 磁盘上已有以它为名的文件 → 直接打开,跳过计算。

于是问题只剩:怎么哈希一个 Python 函数? 答案是 dill——它比标准 pickle 强在能序列化函数的字节码和闭包。函数体改一行,字节码变,指纹变,缓存自动失效。不需要任何手动版本号。


3. 图示:一次 map 的缓存决策

怎么读这张图: 从左到右是 ds.map(fn) 的执行顺序;菱形是缓存判定点,命中就直接跳到结尾。

ds.map(fn, num_proc=4)


① 算新指纹:xxhash(旧指纹 + dill(fn) + 参数)
│ fingerprint.py · update_fingerprint

② 推出缓存文件名:cache-<新指纹>.arrow
│ arrow_dataset.py · _get_cache_file_path

③ 文件已存在?──── 是 ──► ④ Dataset.from_file(缓存文件)
│ 否 内存映射打开,收工

⑤ 切 4 个 shard → mp.Pool 并行跑 _map_single
│ 每个分片写临时文件,成功后原子改名

⑥ 拼接分片 → 新 Dataset(指纹 = 新指纹)

4. 原理演示

# 示意,非源码
import xxhash, dill

def update_fingerprint(old_fp, transform, transform_args):
h = xxhash.xxh64()
h.update(old_fp.encode())
h.update(dill.dumps(transform)) # 函数本体进哈希:改一行代码 = 新键
for key in sorted(transform_args): # 参数按键名排序,顺序无关
h.update(key.encode())
h.update(dill.dumps(transform_args[key]))
return h.hexdigest()

def map(dataset, fn, **kwargs):
new_fp = update_fingerprint(dataset.fingerprint, fn, kwargs)
cache_file = f"cache-{new_fp}.arrow" # 内容寻址:文件名即主键
if os.path.exists(cache_file): # 命中:直接打开
return Dataset.from_file(cache_file)
... # 未命中:真算,写临时文件
os.rename(tmp_file, cache_file) # 原子落盘
return Dataset.from_file(cache_file)

重点看:缓存系统里没有注册表、没有数据库——文件名就是键,文件系统就是存储。这是它能十几年不出大乱子的原因。


5. 真实实现

5.1 指纹的原料:xxhash + dill

Hasher(src/datasets/fingerprint.py:196-223)内部是 xxhash.xxh64(非加密哈希,极快);hash(value) 先用 dumps 把任意 Python 对象序列化成字节再喂进去(:213-214)。dumps 来自 src/datasets/utils/_dill.py,是对 dill 的封装,里面给 regexspacytiktokentorchtransformers 等常见库的不可 pickle 对象注册了专门的序列化器(src/datasets/utils/_dill.py:29-63)。

每个值进哈希前还会加一个 =={type(value)}== 的类型头(:217-220),防止 1"1" 撞键。

5.2 指纹怎么滚动更新

入口是 update_fingerprint(src/datasets/fingerprint.py:253-300):

hasher = Hasher()
hasher.update(fingerprint) # 旧指纹:链式依赖,输入数据变了指纹就变
hasher.update(transform) # 函数本体
for key in sorted(transform_args):
hasher.update(key); hasher.update(transform_args[key])

初始指纹从哪来?两个来源:

  • builder 给的:_get_dataset_fingerprint 把「缓存目录相对路径(含数据集名/配置/版本/哈希)+ split 指令」哈希(src/datasets/builder.py:1105-1111);
  • 现场算的:generate_fingerprint 把 Dataset 的整个 __dict__ 加上各缓存文件的修改时间一起哈希(src/datasets/fingerprint.py:235-246)。

注意「缓存文件 mtime 进指纹」这条:有人在外部改了你的底表文件,指纹自动变,下游缓存作废。

5.3 每个变换方法都戴着指纹装饰器

Dataset.map/filter/select/shuffle 等方法头上都有 @fingerprint_transform(inplace=False, ...)(src/datasets/fingerprint.py:378)。这个装饰器做三件事:

  1. format_kwargs_for_fingerprint(src/datasets/fingerprint.py:333-375)收集要进指纹的参数——并把等于默认值的参数剔掉(:367-375),所以 batched=False 写不写不影响命中。
  2. update_fingerprint 算出 new_fingerprint,塞进被装饰方法的参数里。
  3. 对带随机性的方法(shuffletrain_test_split)有个贴心设计:即使用户不传 seed,也会从 numpy 当前状态抓一个种子放进指纹(randomized_function=True,src/datasets/fingerprint.py:361-365)——于是两次「无种子 shuffle」产生不同缓存文件,互不覆盖。

方法自己还带版本号,实现变了可以手动炸掉旧缓存:filter 头上是 version="2.0.1"(src/datasets/arrow_dataset.py:4140-4142)。

5.4 缓存文件名与命中判定

指纹变成文件名靠 _get_cache_file_path(src/datasets/arrow_dataset.py:3203-3211):

if is_caching_enabled() and self.cache_files:
cache_file_name = "cache-" + fingerprint + ".arrow"
cache_directory = os.path.dirname(self.cache_files[0]["filename"])
else:
cache_file_name = "cache-" + generate_random_fingerprint() + ".arrow"
cache_directory = get_temporary_cache_files_directory()

注意 else 分支:缓存关闭(disable_caching(),src/datasets/fingerprint.py:141)或底表不在磁盘上时,缓存文件用随机名写进临时目录,session 结束被清理——所以 keep_in_memory=True 的 Dataset 做 map,结果不会被复用,这是文档里不明说但代码里很清楚的行为。

命中判定在 map 内部的 load_processed_shard_from_cache(src/datasets/arrow_dataset.py:3461-3470):文件存在且 load_from_cache_file 为真,直接 Dataset.from_file 打开,抛 NonExistentDatasetError 才往下算。

5.5 多进程:先按指纹分片,进程数还能「迁就缓存」

num_proc=4 时(src/datasets/arrow_dataset.py:3559-3573):数据集先 shard(num_shards=4, contiguous=True) 切成 4 片,每片带自己的缓存文件名(..._00000_of_00004.arrow 后缀),丢进 mp.PoolDataset._map_single(:3619-3630),最后 _concatenate_map_style_datasets 拼起来。

这里藏着一个非常实用的设计(:3488-3502):上次用 num_proc=8 缓存过、这次只给 num_proc=4,它不会傻乎乎重算——先 glob 已有缓存文件,按「缺失比例 → 与当前 num_proc 的差距 → 分片数」排序,直接沿用旧的分片数打开。

另外两个工程细节:

  • fork 子进程前强制 TOKENIZERS_PARALLELISM=false(:3536-3550),防止 tokenizers 的 Rust 线程池在 fork 后死锁(注释里给了 tokenizers 仓库的出处)。
  • 每个分片先写 NamedTemporaryFile,成功后 shutil.move 到正式缓存路径(_map_single,src/datasets/arrow_dataset.py:3901-3915:4044-4049)——中途 Ctrl-C 只会留下一个没人引用的临时文件,下次不会误命中半个缓存。

5.6 filter / select / shuffle:不搬数据的索引映射

这是整套机制里最省的一次设计复用。filter 根本不产新数据表,它的实现是把函数包成「输出保留行的行号」,然后调 self.map(...)(src/datasets/arrow_dataset.py:4252-4286):

indices = self.map(
function=partial(get_indices_from_mask_function, function, ...),
features=Features({"indices": Value("uint64")}), # 输出只有一列行号
remove_columns=self.column_names, # 原列全扔掉
...
)
new_dataset = copy.deepcopy(self)
new_dataset._indices = indices.data # 索引表盖在原表上

select/shuffle 同理:shuffle 就是 permutation = generator.permutation(len(self)) 然后调 self.select(permutation)(src/datasets/arrow_dataset.py:4954-4963);select连续区间_select_contiguous 直接切片,非连续才写索引表(_select_with_indices_mapping,src/datasets/arrow_dataset.py:4514-4617)。索引表若已有索引会先组合(self._indices.column(0).take(indices_array),:4586-4591),所以链式 select 不会堆叠多层间接。

读的时候:query_table 看到 _indices 非空,先拿索引表把请求的 key 重映射成物理行号再查主表(src/datasets/formatting/formatting.py:594-636)。

5.7 代价与回头路

索引映射的代价写在 shuffle 的 docstring 里(src/datasets/arrow_dataset.py:4851-4855):随机访问可以慢 10 倍——每次读先查一层索引,而且不再连续读磁盘。回头路是 flatten_indices(src/datasets/arrow_dataset.py:4287):按索引把数据真物化成一张新的连续 Arrow 表。docstring 还给了另一条路:转 IterableDataset,用 buffer shuffle 保持高速(见第 3 章)。


6. 关键细节/坑

  • 函数序列化失败 → 指纹退化为随机值,缓存永远 miss。 update_fingerprint 捕获所有异常后 generate_random_fingerprint()(src/datasets/fingerprint.py:257-276),且警告只打一次。症状是「map 每次都重算」,去日志里找那条 warning。
  • 全局变量不进指纹。 函数闭包外的全局状态(比如一个全局的正则表)改了,函数字节码没变,指纹不变,缓存照样命中旧结果(inferred:指纹只覆盖函数本体与显式参数,update_fingerprint 没有抓取全局命名空间的逻辑)。要安全就把依赖通过 fn_kwargs 显式传进去——fn_kwargs 是进指纹的。
  • load_from_cache_file=False 只是不读,不是不写。 结果仍写到同一个指纹路径并覆盖(inferred:写入路径不检查该开关,_map_single 无条件写 cache_file_name)。
  • num_proc 大于数据集行数会被压回行数(src/datasets/arrow_dataset.py:3413-3417),别指望空数据集上分进程。
  • 异步函数走另一条路: function 是协程时 _map_single 用 asyncio 并发,上限 MAX_NUM_RUNNING_ASYNC_MAP_FUNCTIONS_IN_PARALLEL = 1000(src/datasets/arrow_dataset.py:3927-3948 + src/datasets/config.py:265),filter 遇到异步函数会强制 batch_size=1(src/datasets/arrow_dataset.py:4248-4250)。
  • 磁盘缓存会越攒越多。 每个新指纹一个新文件,旧的不会自动删;清理入口是 Dataset.cleanup_cache_files()(src/datasets/arrow_dataset.py:3166-3201),它只认 cache-*.arrow 命名。

7. 代码地图

主题文件路径符号名
哈希器src/datasets/fingerprint.pyHasherhash_bytes
指纹滚动更新src/datasets/fingerprint.pyupdate_fingerprintgenerate_fingerprintgenerate_random_fingerprint
变换装饰器src/datasets/fingerprint.pyfingerprint_transformformat_kwargs_for_fingerprintformat_transform_for_fingerprint
缓存开关src/datasets/fingerprint.pyenable_cachingdisable_cachingis_caching_enabled_TempCacheDir
dill 封装src/datasets/utils/_dill.pydumps
map 主流程src/datasets/arrow_dataset.pyDataset.map_get_cache_file_pathload_processed_shard_from_cache
单分片执行src/datasets/arrow_dataset.pyDataset._map_single
filter=map 出索引src/datasets/arrow_dataset.pyDataset.filterget_indices_from_mask_function
索引映射src/datasets/arrow_dataset.pyDataset.select_select_with_indices_mapping_new_dataset_with_indicesflatten_indices
shufflesrc/datasets/arrow_dataset.pyDataset.shuffle
builder 侧初始指纹src/datasets/builder.pyDatasetBuilder._get_dataset_fingerprint