back to table · edit · history · where entries came from · files · download
8515 bytes, as of the version from 2026-09-27 17:26 (current). Recorded here, not run.
"""Independent matrix-capacity certificates, using MPFI rather than Arb.
For any exact input distribution pi, form Q=pi W. Then
sum_x pi_x D(W_x || Q) <= C <= max_x D(W_x || Q).
The lower bound is achievable mutual information. The upper bound follows
from sum_x pi'_x D(W_x||Q) = I(pi',W) + D(pi' W||Q) for every pi'.
This is the information-radius characterization in Polyanskiy and Wu,
MIT 6.441 (2016), sections 4.4-4.5, Corollaries 4.2-4.3, printed pp. 45-46.
No formula or optimizer from channel_generate is imported. Two-input matrices
use sign-certified rational bisection on D(W_0||Q)-D(W_1||Q). Larger matrices
try the uniform prior and fail unless the resulting primal/dual gap closes.
Zero transition entries are skipped exactly, including zero output columns.
MPFR endpoints are extracted with exact_rational(), never QQ(endpoint),
which can approximate the endpoint and destroy the outward enclosure.
"""
from fractions import Fraction
import numberdb.sage
from sage.rings.rational_field import QQ
from sage.rings.real_mpfi import RealIntervalField
SOURCE_URL = ('https://ocw.mit.edu/courses/6-441-information-theory-spring-2016/'
'5d8f16adc3385c9ff2975b121bd620e4_MIT6_441S16_course_notes.pdf')
def transition_matrix(channel, shape):
values = [Fraction(part) for part in shape.split(',')]
if channel in ('bsc', 'z'):
if len(values) != 1:
raise ValueError('Incorrect binary shape')
probability = values[0]
matrix = ([[1 - probability, probability], [probability, 1 - probability]]
if channel == 'bsc' else [[1, 0], [probability, 1 - probability]])
elif channel == 'bac':
first, second = values
matrix = [[1 - first, first], [second, 1 - second]]
elif channel == 'qsc':
alphabet, probability = values
if alphabet.denominator != 1 or alphabet < 2:
raise ValueError('Invalid alphabet')
matrix = [[1 - probability if row == column else probability / (alphabet - 1)
for column in range(int(alphabet))] for row in range(int(alphabet))]
else:
raise ValueError('Unknown channel')
return validate_matrix(matrix)
def validate_matrix(matrix):
rows = [[QQ(value) for value in row] for row in matrix]
if not rows or not rows[0] or any(len(row) != len(rows[0]) for row in rows):
raise ValueError('Empty or ragged matrix')
if any(sum(row) != 1 or any(value < 0 for value in row) for row in rows):
raise ValueError('Not a stochastic matrix')
return rows
def exact_matrix_capacity(matrix, unit):
matrix = validate_matrix(matrix)
if all(row == matrix[0] for row in matrix):
return QQ(0), 'Identical conditional output distributions'
supports = [{index for index, value in enumerate(row) if value} for row in matrix]
disjoint = all(not left.intersection(right)
for index, left in enumerate(supports) for right in supports[index + 1:])
count = len(matrix)
if disjoint and unit == 'bits' and count & (count - 1) == 0:
return QQ(count.bit_length() - 1), 'Disjoint output supports; input recoverable'
row_permutations = all(sorted(row) == sorted(matrix[0]) for row in matrix)
column_sums = [sum(row[column] for row in matrix)
for column in range(len(matrix[0]))]
if unit == 'bits' and row_permutations and len(set(column_sums)) == 1:
coefficients = {}
logarithms = [(QQ(len(matrix[0])), QQ(1))]
logarithms.extend((value, value) for value in matrix[0] if value)
for argument, weight in logarithms:
for prime, exponent in argument.numerator().factor():
coefficients[prime] = coefficients.get(prime, QQ(0)) + weight * exponent
for prime, exponent in argument.denominator().factor():
coefficients[prime] = coefficients.get(prime, QQ(0)) - weight * exponent
if all(coefficient == 0 for prime, coefficient in coefficients.items() if prime != 2):
return coefficients.get(2, QQ(0)), (
'Row permutations and equal column sums give uniform optimal output; '
'exact rational prime-log coefficients cancel except log(2)')
return None, None
def exact_endpoints(interval):
try:
lower = interval.lower().exact_rational()
upper = interval.upper().exact_rational()
except (ValueError, TypeError, OverflowError) as error:
raise ArithmeticError('Nonfinite interval endpoint') from error
if lower > upper:
raise ArithmeticError('Reversed interval endpoints')
return lower, upper
def matrix_certificate(matrix, digits=100):
matrix = validate_matrix(matrix)
field = RealIntervalField(numberdb.sage.bits(digits, losing=192))
target = QQ(1) / 10**(digits + 12)
row_terms = [sum((field(value) * field(value).log() for value in row if value), field(0))
for row in matrix]
lower_input, upper_input = QQ(0), QQ(1)
for iteration in range(numberdb.sage.bits(digits, losing=64)):
if len(matrix) == 2:
input_zero = (lower_input + upper_input) / 2
prior = [input_zero, 1 - input_zero]
else:
prior = [QQ(1) / len(matrix)] * len(matrix)
output = [sum(prior[index] * row[column] for index, row in enumerate(matrix))
for column in range(len(matrix[0]))]
log_output = [field(value).log() if value else None for value in output]
divergences = [row_terms[index] - sum(
(field(value) * log_output[column] for column, value in enumerate(row) if value),
field(0)) for index, row in enumerate(matrix)]
mutual = sum((field(weight) * value for weight, value in zip(prior, divergences)), field(0))
lower = mutual.lower()
upper = max(value.upper() for value in divergences)
enclosure = field(lower, upper)
for value in divergences:
exact_endpoints(value)
exact_lower, exact_upper = exact_endpoints(enclosure)
if exact_upper - exact_lower < target:
return enclosure, {
'iterations': iteration + 1,
'input_prior': [str(weight) for weight in prior],
'gap_upper': str(field(upper) - field(lower)),
'capacity_lower_nats': str(exact_lower),
'capacity_upper_nats': str(exact_upper),
'capacity_gap_nats': str(exact_upper - exact_lower),
'endpoint_conversion': 'MPFR exact_rational()',
'method': 'MPFI mutual-information lower / information-radius upper bound',
}
if len(matrix) != 2:
raise ArithmeticError('Uniform prior did not certify capacity')
derivative = divergences[0] - divergences[1]
derivative_lower, derivative_upper = exact_endpoints(derivative)
if derivative_lower > 0:
lower_input = input_zero
elif derivative_upper < 0:
upper_input = input_zero
else:
raise ArithmeticError('Derivative sign unresolved before gap target')
raise ArithmeticError('Matrix-capacity certificate exceeded iteration limit')
def certify_written(matrix, unit, written, certificate):
from decimal import Decimal
if unit not in ('nats', 'bits'):
raise ValueError('Invalid unit')
bound = certificate if unit == 'nats' else certificate / certificate.parent()(2).log()
lower, upper = exact_endpoints(bound)
exact, reason = exact_matrix_capacity(matrix, unit)
if '.' not in written and 'e' not in written.lower():
if exact is None or QQ(written) != exact or not lower <= exact <= upper:
raise AssertionError('Exact entry lacks a matching structural matrix proof')
return {'exact_proof': reason, 'exact_value': str(exact)}
if exact is not None:
raise AssertionError('Exact capacity was unnecessarily decimalized')
decimal = Decimal(written)
centre = QQ(Fraction(decimal))
radius = QQ(10)**decimal.as_tuple().exponent
if not (centre - radius <= lower <= upper <= centre + radius):
raise AssertionError('Stored interval does not contain the independent certificate')
if not upper - lower < radius / 10**8:
raise AssertionError('Certificate is too wide to check the stored digits')
return {'stored_interval_contains_matrix_certificate': True,
'certificate_width_below_1e_minus_8_ulp': True}