Разложить число на кратчайшую сумму квадратов

Задача: пользователь вводит число не больше чем 10 000 000, программа должна разложить число на сумму квадратов так, чтобы этих квадратов было минимальное количество

Пример: 34 = 25+9

35 = 25+9+1

32 = 16+16

39 = 25+9+4+1

Честно, питон знаю на 6+, но сам алгоритм никак не соображу, помогите хотя бы с ним


Ответы (6 шт):

Автор решения: Miha_Taras
input_number = int(input('Введите число:'))

a = input_number - 45
b = 45
res = a + b 

print(a, '+', b,'=',res)

Типо так?

→ Ссылка
Автор решения: Aziz Umarov

Попробуйте подумать в такую сторону

S = summ(a[i]*i*i); где i=1..sqrt(10 000 000), a[i] - множители. Целевая функция min(summ(a[i]))

Понятно что самый длинный вариант это сумма единичек. Взяв за основу единички можно уменьшать целевую функцию добавив квадраты больших чисел. И не смотреть при переборе варианты когда целевая функция принимает большее значение чем текущая.

Идея такая что нужно схлопывать единички до минимума целевой функции. Это как вариант в какую сторону двигаться. У кого есть что нибудь получше пишите мне тоже интересно.

Вот статья в которой доказывается следующий факт.

Если x^2 + y^2 = n , то

(x+y)^2 + (x-y)^2 = 2*n

Имея данную формулу можно быстро уменьшить перебор для исходной целевой функции.

Ещё вот интересная теорема о том что наша целевая функция меньше либо равна 4.

Теорема Лагранжа о сумме четырёх квадратов утверждает, что

Всякое натуральное число можно представить в виде суммы четырёх квадратов целых чисел.

Круг поиска отсюда ещё сузился.

Так как их максимум четыре можно в четыре вложенных циклов получить как тут или вот пример реализации на С.

→ Ссылка
Автор решения: NEStenerus nester

Задачу решил, отталкивался от того о чем писал Aziz Umarov про теорему Лагранжа

import sys

num = int(input())
if int(num ** (1 / 2)) == num ** (1 / 2):
    print(1)
    sys.exit()

k = int(num ** (1 / 2))
i = 0
x = k
j = k
for _ in range(j):
    if i > 1:
        print(4)
        break
    x = int(num ** (1 / 2))
    j = x
    for _ in range(j):
        x = int(num ** (1 / 2))
        for _ in range(x):
            #print(i, j, x)
            if (i * i) + (j * j) + (x * x) == num:
                if i > 0:
                    print(3)
                else:
                    print(2)
                break
            x -= 1
        j -= 1
    i += 1

Для тех кому трудно понять, что делает код:

есть 3 множителя: i, j, x.

i= 0, а x и j = корню из самого большого числа, которое <= введенному

код перебирает все значения множителей так, чтобы i^2 + j^2+ x^2 == введенному числу

если такой случай найден, то если i всё ещё = 0, достаточно 2 множителей, если же i = 1, то множителей нужно 3.

Если же такой случай не найден, то нужно 4 множителя.

→ Ссылка
Автор решения: n1tr0xs

Как вариант, но мне кажется, что этот алгоритм довольно далек от оптимального:

from math import sqrt

def main():
    number = int(input('Enter the number: '))
    for k in range(1, 5):
        decompose = to_sum_of_squares(number, k)
        if decompose:
            print(decompose)
            break

def to_sum_of_squares(n:int, k:'squares count:int')->list:
    if (n < 0) or (k <= 0):
        return []
    maximum = round(sqrt(n))
    if n == maximum*maximum:
        return [n]
    for c in range(1, maximum+1):
        decomposition = to_sum_of_squares((n-c*c), k-1)
        if decomposition:
            return [c*c]+decomposition
        
if __name__ == '__main__':
    main()

Вот "перевод" на Python со статьи, которую указал @Aziz Umarov с некоторыми оптимизациямми вычислений и дополнениями (квадрат ли введенное число; выбор наикратчайшей суммы, а не первой попавшейся):

def main(N):
    sqrt = int(math.sqrt(N))
    if N == sqrt*sqrt:
        return [N]
    res = []
    Y = x = math.ceil(math.sqrt(N)/2)
    x_sq = x*x
    while x_sq <= N:
        while Y and (Y-1)*(Y-1)*3 >= N - x_sq:
            Y -= 1
        y=Z=Y
        y_sq = y*y
        while (y <= x) and (x_sq + y_sq <= N):
            while Z and (Z-1)*(Z-1)*2 >= N - x_sq - y_sq:
                Z -= 1
            z=t=Z
            z_sq = z*z
            while (z <= y) and (x_sq + y_sq + z_sq <= N):
                while t*t > N - x_sq - y_sq - z_sq:
                    t -= 1
                if x_sq + y_sq + z_sq + t*t == N:
                    r = [x_sq, y_sq, z_sq, t*t]
                    res.append(r)
                z += 1
                z_sq = z*z
            y += 1
            y_sq = y*y
        x += 1
        x_sq = x*x
        
    for r in res:
        while 0 in r:
            r.remove(0)
    res.sort(key=lambda x: len(x))
    return res[0]
→ Ссылка
Автор решения: Log Edge

Следующая программа, как мне кажется, вычисляет количество квадратов за O(N * sqrt(N))

def min_squares_sum(n):
    dp = [float('inf')] * (n + 1)
    dp[0] = 0

    for i in range(1, int(n**0.5) + 1):
        for j in range(i * i, n + 1):
            dp[j] = min(dp[j], dp[j - i * i] + 1)

    return dp[n]

def main():
    num = int(input("Введите число (не больше 10000000): "))

    if num > 10000000:
        print("Число слишком большое.")
    else:
        result = min_squares_sum(num)
        print(f"Минимальное количество квадратов для получения {num}: {result}")

if __name__ == "__main__":
    main()
→ Ссылка
Автор решения: Stanislav Volodarskiy

Решение за линейное время с использованием линейной памяти.

  1. Проверяется что число не полный квадрат (O(1)).
  2. Строится множество квадратов и проверяется что число не сумма двух квадратов (O(√n)).
  3. Строится множество сумм пар квадратов (O(n)).
  4. Проверяется что число не сумма трёх квадратов (O(√n)).
  5. Ищется представление числа в виде четырех квадратов (O(n)).
import math


def represent(n):
    n_sqrt = math.isqrt(n)

    if n_sqrt * n_sqrt == n:
        return n_sqrt,

    set1 = set(i * i for i in range(1, n_sqrt + 1))
    for i in set1:
        if n - i in set1:
            return math.isqrt(i), math.isqrt(n - i)

    def gen2():
        for i in set1:
            for j in set1:
                s = i + j
                if s <= n:
                    yield s

    set2 = set(gen2())

    for i in set1:
        if n - i in set2:
            return math.isqrt(i), *represent(n - i)

    for i in set2:
        if n - i in set2:
            return *represent(i), *represent(n - i)

    assert False


def main():
    n = int(input())
    print(n, '=', ' + '.join(f'{i}^2' for i in represent(n)))


main()

Три секунды для десяти миллионов:

$ time echo 9_999_999 | python represent-n-as-sum-of-squares-1.py 
9999999 = 1983^2 + 2111^2 + 725^2 + 1042^2

real  0m2.708s
user  0m2.644s
sys   0m0.060s

Меняя структуры данных можно улучшить константы и для памяти и для скорости. Второе множество заменено на массив байт:

import math


def represent(n):
    n_sqrt = math.isqrt(n)

    if n_sqrt * n_sqrt == n:
        return n_sqrt,

    lst1 = [i * i for i in range(1, n_sqrt + 1)]
    set1 = set(lst1)
    for i in lst1:
        if n - i in set1:
            return math.isqrt(i), math.isqrt(n - i)

    set2 = bytearray(n + 1)
    for i, v in enumerate(lst1):
        for j in range(i + 1):
            s = v + lst1[j]
            if s > n:
                break
            set2[s] = 1

    for i in lst1:
        if set2[n - i] != 0:
            return math.isqrt(i), *represent(n - i)

    for i in range(n + 1):
        if set2[i] != 0 and set2[n - i] != 0:
            return *represent(i), *represent(n - i)

    assert False


def main():
    n = int(input())
    print(n, '=', ' + '.join(f'{i}^2' for i in represent(n)))


main()

Полсекунды до десяти миллионов, четыре секунды до ста миллионов, сорок четыре до миллиарда. Для миллиарда нужен гигабайт памяти:

$ time echo 9_999_999 | python represent-n-as-sum-of-squares-3.py 
9999999 = 1^2 + 5^2 + 242^2 + 3153^2

real  0m0.436s
user  0m0.432s
sys   0m0.000s

$ time echo 99_999_999 | python represent-n-as-sum-of-squares-3.py 
99999999 = 1^2 + 1^2 + 9819^2 + 1894^2

real  0m4.088s
user  0m4.036s
sys   0m0.028s

$ time echo 999_999_999 | python represent-n-as-sum-of-squares-3.py 
999999999 = 2^2 + 3^2 + 6319^2 + 30985^2

real  0m43.468s
user  0m43.216s
sys   0m0.248s

P.S. Множество set2 можно не хранить в явном виде а порождать в виде монотонной последовательности чисел. Это уменьшит требования к памяти до √n. Время работы вырастет до n·log n. И константа подрастёт, конечно.

P.P.S. Если требуется вычислить количество слагаемых в разложении, но не нужно предъявлять само разложение, задача упрощается до O(√n) по времени и O(1) по памяти: Не получается дописать программу.

→ Ссылка