1.3 现场三:用类封装一个迭代器


1.3 现场三:用类封装一个迭代器

本节摘要:面试官要求写一个类,能像内置 range 一样被 for 循环遍历,并支持按批次吐出数据。这道题考查可迭代对象与迭代器的协议分界、__iter__ 与 __next__ 的分工、以及 yield 的惰性求值。它是热身现场的收束题,也是第 5 章 NLP 数据加载器的直接前置。

前两题一个考"用得熟"、一个考"造得出结构",这一题考"懂不懂协议"——Python 的 for 循环背后是一套协议在运转,理解它的人写出的类才能和语言本身咬合。

面试官提问

"写一个类 BatchLoader,构造时传入一个列表和批次大小。要求它能被 for 循环直接遍历,每次迭代吐出一个批次;批次大小必须是正整数,否则构造时就报错。顺便解释一下,for 循环是怎么工作的。"

题目里藏着两个考点:协议方法必须写对;参数校验要放在构造函数里,而不是等到迭代时才炸。

现场推演

候选人先答 for 的工作机制:"for 拿到对象后先调 __iter__ 拿一个迭代器,然后反复调迭代器的 __next__,直到抛出 StopIteration。"面试官追问:"所以可迭代对象和迭代器是同一个东西吗?"候选人答:"不是。可迭代对象实现 __iter__;迭代器实现 __next__ 和 __iter__(返回自身)。列表是可迭代的,但它每次 __iter__ 都产生新的迭代器,所以能被多个 for 独立遍历。"

这个区分答清楚了,代码就顺了。第一版:让类自身充当迭代器。

class BatchLoader: def __init__(self, data, batch_size): if not isinstance(batch_size, int) or batch_size <= 0: raise ValueError("batch_size 必须是正整数") self.data = data self.bs = batch_size self.i = 0 # 游标 def __iter__(self): self.i = 0 # 复位游标,支持再次遍历 return self # 自身即迭代器 def __next__(self): if self.i >= len(self.data): raise StopIteration # for 捕获它后正常结束 batch = self.data[self.i:self.i + self.bs] self.i += self.bs return batch bl = BatchLoader([1, 2, 3, 4, 5, 6, 7], 3) print([b for b in bl]) # 第一次遍历 print([b for b in bl]) # 第二次遍历也正常
[[1, 2, 3], [4, 5, 6], [7]] [[1, 2, 3], [4, 5, 6], [7]]

尾部不足一批时吐出短批次([7]),这与 PyTorch 的 DataLoader 在 drop_last=False 时的行为一致——候选人主动点出这个对应关系,面试官明显来了兴趣。

追问链

第一层:这个设计有什么毛病? 候选人自己指出来了:"自身即迭代器意味着状态只有一份。如果两个地方同时 for 它,游标互相踩。"然后给出第二版:可迭代对象与迭代器分成两个类。

class _BatchIter: def __init__(self, data, bs): self.data, self.bs, self.i = data, bs, 0 def __iter__(self): return self def __next__(self): if self.i >= len(self.data): raise StopIteration batch = self.data[self.i:self.i + self.bs] self.i += self.bs return batch class BatchLoader2: def __init__(self, data, batch_size): if not isinstance(batch_size, int) or batch_size <= 0: raise ValueError("batch_size 必须是正整数") self.data, self.bs = data, batch_size def __iter__(self): return _BatchIter(self.data, self.bs) # 每次遍历独立游标 it1, it2 = iter(BatchLoader2([1,2,3,4], 2)), iter(BatchLoader2([1,2,3,4], 2)) print(next(it1), next(it1), next(it2))
[1, 2] [3, 4] [1, 2]

两个迭代器互不干扰——it2 吐出的是 [1, 2] 而不是接着 it1 的位置。这正是列表能被嵌套遍历的原因。

第二层:用生成器行不行? 行,而且更短。含 yield 的函数调用时返回生成器,生成器自动实现迭代器协议,StopIteration 由函数自然结束触发。

def batch_gen(data, bs): if bs <= 0: raise ValueError("batch_size 必须是正整数") for i in range(0, len(data), bs): yield data[i:i + bs] # 惰性:被 next 到才切这一批 g = batch_gen(list("abcde"), 2) print(next(g), next(g), list(g)) # 第三次直接耗尽剩余
['a', 'b'] ['c', 'd'] [['e']]

注意 list(g) 拿到的只剩 [['e']]——生成器是一次性的,耗尽后再取就是空。面试官追问"惰性有什么实际价值",候选人答:"数据集比内存大时,逐批读文件逐批吐,内存里永远只有一批;第 5 章写文本数据加载器就是这个模式。"

第三层:过滤批次大小? 面试官最后加码:"最后一批只有一个元素,训练时会引入统计噪声,怎么去掉?"候选人答两种:切片前判断剩余长度,或生成器里 if len(batch) == bs: yield batch。前者是 drop_last 参数的思路,一行改动。

失误复盘

常见翻车点:把 __next__ 写成返回 False 表示结束(协议要求抛 StopIteration,for 不会停);__iter__ 里忘记返回 self 或新迭代器;在 __next__ 里用 self.i += 1 而不是 += batch_size,批次变成滑动窗口;构造时不校验参数,batch_size 为零时 for 直接死循环。死循环这个尤其致命——白板上一次,时间就没了。

主线候选人这一场答得最顺,协议区分讲得清楚,还主动连接了 DataLoader。他的复盘写在笔记本上的一句话是:"面试官问 for 怎么工作,不是考背诵,是看你写类之前有没有把语言的地基想清楚。"

关键直觉:可迭代是"能被要一个迭代器",迭代器是"能被要下一个元素"——两个职责拆开写,类才经得起并发遍历的追问。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U