代码dataloaders = zip(labeled_trainloader, [None]*len(labeled_trainloader))含义咨询
问题解答
代码具体含义
你贴的代码是Python中将带标注训练数据加载器与等长空值列表配对的操作,完整代码如下:
dataloaders = zip(labeled_trainloader, [None] * len(labeled_trainloader))
具体执行逻辑:
labeled_trainloader是PyTorch等深度学习框架中常见的带标注训练集数据加载器,本身是可迭代对象,每次迭代返回一批带标签的训练数据[None] * len(labeled_trainloader)生成了一个和labeled_trainloader迭代总步数完全相等的列表,列表所有元素都是Nonezip()将两个可迭代对象按位置一一配对,最终得到的dataloaders也是可迭代对象,每次迭代返回一个二元组:第一个元素是每批带标注训练数据,第二个元素固定为None
生成等长None列表的原因
这个写法几乎都是用于适配半监督训练的统一代码逻辑:
- 常规半监督训练的逻辑会同时遍历带标注数据加载器和无标注数据加载器,一般写法是
zip(labeled_loader, unlabeled_loader),每次同时取一批标注数据和无标注数据输入模型 - 当你当前不需要使用无标注数据、只想跑纯监督训练的基线时,不需要单独修改下游的遍历代码,直接用全
None的列表替代无标注加载器即可。下游代码只要在用到无标注数据的位置判断值是否为None,就能同时兼容监督/半监督两种训练模式,减少冗余分支代码 - 生成和
labeled_trainloader长度完全一致的None列表,是为了匹配zip的迭代规则:zip会以输入的最短可迭代对象的长度为准停止迭代,等长设置可以保证最终dataloaders的迭代步数和原带标注加载器完全一致,不会出现数据截断或者多返回无效值的问题
内容的提问来源于stack exchange,提问作者beginner
相关产品推荐
相关产品推荐

