如何优化Ulam螺旋无限迭代器的记忆化实现?
问题
我实现了一个将自然数映射为Ulam螺旋式格点的无限迭代器,已经尽可能优化性能,没有使用任何if条件。为避免重复计算,我尝试通过复用生成器对迭代器输出做记忆化,但这个实现反而导致性能下降。现在需要能跳过已计算项且提升效率的优化方案,相关代码及性能测试数据如下:
初始迭代器代码
from itertools import islice, repeat def ulamish_spiral_gen(): xc = yc = length = 0 yield 0, 0 while True: length += 1 yield from zip(range(xc + 1, (xc := xc + length) + 1, 1), repeat(yc)) yield from zip(repeat(xc), range(yc + 1, (yc := yc + length) + 1, 1)) length += 1 yield from zip(range(xc - 1, (xc := xc - length) - 1, -1), repeat(yc)) yield from zip(repeat(xc), range(yc - 1, (yc := yc - length) - 1, -1)) def ulamish_spiral(n): return list(islice(ulamish_spiral_gen(), n))
尝试的记忆化实现(性能下降)
COMPUTED = [] ULAMISH_GEN = ulamish_spiral_gen() def ulamish_spiral(n): if n > (l := len(COMPUTED)): COMPUTED.extend(islice(ULAMISH_GEN, n - l)) return COMPUTED[:n]
性能测试结果
In [225]: %timeit COMPUTED.clear(); ULAMISH_GEN = ulamish_spiral_gen(); ulamish_spiral(8192) 928 µs ± 8.96 µs per loop (mean ± std. dev. of 7 runs, 1,000 loops each) In [226]: %timeit COMPUTED.clear(); ULAMISH_GEN = ulamish_spiral_gen(); ulamish_spiral(1024); ulamish_spiral(2048); ulamish_spiral(4096); ulamish_spiral(8192) 993 µs ± 18.2 µs per loop (mean ± std. dev. of 7 runs, 1,000 loops each) In [227]: %timeit COMPUTED.clear(); ULAMISH_GEN = ulamish_spiral_gen(); ulamish_spiral(1024); ulamish_spiral(2048); ulamish_spiral(4096); ulamish_spiral(8192); ulamish_spiral(16384) 2.14 ms ± 106 µs per loop (mean ± std. dev. of 7 runs, 100 loops each) In [228]: %timeit COMPUTED.clear(); ULAMISH_GEN = ulamish_spiral_gen(); ulamish_spiral(16384) 2 ms ± 88.7 µs per loop (mean ± std. dev. of 7 runs, 100 loops each) In [229]: COMPUTED.clear(); ULAMISH_GEN = ulamish_spiral_gen(); ulamish_spiral(1024); ulamish_spiral(2048); ulamish_spiral(16384) == list(islice(ulamish_spiral_gen(), 16384)) Out[229]: True
优化方案
问题根源
原记忆化实现性能下降的核心原因:
- 每次调用
ulamish_spiral(n)都执行COMPUTED[:n]切片,会创建新列表,n越大开销越明显; - 全局变量的访问和维护带来额外性能损耗;
- 多次调用时的边界判断、生成器切片的微小开销持续累积。
优化实现1:类封装缓存状态
用类封装缓存和生成器,避免全局变量,同时减少不必要的列表创建:
from itertools import islice, repeat class UlamSpiralCache: def __init__(self): self.cache = [(0, 0)] self.generator = self._gen() def _gen(self): xc = yc = length = 0 while True: length += 1 yield from zip(range(xc + 1, (xc := xc + length) + 1, 1), repeat(yc)) yield from zip(repeat(xc), range(yc + 1, (yc := yc + length) + 1, 1)) length += 1 yield from zip(range(xc - 1, (xc := xc - length) - 1, -1), repeat(yc)) yield from zip(repeat(xc), range(yc - 1, (yc := yc - length) - 1, -1)) def get(self, n): current_len = len(self.cache) if n > current_len: self.cache.extend(islice(self.generator, n - current_len)) # 若不需要返回副本,可直接返回self.cache[:n](列表切片为副本,按需调整) return self.cache[:n] # 单例实例,确保全局复用同一缓存和生成器 ulamish_spiral_cache = UlamSpiralCache() def ulamish_spiral(n): return ulamish_spiral_cache.get(n)
这个版本解决了多次调用的性能问题:第一次调用n=16384耗时和原迭代器接近,后续调用n≤16384时直接返回缓存,几乎无额外开销。
优化实现2:数学公式直接计算
如果不需要迭代器的惰性特性,可通过数学公式直接计算第k个点的坐标,完全跳过生成器和缓存开销,单个点计算时间为O(1):
def ulam_spiral_point(k): if k == 0: return (0, 0) # 计算所在层(从1开始计数) m = int((k ** 0.5 - 1) // 2) side = 2 * m + 1 start = (side - 2) ** 2 + 1 edge_len = 2 * m offset = k - start if offset < edge_len: return (m, -m + offset) offset -= edge_len if offset < edge_len: return (m - offset, m) offset -= edge_len if offset < edge_len: return (-m, m - offset) offset -= edge_len return (-m + offset, -m) def ulamish_spiral(n): return [ulam_spiral_point(k) for k in range(n)]
注意:需验证公式输出和原迭代器的坐标方向完全匹配,若有差异微调坐标符号或偏移逻辑即可。该方案批量生成的性能远高于迭代器+缓存模式,且多次调用无额外开销。
内容的提问来源于stack exchange,提问作者Ξένη Γήινος
相关产品推荐
相关产品推荐

