ProCreations's picture
Audit released TableLong retrieval outputs
661b19b
Raw
History Blame Contribute Delete
2.4 kB
#!/usr/bin/env python3
"""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()