mediatok-player / entropy /context.py
Daankular's picture
Upload folder using huggingface_hub
20857b0 verified
Raw
History Blame Contribute Delete
5.29 kB
"""Context modelling and spatial prediction for neural tokens.
Instead of storing absolute token values, predict each token from its
spatial neighbours (left, top, top-left) and store only the residual.
Left neighbour
Top neighbour
Previous frame (temporal)
Residuals concentrate around zero == much lower entropy.
"""
def _zigzag(v: int) -> int:
return (v << 1) ^ (v >> 31)
def _unzigzag(z: int) -> int:
return (z >> 1) ^ -(z & 1)
def context_predict_2d(tokens_2d: list[list[int]]) -> list[int]:
"""Predict tokens using left + top neighbours in a 2D grid.
tokens_2d: list of rows, each row is a list of tokens.
Returns zigzag-encoded residuals.
Prediction: P(x,y) = (left + top) // 2
Residual: R(x,y) = token(x,y) - P(x,y)
First row: top unavailable, use only left neighbour.
First col: left unavailable, use only top neighbour.
[0,0]: stored as-is.
"""
if not tokens_2d or not tokens_2d[0]:
return []
rows = len(tokens_2d)
cols = len(tokens_2d[0])
residuals: list[int] = []
for r in range(rows):
for c in range(cols):
tok = tokens_2d[r][c]
if r == 0 and c == 0:
pred = 0
elif r == 0:
pred = tokens_2d[r][c - 1]
elif c == 0:
pred = tokens_2d[r - 1][c]
else:
pred = (tokens_2d[r][c - 1] + tokens_2d[r - 1][c]) // 2
residual = tok - pred
residuals.append(_zigzag(residual))
return residuals
def context_unpredict_2d(residuals_zigzag: list[int], rows: int, cols: int) -> list[list[int]]:
"""Reverse context_predict_2d: residuals -> original tokens."""
tokens_2d: list[list[int]] = [[0] * cols for _ in range(rows)]
idx = 0
for r in range(rows):
for c in range(cols):
res = _unzigzag(residuals_zigzag[idx])
if r == 0 and c == 0:
pred = 0
elif r == 0:
pred = tokens_2d[r][c - 1]
elif c == 0:
pred = tokens_2d[r - 1][c]
else:
pred = (tokens_2d[r][c - 1] + tokens_2d[r - 1][c]) // 2
tokens_2d[r][c] = pred + res
idx += 1
return tokens_2d
def context_predict_temporal(frame_tokens: list[list[int]]) -> list[int]:
"""Predict tokens from previous frame + spatial neighbours.
Prediction: P(t,x,y) = token(t-1, x, y) + spatial_correction
where spatial_correction = (left + top - top_left) // 3
This is a simple learned-adjacent predictor:
- Temporal prediction handles stationary / slow-moving content
- Spatial prediction handles edges / gradients
Returns zigzag residuals in row-major order.
"""
if not frame_tokens:
return []
rows = len(frame_tokens[0])
cols = len(frame_tokens[0][0])
residuals: list[int] = []
prev_frame: list[list[int]] = [[0] * cols for _ in range(rows)]
for t, frame in enumerate(frame_tokens):
for r in range(rows):
for c in range(cols):
tok = frame[r][c]
temporal_pred = prev_frame[r][c] if t > 0 else 0
if r == 0 and c == 0:
spatial_correction = 0
elif r == 0:
spatial_correction = (frame[r][c - 1] - prev_frame[r][c - 1]) // 2
elif c == 0:
spatial_correction = (frame[r - 1][c] - prev_frame[r - 1][c]) // 2
else:
left_delta = frame[r][c - 1] - prev_frame[r][c - 1]
top_delta = frame[r - 1][c] - prev_frame[r - 1][c]
spatial_correction = (left_delta + top_delta) // 2
pred = temporal_pred + spatial_correction
residual = tok - pred
residuals.append(_zigzag(residual))
prev_frame = [row[:] for row in frame]
return residuals
def context_unpredict_temporal(residuals_zigzag: list[int], rows: int, cols: int,
n_frames: int) -> list[list[list[int]]]:
"""Reverse context_predict_temporal."""
frames: list[list[list[int]]] = []
prev_frame: list[list[int]] = [[0] * cols for _ in range(rows)]
idx = 0
for _ in range(n_frames):
frame: list[list[int]] = [[0] * cols for _ in range(rows)]
for r in range(rows):
for c in range(cols):
res = _unzigzag(residuals_zigzag[idx])
temporal_pred = prev_frame[r][c]
if r == 0 and c == 0:
spatial_correction = 0
elif r == 0:
spatial_correction = (frame[r][c - 1] - prev_frame[r][c - 1]) // 2
elif c == 0:
spatial_correction = (frame[r - 1][c] - prev_frame[r - 1][c]) // 2
else:
left_delta = frame[r][c - 1] - prev_frame[r][c - 1]
top_delta = frame[r - 1][c] - prev_frame[r - 1][c]
spatial_correction = (left_delta + top_delta) // 2
pred = temporal_pred + spatial_correction
frame[r][c] = pred + res
idx += 1
frames.append(frame)
prev_frame = [row[:] for row in frame]
return frames