HuggingFace Datasets中lambda无法被pickle/dill序列化致缓存失效
HuggingFace Datasets map 函数哈希失败、缓存失效问题解决
问题复现
预处理CIFAR100数据集时,使用lambda作为dataset.map()的转换函数,代码如下:
from datasets.load import load_dataset from datasets import Features, Array3D from transformers.models.vit.feature_extraction_vit import ViTFeatureExtractor # Resampling & Normalization feature_extractor = ViTFeatureExtractor.from_pretrained('google/vit-base-patch16-224-in21k') dataset = load_dataset('cifar100', split='train[:100]') features = Features({ 'pixel_values': Array3D(dtype="float32", shape=(3, 224, 224)), **dataset.features, }) dataset = dataset.map(lambda batch, col_name: feature_extractor(batch[col_name]), features=features, fn_kwargs={'col_name': 'img'}, batched=True)
运行后抛出哈希警告,提示转换函数无法被正确序列化,缓存机制失效:
Reusing dataset cifar100 (/home/qys/.cache/huggingface/datasets/cifar100/cifar100/1.0.0/f365c8b725c23e8f0f8d725c3641234d9331cd2f62919d1381d1baa5b3ba3142) Parameter 'function'=<function <lambda> at 0x7f3279f3eef0> of the transform datasets.arrow_dataset.Dataset._map_single couldn't be hashed properly, a random hash was used instead. Make sure your transforms and parameters are serializable with pickle or dill for the dataset fingerprinting and caching to work. If you reuse this transform, the caching mechanism will consider it to be different from the previous calls and recompute everything. This warning is only showed once. Subsequent hashing failures won't be showed.
测试序列化能力时发现,顶层定义的命名函数可以正常被datasets.fingerprint.Hasher计算哈希,但功能完全一致的lambda函数会序列化失败,哪怕将lambda赋值给顶层变量也无法解决,报错为pickle阶段的KeyError。
根本原因
datasets的缓存指纹生成逻辑,依赖dill/pickle对传入map的转换函数、关联参数做全量序列化,序列化结果的哈希值就是缓存的唯一标识。只要序列化失败,就会随机生成哈希值,导致缓存无法命中,每次运行都重算转换逻辑。- lambda属于匿名函数,没有模块级可查找的固定限定名(
__qualname__),跨进程反序列化时无法定位到函数定义;如果lambda是闭包、引用了外层作用域的变量(比如代码里的feature_extractor、测试用例里的foo),dill和datasets自定义的函数序列化逻辑组合时会出现memo表引用查找错误,直接抛出KeyError。 - 哪怕lambda没有引用外部变量,dill对匿名函数的序列化稳定性也远低于顶层命名函数,不适合作为需要长期缓存的转换逻辑传入
map。
解决方案
方案1:使用顶层命名函数(最稳定,生产环境推荐)
放弃lambda写法,将转换逻辑定义为模块顶层可直接访问的普通函数,需要固定的参数用functools.partial绑定(partial对象可正常序列化),示例代码:
from datasets.load import load_dataset from datasets import Features, Array3D from transformers.models.vit.feature_extraction_vit import ViTFeatureExtractor from functools import partial # 所有依赖的对象也定义在顶层可访问位置 feature_extractor = ViTFeatureExtractor.from_pretrained('google/vit-base-patch16-224-in21k') # 顶层定义命名转换函数 def preprocess_batch(batch, col_name, extractor): return extractor(batch[col_name]) if __name__ == "__main__": dataset = load_dataset('cifar100', split='train[:100]') features = Features({ 'pixel_values': Array3D(dtype="float32", shape=(3, 224, 224)), **dataset.features, }) # 绑定固定参数 map_fn = partial(preprocess_batch, col_name="img", extractor=feature_extractor) # 可以提前验证哈希是否正常 # from datasets.fingerprint import Hasher # print(Hasher.hash(map_fn)) dataset = dataset.map(map_fn, features=features, batched=True)
该写法下函数和参数都可以被正常序列化哈希,缓存可以稳定命中,不会再出现警告。
方案2:调试阶段临时禁用缓存
如果只是快速验证逻辑、不需要持久化缓存,可以在map调用时传入参数跳过缓存检查:
dataset = dataset.map(..., load_from_cache_file=False)
该方法会让每次运行都重新执行转换,仅适合小样本调试场景,大数据集下会严重拖慢效率。
方案3:升级依赖修复已知兼容问题
部分老版本datasets、dill存在lambda序列化的已知bug,可以升级到最新稳定版尝试修复:
pip install -U datasets dill transformers
注意升级后闭包形式的lambda依然可能出现哈希不稳定问题,长期使用还是优先选择方案1。
内容的提问来源于stack exchange,提问作者nalzok
相关产品推荐
相关产品推荐

