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()