Diffusers
Safetensors
AR / check_data_num.py
xfcghj's picture
Upload folder using huggingface_hub
f0fc238 verified
Raw
History Blame Contribute Delete
1.11 kB
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}")