如何将Hugging Face数据集的Sequence(Image)直接转为Array4D?
解决Hugging Face数据集Sequence(Image)转Array4D的性能问题
问题背景
现有Hugging Face数据集,ImageData列是Sequence(Image)类型(固定长度16的图片序列),需要转换为PyTorch标准4D张量(形状(V, C, H, W))。常规方法通过map处理后,该列会生成嵌套的float序列,存储与处理速度极慢——1000个样本在计算集群上需耗时数小时,核心原因是嵌套序列的构建为单线程操作,且集群单核性能较差。需要直接将该列转换为Array4D格式,完全规避嵌套序列带来的性能开销。
解决方案
方法1:直接用cast_column自定义转换
跳过set_format和map的中间步骤,直接将Sequence(Image)列转换为Array4D,通过自定义函数一次性完成图片序列到4D数组的转换:
import datasets import PIL.Image import numpy as np V = 16 H, W, C = 244, 244, 3 def get_ds(): """外部提供的数据集生成函数""" N = 10 data = [ {"ImageData": [PIL.Image.new("RGB", (W, H)) for _ in range(V)]} for _ in range(N) ] ds = datasets.Dataset.from_list(data) ds = ds.cast( datasets.Features({"ImageData": datasets.Sequence(datasets.Image(), length=V)}) ) return ds ds = get_ds() print(f"转换前特征: {ds.features=}") # 自定义转换函数:将图片序列转为4D数组 def images_seq_to_array4d(images_seq): # 堆叠所有图片为(V,H,W,C)数组,再调整通道维度到PyTorch习惯的位置 arr = np.stack([np.array(img) for img in images_seq], axis=0) arr = arr.transpose(0, 3, 1, 2) arr = arr.astype(np.float32) / 255.0 return arr # 直接转换列类型并应用转换逻辑 ds = ds.cast_column( "ImageData", datasets.Array4D(shape=(V, C, H, W), dtype="float32"), function=images_seq_to_array4d ) print(f"转换后特征: {ds.features=}") # 验证结果格式 ds.set_format(type="torch") assert next(iter(ds))['ImageData'].shape == (V, C, H, W)
方法2:利用with_format批量处理
临时设置Torch格式批量读取数据,处理后直接用Array4D特征重建数据集,避免map的单线程嵌套序列构建:
import datasets import PIL.Image import torch V = 16 H, W, C = 244, 244, 3 def get_ds(): """外部提供的数据集生成函数""" N = 10 data = [ {"ImageData": [PIL.Image.new("RGB", (W, H)) for _ in range(V)]} for _ in range(N) ] ds = datasets.Dataset.from_list(data) ds = ds.cast( datasets.Features({"ImageData": datasets.Sequence(datasets.Image(), length=V)}) ) return ds ds = get_ds() print(f"转换前特征: {ds.features=}") # 临时设置Torch格式,批量读取所有数据并处理 with ds.with_format(type="torch"): all_data = ds[:] all_data["ImageData"] = all_data["ImageData"].float() / 255.0 # 用处理后的数据重建数据集,直接指定Array4D特征 ds = datasets.Dataset.from_dict( all_data, features=datasets.Features({"ImageData": datasets.Array4D(shape=(V, C, H, W), dtype="float32")}) ) print(f"转换后特征: {ds.features=}") assert next(iter(ds))['ImageData'].shape == (V, C, H, W)
原方法性能差的原因
原流程中,map函数处理每个样本时会将PyTorch张量拆分为嵌套的float序列存储,Hugging Face对嵌套Sequence的序列化是单线程操作,且嵌套结构的存储开销极大——大规模数据下,单线程处理会直接导致集群性能灾难。上述两种方法均跳过了嵌套序列的生成步骤,直接生成Array4D格式的连续存储,利用数组的批量处理能力大幅提升速度。
内容的提问来源于stack exchange,提问作者LudvigH
相关产品推荐
相关产品推荐

