如何让我用来生成正切数的这个类具备可恢复性?
我决定把实现代码用于生成切线数,作为一次自我设定的编程挑战。我已经成功实现了,但并没有达到我期望的效率。
什么是切线数?它们就是下面这个级数:
1
2
16
272
7936
353792
22368256
1903757312
209865342976
29088885112832
4951498053124096
1015423886506852352
246921480190207983616
70251601603943959887872
23119184187809597841473536
8713962757125169296170811392
3729407703720529571097509625856
1798651693450888780071750349094912
970982810785059112379399707952152576
583203324917310043943191641625494290432
在 OEIS 上这是A000182。
我实现了一个无穷的切线数生成器,使用的公式如下:a(n) = Sum_{k = 1..n-1} binomial(2 * n-2, 2 * k-1) * a(k) * a(n-k),且a(1) = 1,但我不会把它贴在这里,因为它相当长,且并非本题的重点;它并不如下面将要贴出的函数高效。为了让这个问题保持聚焦于一个具体主题,我不会在这里粘贴代码,但你可以通过这个 链接 查看。
我从 Peter Luschny 的页面中反向工程了以下函数,我在同一输入下它大约只需要原始生成器的2/3的时间来完成。
函数如下:
def A000182_list(n):
T = [0 for i in range(1, n+2)]
T[1] = 1
for k in range(2, n+1):
T[k] = (k-1)*T[k-1]
for k in range(2, n+1):
for j in range(k, n+1):
T[j] = (j-k)*T[j-1]+(j-k+2)*T[j]
return T[1:]
我当然知道它为何有效,但我知道它在做什么。我将逐行解释这段代码,但我对其背后的数学原理一无所知,尽管我知道在我自己生成器中使用的每一个恒等式和对称性。
这一行:T = [0 for i in range(1, n+2)] 起什么作用?显然,它初始化了一个长度为 n-1 的全0 列表。并且它效率很低。
第二行是:T[1] = 1,因此把下标为1 的元素设为1。我们可以看到下标0 在代码中从未被使用或访问过。所以我们只需要一个长度为n 的列表。
因此我把它重构成这样:result = [1] * n。通过 * 运算符对列表进行重复,是在Python中获得长度为n 的列表的最有效方法。
下面这个循环做了什么?
for k in range(2, n+1):
T[k] = (k-1)*T[k-1]
显然它从k=2循环到k=n,并把 T[K] 设为某个数。好,但那个数是什么?很明显,k-1从 1开始,随着k 的变化依次是2、3、4、5、6、7,当k 分别取2、3、4、5、6、7、8时。表达式k-1使用了两次,第二次使用会带来冗余计算。在k = 2时,它计算1 * T[1] = 1,并把结果赋给T[2]。在k = 3时,它计算2 * T[2] = 2,并把结果赋给T[3]。在k = 4时,它计算3 * T[3] = 6,并把结果赋给T[4]。在k = 5时,它计算4 * T[4] = 24,并把结果赋给T[5];在k = 6时,它计算5 * 24 = 120……换句话说,它只是计算阶乘,每一项都是前一项与随迭代逐步增大的正整数的乘积。由于这些项在下一次迭代中立即使用,并且恰好一次使用,T[k-1] 会产生冗余的列表索引。
我把代码重构为这样的:
result = [1] * n
last = 1
for i in range(1, n):
result[i] = last = last * i
下面这个嵌套循环又在做什么?看起来相当复杂:
for k in range(2, n+1):
for j in range(k, n+1):
T[j] = (j-k)*T[j-1]+(j-k+2)*T[j]
但实际上嵌套的for循环相当简单。首先,显然k 只是从2 开始到n 的连续正整数(2、3、4、5、6、7、8、9……n)。那么j 呢?j只是从当前迭代中k 的值开始的连续整数,直到n 结束。因此例如当n = 10时,所有迭代中的k、j的取值为:
In [31]: [list(range(i, 11)) for i in range(2, 11)]
Out[31]:
[[2, 3, 4, 5, 6, 7, 8, 9, 10],
[3, 4, 5, 6, 7, 8, 9, 10],
[4, 5, 6, 7, 8, 9, 10],
[5, 6, 7, 8, 9, 10],
[6, 7, 8, 9, 10],
[7, 8, 9, 10],
[8, 9, 10],
[9, 10],
[10]]
那么j-k是什么呢?由于j 从k 开始,j-k为 0;在下一次迭代j 增加1,因此j-k为 1;再下一次j 再增1,j-k为 2……显然j-k是从0 开始的连续整数,依次是0、1、2、3、4、5、6、7、8、9……因此j-k+2只是2、3、4、5、6、7、8、9、10、11……
和阶乘的模式一样,每一项都是通过前一项再索引来计算的。若把这个赋值给一个变量,我们就不需要再进行索引 T[j - 1]。
于是,这是我重构后的代码:
def A000182_fast(n):
result = [1] * n
last = 1
for i in range(1, n):
result[i] = last = last * i
for k in range(1, n):
b = 2
last = result[k - 1]
for a, c in enumerate(range(k, n)):
result[c] = last = a * last + b * result[c]
b += 1
return result
性能:
In [34]: A000182_fast(32)
Out[34]:
[1,
2,
16,
272,
7936,
353792,
22368256,
1903757312,
209865342976,
29088885112832,
4951498053124096,
1015423886506852352,
246921480190207983616,
70251601603943959887872,
23119184187809597841473536,
8713962757125169296170811392,
3729407703720529571097509625856,
1798651693450888780071750349094912,
970982810785059112379399707952152576,
583203324917310043943191641625494290432,
387635983772083031828014624002175135645696,
283727921907431909304183316295787837183229952,
227681379129930886488600284336316164603920777216,
199500252157859031027160499643195658166340757225472,
190169564657928428175235445073924928592047775873499136,
196535694915671808914892880726989984967498805398829268992,
219523439106761591280258358007964245987752702449505540243456,
264239411287900883270178745605712648488731058170551223831232512,
341838301335718580350174449297951396847081443826785448952307122176,
474090194351342155974522582010891145370303013658973457042263121068032,
703237958001393736999896827714634659411015090272684227831001161763127296,
1113255345330866700339746218047088280783690575394538814699382505942970007552]
In [35]: A000182_fast(32) == A000182_list(32)
Out[35]: True
In [36]: %timeit A000182_fast(256)
15.1 ms ± 109 μs per loop (mean ± std. dev. of 7 runs, 100 loops each)
In [37]: %timeit A000182_list(256)
15.6 ms ± 211 μs per loop (mean ± std. dev. of 7 runs, 100 loops each)
你可以看到,我的代码在严格意义上比原始代码更快,但提升并不明显。
那么我就此止步了吗?当然没有。尽管我已经有了可工作的高效代码,但它并没有解决我想解决的问题。
我想按需生成切线数,也就是说我想要像 tangent_numbers(n) 这样的东西,返回前n 个切线数,这点我已经实现,但我还希望代码完全避免冗余计算,我希望代码记住已经生成的切线数,并把它存储在一个变量中。这样如果 n 小于或等于已经生成的切线数的数量,代码就不应进行任何实际计算,而应立即返回全局变量的一个切片。
然后,对于一个大于已生成序列长度的数字,代码应用 [0] * (number - length) 来填充全局切线数序列,然后从下标length开始,紧接着从它上一次停止的地方继续,用已生成的项来生成下一个项,并把新项存入全局变量,完成后再返回该全局变量。
我已经用我的无限生成器和 itertools.islice 实现了恢复功能,但正如我所解释的那样,它并不像有限函数那样快。
我曾试图把代码包装成一个类,但它不起作用:
OPS = []
class Tangent_numbers:
def __init__(self, n: int = 0):
if not isinstance(n, int) or n < 0:
raise ValueError("argument n must be a nonnegative integer")
self.series = [1]
self.factorial = 1
self.length = 1
if n > 1:
self.extend(n)
def extend(self, n: int) -> None:
if not isinstance(n, int) or n <= 0:
raise ValueError("argument n must be a positive integer")
(series := self.series).extend([1] * (n - (length := self.length)))
fact = self.factorial
for i, k in enumerate(range(length, n), start=length):
series[k] = fact = fact * i
self.factorial = fact
start = length
for k in range(1, n):
start -= start > 0
b = start + 2
for a, c in enumerate(range(k + start, n), start=start):
series[c] = a * series[c - 1] + b * series[c]
OPS.append((k, a, b, c))
b += 1
self.length = (length, n)[length < n]
基本上,这个类将序列扩展到所需的长度,预先用阶乘进行填充,然后再对已经遇到的k 重新开始计算,这些k 是从1 开始的前length - 1个整数,对于这些数字它们并不是从 k 开始,而是从 k + start,start 逐次递减1,初始时设为 length。所以例如,如果 length 是6,那么偏移大于0 的k 的成对是 [(1, 5), (2, 4), (3, 3), (4, 2), (5, 1)],从下标等于length的位置开始,循环结构和下标与基准情况(A000182_fast)完全相同。
而且它不起作用:
In [39]: tan1 = Tangent_numbers()
In [40]: tan1.extend(32)
In [41]: tan1.series
Out[41]:
[1,
2,
16,
272,
7936,
353792,
22368256,
1903757312,
209865342976,
29088885112832,
4951498053124096,
1015423886506852352,
246921480190207983616,
70251601603943959887872,
23119184187809597841473536,
8713962757125169296170811392,
3729407703720529571097509625856,
1798651693450888780071750349094912,
970982810785059112379399707952152576,
583203324917310043943191641625494290432,
387635983772083031828014624002175135645696,
283727921907431909304183316295787837183229952,
227681379129930886488600284336316164603920777216,
199500252157859031027160499643195658166340757225472,
190169564657928428175235445073924928592047775873499136,
196535694915671808914892880726989984967498805398829268992,
219523439106761591280258358007964245987752702449505540243456,
264239411287900883270178745605712648488731058170551223831232512,
341838301335718580350174449297951396847081443826785448952307122176,
474090194351342155974522582010891145370303013658973457042263121068032,
703237958001393736999896827714634659411015090272684227831001161763127296,
1113255345330866700339746218047088280783690575394538814699382505942970007552]
In [42]: good_ops = OPS.copy()
In [43]: OPS.clear()
In [44]: tan2 = Tangent_numbers()
In [45]: tan2.extend(16)
In [46]: tan2.extend(32)
In [47]: bad_ops = OPS.copy()
In [48]: OPS.clear()
In [49]: tan2.series
Out[49]:
[1,
2,
16,
272,
7936,
353792,
22368256,
1903757312,
209865342976,
29088885112832,
4951498053124096,
1015423886506852352,
246921480190207983616,
70251601603943959887872,
23119184187809597841473536,
8713962757125169296170811392,
2904913076573872266203498954016294449446912,
4691853834639623592590771373697685743890071552,
5130358638144401639057729513343066448970081370112,
4915808771679000034913126880856135752591161113444352,
4532390344824939077129321365233314869866731663798566912,
4203608535388999900804817047351188778825742479186567626752,
4016075105391784162392416168597886026774711180596874221453312,
4006317907309964706114200066419926732731565593907176950571466752,
4206222056177413192157298446480182062802619739270396407035369357312,
4669250237392028942296888593639384792643032067338268907267585686372352,
5494543666053818716762261654491210913657871572416499541257406325993766912,
6862845392986161471745129519783028922048665371026568262534783173188021911552,
9102416468449937394166752814219701116933808763260683657119736372726144394330112,
12818533459855341263183886505243846034060346753317615985094136744918816751827812352,
19157416044307854900780627577605126931202540354102138452037843767399405302510530854912,
30362182612272673641771108733390016965056260526780787823865261029201116235800577941962752]
In [50]: set(good_ops) ^ set(bad_ops)
Out[50]: set()
In [51]: from collections import Counter
In [52]: Counter(good_ops) - Counter(bad_ops)
Out[52]: Counter()
In [53]: Counter(bad_ops) - Counter(good_ops)
Out[53]: Counter()
它在一次遍历中正确地生成了前32个切线数,符合预期,但在两遍遍历中并不能正确生成前32个切线数,恢复功能出错。但你可以清楚地看到,在这两种情况下循环结构和下标组合都是相同的,但结果却不同。
我不知道为什么会这样,相关的数学对我来说太难,我没能成功逆向推导所用的数学。我无法修复这个错误。
我该如何正确实现恢复功能?
解决方案
我已经解决了这个问题。
我使用以下代码来分析这个过程:
DATA = []
def A000182_analysis(n):
result = [1] * n
last = 1
for i, k in enumerate(range(1, n), start=1):
result[k] = last = last * i
DATA.clear()
for k in range(1, n):
b = 2
last = result[k - 1]
for a, c in enumerate(range(k, n)):
old = last
result[c] = last = a * last + b * (prev := result[c])
DATA.append((k, a, b, c, old, prev, last))
b += 1
return result
结果如下:
In [49]: A000182_analysis(6)
Out[49]: [1, 2, 16, 272, 7936, 353792]
In [50]: D6 = DATA.copy()
In [51]: A000182_analysis(7)
Out[51]: [1, 2, 16, 272, 7936, 353792, 22368256]
In [52]: D7 = DATA.copy()
In [53]: A000182_analysis(8)
Out[53]: [1, 2, 16, 272, 7936, 353792, 22368256, 1903757312]
In [54]: D8 = DATA.copy()
In [55]: D6
Out[55]:
[(1, 0, 2, 1, 1, 1, 2),
(1, 1, 3, 2, 2, 2, 8),
(1, 2, 4, 3, 8, 6, 40),
(1, 3, 5, 4, 40, 24, 240),
(1, 4, 6, 5, 240, 120, 1680),
(2, 0, 2, 2, 2, 8, 16),
(2, 1, 3, 3, 16, 40, 136),
(2, 2, 4, 4, 136, 240, 1232),
(2, 3, 5, 5, 1232, 1680, 12096),
(3, 0, 2, 3, 16, 136, 272),
(3, 1, 3, 4, 272, 1232, 3968),
(3, 2, 4, 5, 3968, 12096, 56320),
(4, 0, 2, 4, 272, 3968, 7936),
(4, 1, 3, 5, 7936, 56320, 176896),
(5, 0, 2, 5, 7936, 176896, 353792)]
In [56]: D7
Out[56]:
[(1, 0, 2, 1, 1, 1, 2),
(1, 1, 3, 2, 2, 2, 8),
(1, 2, 4, 3, 8, 6, 40),
(1, 3, 5, 4, 40, 24, 240),
(1, 4, 6, 5, 240, 120, 1680),
(1, 5, 7, 6, 1680, 720, 13440),
(2, 0, 2, 2, 2, 8, 16),
(2, 1, 3, 3, 16, 40, 136),
(2, 2, 4, 4, 136, 240, 1232),
(2, 3, 5, 5, 1232, 1680, 12096),
(2, 4, 6, 6, 12096, 13440, 129024),
(3, 0, 2, 3, 16, 136, 272),
(3, 1, 3, 4, 272, 1232, 3968),
(3, 2, 4, 5, 3968, 12096, 56320),
(3, 3, 5, 6, 56320, 129024, 814080),
(4, 0, 2, 4, 272, 3968, 7936),
(4, 1, 3, 5, 7936, 56320, 176896),
(4, 2, 4, 6, 176896, 814080, 3610112),
(5, 0, 2, 5, 7936, 176896, 353792),
(5, 1, 3, 6, 353792, 3610112, 11184128),
(6, 0, 2, 6, 353792, 11184128, 22368256)]
In [57]: D8
Out[57]:
[(1, 0, 2, 1, 1, 1, 2),
(1, 1, 3, 2, 2, 2, 8),
(1, 2, 4, 3, 8, 6, 40),
(1, 3, 5, 4, 40, 24, 240),
(1, 4, 6, 5, 240, 120, 1680),
(1, 5, 7, 6, 1680, 720, 13440),
(1, 6, 8, 7, 13440, 5040, 120960),
(2, 0, 2, 2, 2, 8, 16),
(2, 1, 3, 3, 16, 40, 136),
(2, 2, 4, 4, 136, 240, 1232),
(2, 3, 5, 5, 1232, 1680, 12096),
(2, 4, 6, 6, 12096, 13440, 129024),
(2, 5, 7, 7, 129024, 120960, 1491840),
(3, 0, 2, 3, 16, 136, 272),
(3, 1, 3, 4, 272, 1232, 3968),
(3, 2, 4, 5, 3968, 12096, 56320),
(3, 3, 5, 6, 56320, 129024, 814080),
(3, 4, 6, 7, 814080, 1491840, 12207360),
(4, 0, 2, 4, 272, 3968, 7936),
(4, 1, 3, 5, 7936, 56320, 176896),
(4, 2, 4, 6, 176896, 814080, 3610112),
(4, 3, 5, 7, 3610112, 12207360, 71867136),
(5, 0, 2, 5, 7936, 176896, 353792),
(5, 1, 3, 6, 353792, 3610112, 11184128),
(5, 2, 4, 7, 11184128, 71867136, 309836800),
(6, 0, 2, 6, 353792, 11184128, 22368256),
(6, 1, 3, 7, 22368256, 309836800, 951878656),
(7, 0, 2, 7, 22368256, 951878656, 1903757312)]
请注意,在每次调用中,起始值都是相同的,但后面的值不同。
究竟哪些值不同?
In [58]: set(D7) - set(D6)
Out[58]:
{(1, 5, 7, 6, 1680, 720, 13440),
(2, 4, 6, 6, 12096, 13440, 129024),
(3, 3, 5, 6, 56320, 129024, 814080),
(4, 2, 4, 6, 176896, 814080, 3610112),
(5, 1, 3, 6, 353792, 3610112, 11184128),
(6, 0, 2, 6, 353792, 11184128, 22368256)}
In [59]: set(D8) - set(D7)
Out[59]:
{(1, 6, 8, 7, 13440, 5040, 120960),
(2, 5, 7, 7, 129024, 120960, 1491840),
(3, 4, 6, 7, 814080, 1491840, 12207360),
(4, 3, 5, 7, 3610112, 12207360, 71867136),
(5, 2, 4, 7, 11184128, 71867136, 309836800),
(6, 1, 3, 7, 22368256, 309836800, 951878656),
(7, 0, 2, 7, 22368256, 951878656, 1903757312)}
In [60]: set(D8) - set(D6)
Out[60]:
{(1, 5, 7, 6, 1680, 720, 13440),
(1, 6, 8, 7, 13440, 5040, 120960),
(2, 4, 6, 6, 12096, 13440, 129024),
(2, 5, 7, 7, 129024, 120960, 1491840),
(3, 3, 5, 6, 56320, 129024, 814080),
(3, 4, 6, 7, 814080, 1491840, 12207360),
(4, 2, 4, 6, 176896, 814080, 3610112),
(4, 3, 5, 7, 3610112, 12207360, 71867136),
(5, 1, 3, 6, 353792, 3610112, 11184128),
(5, 2, 4, 7, 11184128, 71867136, 309836800),
(6, 0, 2, 6, 353792, 11184128, 22368256),
(6, 1, 3, 7, 22368256, 309836800, 951878656),
(7, 0, 2, 7, 22368256, 951878656, 1903757312)}
好吧,正如你所看到的,每一行有7 个数字,将它们标记为数列A、B、C、D、E、F、G。
我们有这样的规律,对于输入N,A的序列是1、2、3、4、5……N-1。对于每个A,B从 0开始,依次是0、1、2、3……,换句话说,B每次迭代增加1。对于每个A,C从 2开始,依次是2、3、4、5、6、7……,它也在每次迭代时增加1。D从 A开始并且每次迭代递增1。
B、C、D在同一个A 的循环中增加,它们在A 增加后都会重置。每个A 能持续多久直到增加?很显然,每个A 恰好持续N - A次迭代,并且对每个A,最后一个D 始终是N-1。
最后一个A 是N-1,意味着对于输入大于N 的情况,第一组A 是N,这些级别还没有被完全处理。
现在看看数字E、F、G。F无关紧要。请注意,对每个A,当前的G 是下一步的E。G是存储在下标D 处的值,D的值由下标D-1的元素计算得出,而下标D+1的元素则由下标D 的元素计算得到,即G。因此G 是下一个E。
从集合的差异来看,你会发现对于A >= current_length,它们是完整地计算出来的。但对于A < current_length,从D = length开始的元素尚未被处理,因此我们从D = length开始,我们需要对应A 的length-1的元素。所说的元素就是该A 的该下标已计算出的最后一个值。
因此我们只需要记住每个A 对应的G 的最后一个值,就可以实现恢复功能。任务完成。
class Tangent_numbers:
def __init__(self, n: int = 0):
if not isinstance(n, int) or n < 0:
raise ValueError("argument n must be a nonnegative integer")
self.series = [1]
self.factorial = 1
self.length = 1
self.lasts = []
if n > 1:
self.extend(n)
def extend(self, n: int) -> None:
if not isinstance(n, int) or n <= 0:
raise ValueError("argument n must be a positive integer")
if n > (length := self.length):
(series := self.series).extend([0] * (n - length))
fact = self.factorial
for i, k in enumerate(range(length, n), start=length):
series[k] = fact = fact * i
self.factorial = fact
i = 0
k = length + 1
lasts = [0] * (n - 1)
for j, last in zip(range(length - 1, 0, -1), self.lasts):
b = k
for a, c in zip(range(j, n), range(length, n)):
series[c] = last = a * last + b * series[c]
b += 1
lasts[i] = last
k -= 1
i += 1
for k in range(length, n):
b = 2
last = series[k - 1]
for a, c in enumerate(range(k, n)):
series[c] = last = a * last + b * series[c]
b += 1
lasts[i] = last
i += 1
self.lasts = lasts
self.length = n
In [61]: tan = Tangent_numbers()
In [62]: tan.extend(16)
In [63]: tan.extend(32)
In [64]: tan.extend(32)
In [65]: tan.series
Out[65]:
[1,
2,
16,
272,
7936,
353792,
22368256,
1903757312,
209865342976,
29088885112832,
4951498053124096,
1015423886506852352,
246921480190207983616,
70251601603943959887872,
23119184187809597841473536,
8713962757125169296170811392,
3729407703720529571097509625856,
1798651693450888780071750349094912,
970982810785059112379399707952152576,
583203324917310043943191641625494290432,
387635983772083031828014624002175135645696,
283727921907431909304183316295787837183229952,
227681379129930886488600284336316164603920777216,
199500252157859031027160499643195658166340757225472,
190169564657928428175235445073924928592047775873499136,
196535694915671808914892880726989984967498805398829268992,
219523439106761591280258358007964245987752702449505540243456,
264239411287900883270178745605712648488731058170551223831232512,
341838301335718580350174449297951396847081443826785448952307122176,
474090194351342155974522582010891145370303013658973457042263121068032,
703237958001393736999896827714634659411015090272684227831001161763127296,
1113255345330866700339746218047088280783690575394538814699382505942970007552]