Number Theory & Algebra contents
Binary Exponentiation
Compute powers in O(log n) multiplications by squaring, and reuse the same trick for modular powers, matrices and more.
Binary exponentiation (also called exponentiation by squaring) computes for a non-negative integer using only multiplications, instead of the multiplications of the naive loop.
The trick works for any operation that is associative:
so it applies to modular multiplication, matrix multiplication, function composition (permutations) and more, as we see below.
The idea
Two facts are enough:
Write the exponent in binary. For example, , so
The powers are cheap to produce because each one is the square of the previous one:
The number has bits, so we compute that many squares and multiply together those whose bit in is set (here we skip because bit 1 of 13 is zero):
The total work is at most squarings plus multiplications: .
The same idea as a recurrence:
Implementation
The recursive version is a direct translation of the recurrence:
def binpow_rec(a, n):
if n == 0:
return 1
half = binpow_rec(a, n // 2)
return half * half * a if n % 2 else half * halfThe iterative version walks over the bits of from the lowest to the highest. a always holds for the current bit . It has the same complexity but avoids the function-call overhead, which matters in Python:
def binpow(a, n):
result = 1
while n > 0:
if n & 1: # this bit of n is set: multiply it in
result *= a
a *= a # a = a^(2^i) -> a^(2^(i+1))
n >>= 1
return result
assert binpow(3, 13) == 1_594_323
assert binpow(2, 100) == 2 ** 100
assert binpow_rec(7, 20) == 7 ** 20 == binpow(7, 20)Modular exponentiation
A very common task is computing (for example, to find a modular inverse). Because
we can reduce after every multiplication, so the numbers never exceed :
def binpow_mod(a, n, m):
a %= m
result = 1 % m
while n > 0:
if n & 1:
result = result * a % m
a = a * a % m
n >>= 1
return result
assert binpow_mod(3, 200, 1_000_000_007) == pow(3, 200, 1_000_000_007)
assert binpow_mod(2, 10**18, 998_244_353) == pow(2, 10**18, 998_244_353)If the exponent is enormous compared to the modulus, you can shrink it first. For a prime and , Fermat's little theorem gives ; for composite replace by Euler's (see Euler's totient function).
Applications
Any associative operation: matrix power
The loop above never used the fact that a is a number, only that multiplication is associative and has an identity. Replace numbers by matrices and you can raise a matrix to a power in for a matrix.
The classic use is Fibonacci numbers. The pair becomes under a fixed linear map, so is an entry of :
def mat_mul(A, B, mod):
n, m, p = len(A), len(B), len(B[0])
return [
[sum(A[i][k] * B[k][j] for k in range(m)) % mod for j in range(p)]
for i in range(n)
]
def mat_pow(M, n, mod):
size = len(M)
result = [[int(i == j) for j in range(size)] for i in range(size)] # identity
while n > 0:
if n & 1:
result = mat_mul(result, M, mod)
M = mat_mul(M, M, mod)
n >>= 1
return result
def fib(n, mod=10**9 + 7):
return mat_pow([[1, 1], [1, 0]], n, mod)[0][1]
assert [fib(i) for i in range(10)] == [0, 1, 1, 2, 3, 5, 8, 13, 21, 34]
assert fib(90, 10**30) == 2880067194370816120Number of paths of length in a graph
If is the adjacency matrix of a directed graph, then is the number of paths with exactly edges from to . This costs . Replacing "multiply and add" by "add and take the minimum" gives the shortest path that uses exactly edges.
Applying a permutation times
Composing a permutation with itself is associative, so we can raise it to the -th power by squaring:
def apply_permutation(sequence, permutation):
return [sequence[p] for p in permutation]
def permute(sequence, permutation, k):
while k > 0:
if k & 1:
sequence = apply_permutation(sequence, permutation)
permutation = apply_permutation(permutation, permutation)
k >>= 1
return sequence
seq = list("abcde")
perm = [1, 2, 0, 4, 3] # a 3-cycle and a 2-cycle
once = apply_permutation(seq, perm)
assert permute(seq, perm, 1) == once
assert permute(seq, perm, 6) == seq # lcm(3, 2) = 6 returns to the start
assert permute(seq, perm, 7) == onceThe cost is . It can be done in by decomposing the permutation into cycles and taking modulo each cycle length.
Multiplying two numbers modulo without overflow
In C++ the product may not fit into 64 bits even though do. The fix is to apply the same doubling idea to addition:
def mul_mod(a, b, m):
result = 0
a %= m
while b > 0:
if b & 1:
result = (result + a) % m
a = (a + a) % m
b >>= 1
return result
m = 10**18 + 9
assert mul_mod(123456789012345678, 987654321098765432, m) == 123456789012345678 * 987654321098765432 % mCommon mistakes
- Forgetting to reduce
afirst. Inbinpow_mod, start witha %= m, otherwise the first squaring works with a number larger than (correct, but slower and, in C++, dangerous). m = 1. Every result must be0; that is why the code starts from1 % m, not1.- Negative exponents. The loops assume . For , compute the modular inverse first (in Python 3.8+:
pow(a, -n, m)does exactly that when ).
Practice problems
- UVa 1230 - MODEX
- UVa 374 - Big Mod
- UVa 11029 - Leading and Trailing
- Codeforces - Parking Lot
- leetcode - Count good numbers
- Codechef - Chef and Riffles
- Codeforces - Decoding Genome
- Codeforces - Neural Network Country
- Codeforces - Magic Gems
- SPOJ - The last digit
- SPOJ - Locker
- LA - 3722 Jewel-eating Monsters
- SPOJ - Just add it
- Codeforces - Stairs and Lines