Diffusers
Safetensors
AR / split.py
xfcghj's picture
Upload folder using huggingface_hub
f0fc238 verified
Raw
History Blame Contribute Delete
3.88 kB
import os
import tarfile
import zstandard
class WritingCounter:
"""一个轻量级的包装器,用于在流式写入时精确统计写入磁盘的字节数"""
def __init__(self, fileobj):
self.fileobj = fileobj
self.count = 0
def write(self, buf):
self.fileobj.write(buf)
self.count += len(buf)
def close(self):
self.fileobj.close()
def split_big_dataset(source_path, output_dir, samples_per_shard=10000, max_shard_size_gb=3.2):
"""
将一个巨大的 tar.zst 文件切分成多个固定的 shard_xxxx.tar.zst
条件:每隔 samples_per_shard 个样本 OR 当压缩包大小达到 max_shard_size_gb 时切分
"""
os.makedirs(output_dir, exist_ok=True)
# 将 GB 转换为字节数 (Byte)
MAX_SHARD_BYTES = int(max_shard_size_gb * 1024 * 1024 * 1024)
dctx = zstandard.ZstdDecompressor()
cctx = zstandard.ZstdCompressor(level=3) # ✨ 修复:改回正确的 ZstdCompressor
shard_idx = 0
sample_in_shard = 0
current_shard_fh = None
current_counter = None
current_zstd_stream = None
current_tar = None
def close_current_shard():
nonlocal current_tar, current_zstd_stream, current_counter, current_shard_fh
# 按照反向顺序关闭,确保数据完全刷入磁盘
if current_tar:
current_tar.close()
current_tar = None
if current_zstd_stream:
current_zstd_stream.close()
current_zstd_stream = None
if current_counter:
current_counter.close()
current_counter = None
try:
with open(source_path, 'rb') as fh:
with dctx.stream_reader(fh) as reader:
with tarfile.open(fileobj=reader, mode='r|') as src_tar:
for member in src_tar:
if not member.isfile():
continue
src_f = src_tar.extractfile(member)
if src_f is None:
continue
# ✨ 修复:移除了上一版残留的错误隐藏行
if current_tar is None or sample_in_shard >= samples_per_shard or (current_counter and current_counter.count >= MAX_SHARD_BYTES):
close_current_shard()
shard_name = f"shard_{shard_idx:04d}.tar.zst"
shard_path = os.path.join(output_dir, shard_name)
print(f"📦 正在创建分片: {shard_name} ...")
current_shard_fh = open(shard_path, 'wb')
current_counter = WritingCounter(current_shard_fh)
current_zstd_stream = cctx.stream_writer(current_counter)
current_tar = tarfile.open(fileobj=current_zstd_stream, mode='w|')
shard_idx += 1
sample_in_shard = 0
# 写入到新的分片中
current_tar.addfile(member, src_f)
sample_in_shard += 1
finally:
close_current_shard()
print(f"🎉 切分完成!共生成 {shard_idx} 个分片文件,存放于 {output_dir}")
# 使用示例
if __name__ == "__main__":
split_big_dataset(
source_path="/home/dataset-assist-0/usr/lh/ysh/dw/RL/AR/data/DyMesh_complete.tar.zst",
output_dir="/home/dataset-assist-0/usr/lh/ysh/dw/RL/AR/data/shards/",
samples_per_shard=10000, # 阈值 1:每 10000 个样本
max_shard_size_gb=3.2 # 阈值 2:或者每满 3.2 GB
)