如何让我用来生成正切数的这个类具备可恢复性?

编程语言 2026-07-12

我决定把实现代码用于生成切线数,作为一次自我设定的编程挑战。我已经成功实现了,但并没有达到我期望的效率。

什么是切线数?它们就是下面这个级数:

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 + startstart 逐次递减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]
站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。

相关文章