Diffusers
Safetensors
AR / test_dataloader.py
xfcghj's picture
Upload folder using huggingface_hub
f0fc238 verified
Raw
History Blame Contribute Delete
3.84 kB
import os
import glob
import tarfile
import pickle
import zstandard
import torch
import torch.distributed as dist
from torch.utils.data import IterableDataset, DataLoader
class DyMeshIterableDataset(IterableDataset):
def __init__(self, data_dir, rank=0, world_size=1):
super(DyMeshIterableDataset, self).__init__()
self.tar_paths = sorted(glob.glob(os.path.join(data_dir, "shard_*.tar.zst")))
self.my_tar_paths = [p for i, p in enumerate(self.tar_paths) if i % world_size == rank]
print(f"[Rank {rank}] 负责处理 {len(self.my_tar_paths)} 个分片")
def __iter__(self):
worker_info = torch.utils.data.get_worker_info()
if worker_info is None:
curr_paths = self.my_tar_paths
else:
per_worker = (len(self.my_tar_paths) + worker_info.num_workers - 1) // worker_info.num_workers
wid = worker_info.id
curr_paths = self.my_tar_paths[wid * per_worker : (wid + 1) * per_worker]
dctx = zstandard.ZstdDecompressor()
for tar_path in curr_paths:
try:
with open(tar_path, 'rb') as fh:
with dctx.stream_reader(fh) as reader:
with tarfile.open(fileobj=reader, mode='r|') as tar:
for member in tar:
if member.isfile():
f = tar.extractfile(member)
if f:
data = pickle.load(f)
yield {'vertices': data['vertices'], 'faces': data['faces'], 'caption': data.get('caption', '无')}
except Exception: continue
def get_mesh_dataloader(data_dir, batch_size=4, num_workers=2):
rank = int(os.environ.get("RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
dataset = DyMeshIterableDataset(data_dir, rank=rank, world_size=world_size)
return DataLoader(dataset, batch_size=batch_size, collate_fn=lambda x: x, num_workers=num_workers, pin_memory=True)
def main():
dist.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
global_rank = dist.get_rank()
world_size = dist.get_world_size()
torch.cuda.set_device(local_rank)
dataloader = get_mesh_dataloader("/home/dataset-assist-0/usr/lh/ysh/dw/RL/AR/data/shards/", batch_size=40, num_workers=8)
local_batch_count = 0
local_sample_count = 0
for batch_idx, batch in enumerate(dataloader):
local_batch_count += 1
local_sample_count += len(batch) # 统计实际样本数
if batch_idx % 50 == 0:
print(f"[Rank {global_rank}] 进度: 处理至第 {local_batch_count} 个 batch (累计样本: {local_sample_count})")
dist.barrier()
# 汇总 Batch 数量
batch_counts = torch.zeros(world_size, dtype=torch.long, device=f"cuda:{local_rank}")
batch_counts[global_rank] = local_batch_count
dist.all_reduce(batch_counts, op=dist.ReduceOp.SUM)
# 汇总 Sample 数量
sample_counts = torch.zeros(world_size, dtype=torch.long, device=f"cuda:{local_rank}")
sample_counts[global_rank] = local_sample_count
dist.all_reduce(sample_counts, op=dist.ReduceOp.SUM)
if global_rank == 0:
print("\n" + "="*50)
print("📊 分布式训练统计报告")
for i in range(world_size):
print(f" Rank {i}: 处理了 {batch_counts[i].item()} 个 Batch, 共 {sample_counts[i].item()} 个样本")
print("-" * 50)
print(f" 全局总 Batch 数: {batch_counts.sum().item()}")
print(f" 全局总样本数 : {sample_counts.sum().item()}")
print("="*50 + "\n")
dist.destroy_process_group()
if __name__ == "__main__":
main()