为分布式训练准备数据集
您已经为一个精准农业系统预处理了数据集,以帮助农民监测作物健康。现在,您将通过创建 DataLoader 来加载数据,并在可用时将数据放到 GPU 上进行分布式训练。请注意,本练习实际使用的是 CPU,但在 CPU 和 GPU 上代码是相同的。
已预加载部分数据:
- 一个包含农业影像的示例
dataset - 来自
accelerate库的Accelerator类 DataLoader类
本练习是课程的一部分
使用 PyTorch 高效训练 AI 模型
练习说明
- 为预定义的
dataset创建一个dataloader。 - 使用
accelerator对象将dataloader放置到可用设备上。
交互式实操练习
通过完成这段示例代码来试试这个练习。
accelerator = Accelerator()
# Create a dataloader for the pre-defined dataset
dataloader = ____(____, batch_size=32, shuffle=True)
# Place the dataloader on available devices
dataloader = accelerator.____(____)
print(accelerator.device)