为什么BatchIndices对象不是迭代器?
为什么BatchIndices对象不是迭代器?
好问题!咱们来一步步拆解这个问题,先看看你提供的代码:
import threading import numpy as np class BatchIndices(object): def __init__(self, n, bs, shuffle=False): self.n,self.bs,self.shuffle = n,bs,shuffle print(n,bs) self.lock = threading.Lock() self.reset() def reset(self): self.idxs = (np.random.permutation(self.n) if self.shuffle else np.arange(0, self.n)) self.curr = 0 def __next__(self): with self.lock: if self.curr >= self.n: self.reset() ni = min(self.bs, self.n-self.curr) res = self.idxs[self.curr:self.curr+ni] self.curr += ni print(res) return res # 示例调用 bi = BatchIndices(10, 3)
核心原因:缺少__iter__()方法
在Python里,一个对象要被认定为迭代器,必须同时满足两个硬性条件:
- 实现
__next__()方法:你的类已经做到了,这个方法负责返回下一个元素(你的代码里还做了循环重置,属于无限迭代的设计) - 实现
__iter__()方法:这个方法必须返回迭代器对象本身(也就是self)——而你的类完全没写这个方法,所以Python不把它当作迭代器。
怎么修复?
只需要给BatchIndices类添加一个简单的__iter__方法就行:
def __iter__(self): return self
添加之后,你的对象就变成了标准迭代器,能直接用在for循环这类迭代场景里了,比如:
for batch in bi: print("正在处理批次:", batch) # 因为你的类是无限迭代的,记得加终止条件 if some_stop_condition: break
额外提一句:如果只是可迭代对象(比如列表、元组),只需要实现__iter__()返回一个迭代器就行,但迭代器本身必须同时具备__iter__和__next__两个方法。你的类之前缺了关键的__iter__,所以既不是可迭代对象也不是迭代器。
内容的提问来源于stack exchange,提问作者waseem abbas
相关产品推荐
相关产品推荐

