当前位置: 首页 > news >正文

Epoch、Batch 与 DataLoader

很多刚开始阅读 PyTorch 推荐系统训练代码的同学,都会卡在一组名词:
epochbatchDataLoadershuffleloss.backward()optimizer.step()testHRNDCG

单独看每个概念不难,但放进训练循环里,很容易分不清整条流水线:
数据从哪里读取?如何切分成 batch?模型什么时候更新参数?测试阶段会不会偷偷修改权重?

本文用最简单的数字样本,完整拆解整条训练链路。

✅ 核心结论

  1. 一个epoch= 完整遍历一遍全部训练数据集。
  2. 单个epoch内部,循环读取多个batch
  3. 训练阶段,每拿到一个 batch,标准流程:
    前向传播 → 计算损失 → 清空梯度 → 反向传播 → 更新模型参数。
  4. verbose用来控制:每隔多少个 epoch,执行日志打印/离线评估。
  5. 测试评估阶段:仅做预测、计算指标(HR@K、NDCG@K),不会更新模型参数
  6. shuffle=True只改变样本读取顺序,不会修改 Dataset 内部原始数据。
  7. shuffle=True每一轮epoch开始,重新生成随机读取顺序
  8. shuffle=False:所有epoch读取样本顺序完全固定。

一句话概括分工:

Dataset管理「有哪些数据」;DataLoader管理「以什么顺序、多少条一组取出数据」;训练循环控制「拿到batch后如何迭代优化模型」。


🔢 极简数字数据集演示

假设我们一共有10条训练样本,使用索引编号代表样本:

原始数据集索引:[0, 1, 2, 3, 4, 5, 6, 7, 8, 9]

参数设置:

batch_size=4

含义:一次性取出4条样本,封装为1个batch。

场景1:shuffle=False(不打乱顺序)

数据会按原生索引切分batch:

第1个 batch: [0, 1, 2, 3] 第2个 batch: [4, 5, 6, 7] 第3个 batch: [8, 9]

总数10无法被4整除,最后一组为不完整batch。
如果开启drop_last=True,最后这个不完整batch[8, 9]会直接丢弃。

shuffle=False:所有epoch顺序固定
读取索引永远不变,连续两轮epoch批次完全相同:

epoch 1: batch 1: [0, 1, 2, 3] batch 2: [4, 5, 6, 7] batch 3: [8, 9] epoch 2: batch 1: [0, 1, 2, 3] batch 2: [4, 5, 6, 7] batch 3: [8, 9]

场景2:shuffle=True(开启打乱)

❗重点:打乱的是索引读取顺序,原始Dataset内的数据本身保持不变。
原始数据依旧是[0, 1, 2, 3, 4, 5, 6, 7, 8, 9]

关键点总结:
shuffle=True≠ 永久修改数据;
每一轮epoch启动时,DataLoader重新规划本轮样本读取次序。


📦 DataLoader 到底如何生成一个 batch?

这是新手最容易混淆的环节。我们拆成完整5步理解。

自定义数据集模板:

classNumberDataset(torch.utils.data.Dataset):def__init__(self):self.samples=[{"idx":0,"feature":[0.0,0.5],"label":0},{"idx":1,"feature":[1.0,1.5],"label":1},{"idx":2,"feature":[2.0,2.5],"label":0},# ...更多样本]def__len__(self):returnlen(self.samples)def__getitem__(self,index):# 根据索引返回【单条样本】returnself.samples[index]

训练循环代码:

forbatchindata_loader:...

循环背后完整流程:

第1步:生成索引顺序

  • shuffle=False:原生索引[0,1,2,3,...]
  • shuffle=True:随机打乱索引序列

注意:数字只是样本索引,不是样本内容。后续会调用dataset[index]获取单条数据。

第2步:按batch_size分组索引

索引序列[3, 7, 1, 9, 0, 6, 2, 8, 4, 5]
batch_size=4 → 分组:
[3,7,1,9][0,6,2,8][4,5]
当前仅规划索引,还没有读取真实样本

第3步:逐条调用Dataset.__getitem__

取索引组[3,7,1,9]
依次执行:

sample_1=dataset[3]sample_2=dataset[7]sample_3=dataset[1]sample_4=dataset[9]

得到4条独立样本字典:

[{"idx":3,"feature":[3.0,3.5],"label":1},{"idx":7,"feature":[7.0,7.5],"label":1},{"idx":1,"feature":[1.0,1.5],"label":1},{"idx":9,"feature":[9.0,9.5],"label":1},]

第4步:collate_fn 打包,拼接成batch张量


默认default_collate负责打包:把多条样本同字段堆叠
打包完成后:

{"idx":tensor([3,7,1,9]),"feature":tensor([[3.0,3.5],[7.0,7.5],[1.0,1.5],[9.0,9.5],]),"label":tensor([1,1,1,1]),}

这就是训练循环拿到的batch。数据类型本质上是:

dict[str,torch.Tensor]

通俗理解:
Dataset产出一条条独立样本;collate_fn打包员,把多条样本组装成一个batch张量。

第5步:训练循环接收batch,送入模型

外部代码直接使用:

forbatchintrain_loader:features=batch["feature"]labels=batch["label"]scores=model(features)

整条链路简化:
生成索引顺序 → 索引分组 → __getitem__逐条读取样本 → collate_fn拼接batch → 返回循环


🔁 完整训练循环执行流程

示例标准训练代码:

epochs=3batch_size=4verbose=2forepochinrange(1,epochs+1):model.train()# 内层循环遍历所有batchforbatchintrain_loader:features=batch["feature"]labels=batch["label"]scores=model(features)loss=loss_fn(scores,labels)optimizer.zero_grad()loss.backward()optimizer.step()# 间隔verbose个epoch执行测试ifepoch%verbose==0:evaluate(model,test_loader)

流程翻译:

  • Epoch 1:完整遍历训练集,所有batch更新参数;1%2≠0不测试
  • Epoch 2:完整遍历训练集,所有batch更新参数;2%2=0执行测试评估
  • Epoch 3:完整遍历训练集,所有batch更新参数;3%2≠0不测试

单个batch内部标准训练闭环

model.train()# 切换训练模式(Dropout/BatchNorm生效)scores=model(features)# 前向传播,得到预测值loss=loss_fn(scores,labels)# 计算损失optimizer.zero_grad()# 清空上一轮梯度loss.backward()# 反向传播,计算参数梯度optimizer.step()# 使用梯度更新模型权重

⚖️ 训练阶段 VS 测试评估阶段

阶段是否计算梯度是否反向传播是否更新参数核心目的
训练阶段✅ 是✅ 是✅ 是迭代优化模型
测试评估❌ 否❌ 否❌ 否观测模型泛化效果

训练代码模板

model.train()forbatchintrain_loader:scores=model(batch["feature"])loss=loss_fn(scores,batch["label"])optimizer.zero_grad()loss.backward()optimizer.step()

测试评估模板

model.eval()withtorch.no_grad():# 关闭梯度计算,节省显存forbatchintest_loader:scores=model(batch["feature"])# 根据预测分数计算HR、NDCG指标

torch.no_grad():告知PyTorch仅推理,不需要构建梯度计算图。


📊 推荐系统指标:HR@K 和 NDCG@K

推荐任务离线评估最常用两个指标:

HR@K(Hit Ratio,命中率)

含义:给用户推荐Top-K物品,真实交互物品是否出现在推荐列表内

  • 命中:HR=1
  • 未命中:HR=0

示例:
推荐Top5列表:[item_8, item_2, item_6, item_1, item_9]
用户真实喜爱物品:item_6
item_6 在列表中 → HR@5 = 1。

NDCG@K

HR只关心「有没有命中」;NDCG额外关注命中物品的排序位置
案例:真实物品为 item_6

  • 推荐A:[item_6, item_2, item_8, item_1, item_9](命中排在第1位)
  • 推荐B:[item_8, item_2, item_1, item_9, item_6](命中排在第5位)

两者HR都等于1,但推荐A的NDCG更高。
简单理解:

HR:是否猜中;NDCG:猜中之后,排得够不够靠前。


🛠️ 可直接运行的演示代码

复制运行,直观观察__getitem__的调用时机:

importtorchfromtorch.utils.dataimportDataset,DataLoaderclassNumberDataset(Dataset):def__init__(self):self.samples=[]foridxinrange(10):self.samples.append({"idx":idx,"feature":torch.tensor([float(idx),float(idx)+0.5]),"label":torch.tensor(idx%2,dtype=torch.float32),})def__len__(self):returnlen(self.samples)def__getitem__(self,index):print(f" __getitem__ 被调用,index={index}")returnself.samples[index]dataset=NumberDataset()loader=DataLoader(dataset,batch_size=4,shuffle=False)forbatch_id,batchinenumerate(loader,start=1):print(f"\n第{batch_id}个 batch")print("idx:",batch["idx"])print("feature:",batch["feature"])print("label:",batch["label"])

运行输出可以验证:

  1. DataLoader 逐个调用__getitem__获取单条样本
  2. 收集足够样本后,collate_fn自动合并为batch张量
  3. 训练循环拿到的是封装完成的批量数据,而非单样本

💡 收尾总结

看到推荐系统经典双层循环代码:

forepochinrange(epochs):forbatchintrain_loader:...

脑海中自动翻译:

外层epoch循环:控制一共完整训练多少轮;
内层batch循环:一轮训练中分批读取数据;
每一个训练batch都会更新模型参数;
到达指定epoch间隔,进入评估模式,只预测、计算指标,不更新权重。

打通这套流程后,再阅读 BPR、NeuMF、LightGCN 等推荐模型训练代码,理解门槛会大幅降低。

http://www.cnnetsun.cn/news/3776705.html

相关文章:

  • 点击化学:从CuAAC到SPAAC,掌握模块化分子连接的底层逻辑与实战指南
  • C++ vector多维数组初始化:一行代码实现高效内存管理
  • 同样的 Agent,换了一套提示词,效果翻了 5 倍:Skill 工程实战指南
  • GPT文本生成原理与采样策略优化实践
  • 工业级PID控制器C语言实现:从离散化到抗饱和与参数整定
  • MATLAB图像处理实战:空域与频域方法消除条纹干扰
  • Unity与Cocos2d-x双引擎实现Flappy Bird:源码对比与实战解析
  • 商用AI主机如何解决Token成本与稳定性难题,赋能本地大模型应用开发
  • LDO与DC-DC选型指南:从压差、功耗到锂电池供电的实战解析
  • ComfyUI UltimateSDUpscale安装问题深度解析:从模块缺失到完美修复
  • K8s StatefulSet 持久化存储:PV 绑定、扩容与快照备份
  • 推挽与开漏输出电路原理详解:从MOSFET结构到I2C总线应用
  • AI Agent如何自动化生成PPT:从技术原理到实践应用
  • STM32开发中“Not a genuine ST Device!”错误排查与解决指南
  • YimMenu终极指南:3步打造GTA5最强防崩溃游戏菜单
  • LangChain技能全景:从基础连接到生产级智能体部署全解析
  • Arduino入门指南:从环境搭建到项目实战,快速上手物联网开发
  • 电子工程师必备:电容选型实战指南与高频特性深度解析
  • 基于 Free Pascal 从零编写裸机操作系统(一)
  • CST同轴线仿真全流程:从建模优化到高频连接器设计实践
  • 大模型思考过程加密:技术原理、行业影响与工程应对策略
  • PCB板HDI1/HDI2/HDI3/HDI…、ELIC指的是什么?
  • 微信聊天记录AI分析:原理、应用与隐私安全实践指南
  • 主流 Agent 架构分析
  • 步进电机步距角与细分驱动详解:从原理到实战,告别抖动与丢步
  • 实用的工艺品设计服务受青睐,优质选择不容错过
  • AI助手APP竞争格局解析:从通用到垂直,如何选择与高效使用?
  • 首届OPC-AI赋能实战班在甬举办,分享AI超级个体实践思考
  • MAXQDA 2020安装与核心功能详解:从环境配置到定性数据分析实战
  • 深入解析ARM SWD协议:从原理到实战的嵌入式调试核心