| |
| """Exact CPU scope sweep for the periodic table MI mechanism.""" |
|
|
| from __future__ import annotations |
|
|
| import json |
| import math |
| from fractions import Fraction |
|
|
|
|
| def mutual_information(joint: list[list[Fraction]]) -> float: |
| left = [sum(row) for row in joint] |
| right = [sum(joint[i][j] for i in range(len(joint))) for j in range(len(joint))] |
| result = 0.0 |
| for i, row in enumerate(joint): |
| for j, value in enumerate(row): |
| if value: |
| result += float(value) * math.log(float(value / (left[i] * right[j]))) |
| return result |
|
|
|
|
| def one_dimension(m: int) -> dict: |
| diagonal = Fraction(7, 10) |
| off_diagonal = Fraction(3, 10 * (m - 1)) |
| columns = [ |
| [diagonal if i == j else off_diagonal for i in range(m)] |
| for j in range(m) |
| ] |
| same = [ |
| [sum(columns[column][i] * columns[column][j] for column in range(m)) / m for j in range(m)] |
| for i in range(m) |
| ] |
| equal = [[Fraction(1, m * m) for _ in range(m)] for _ in range(m)] |
| lags = [m * k for k in (1, 2, 4, 8, 16, 32, 64)] |
| return { |
| "m": m, |
| "lags": lags, |
| "exact_joint_nonzero_cells": sum(value != 0 for row in same for value in row), |
| "same_column_mi_nats": mutual_information(same), |
| "equal_column_control_nats": mutual_information(equal), |
| } |
|
|
|
|
| def power_law_slopes() -> dict[str, float]: |
| horizons = [1, 10, 100, 1000, 10000, 100000] |
| output = {} |
| for alpha in (0.4, 0.7, 1.0): |
| x = [math.log10(k) for k in horizons] |
| y = [math.log10(0.166613495469 / (0.8 * k ** (-alpha))) for k in horizons] |
| xbar, ybar = sum(x) / len(x), sum(y) / len(y) |
| output[str(alpha)] = sum((a - xbar) * (b - ybar) for a, b in zip(x, y)) / sum( |
| (a - xbar) ** 2 for a in x |
| ) |
| return output |
|
|
|
|
| def main() -> None: |
| rows = [one_dimension(m) for m in (2, 3, 4, 5, 8, 16)] |
| if any(row["same_column_mi_nats"] <= 0 for row in rows): |
| raise AssertionError(rows) |
| if any(abs(row["equal_column_control_nats"]) > 1e-15 for row in rows): |
| raise AssertionError(rows) |
| slopes = power_law_slopes() |
| if any(abs(float(alpha) - slope) > 1e-12 for alpha, slope in slopes.items()): |
| raise AssertionError(slopes) |
| print(json.dumps({"periodic_scope": rows, "power_law_slopes": slopes}, sort_keys=True)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|