import torch from tqdm import tqdm # 推荐安装 tqdm 以显示进度条 from utils.autodataloader_v3 import get_mesh_dataloader def count_dataset_samples(dataloader): print("正在统计数据总数,请稍候...") count = 0 # 使用 tqdm 包装 dataloader 可以直观看到进度,如果不方便安装,可去掉 tqdm for batch in tqdm(dataloader): # 这里的 batch 是由 collate_fn 返回的字典 # 由于 batch_size 可能大于 1,每次累加当前 batch 的大小 # 假设 batch['vertices'] 是一个列表,其长度即为当前批次的样本数 batch_size = len(batch['vertices']) # print(f"batch_size:{batch_size}") count += batch_size return count if __name__ == "__main__": # 实例化 DataLoader # 注意:如果数据量非常巨大,请确保 num_workers 设置合理,否则可能会因为读取速度瓶颈卡住 dataloader = get_mesh_dataloader(batch_size=32, num_workers=16) total_samples = count_dataset_samples(dataloader) print(f"\n✅ 数据集总样本数: {total_samples}")