Modilify-Mk2-preview-mlx / modilify_mk2 /mlx_commit_policy.py
ydy9038074's picture
Publish Modilify Mk2 Preview MLX
e4f7326 verified
Raw History Blame Contribute Delete
10.6 kB
"""Native MLX schema25 confidence fusion and sample-independent commit policy."""
from __future__ import annotations
import math
from collections.abc import Sequence
from dataclasses import dataclass
import mlx.core as mx
JUMP_FAILURE_BUDGET = 2.0
FUSED_EPS = 1.0e-6
COMMIT_REASON_NONE = 0
COMMIT_REASON_NORMAL = 1
COMMIT_REASON_FORCED_JUMP = 2
COMMIT_REASON_TERMINAL = 3
def commit_target_confidence_bias(
target_confidence: float | None,
failure_budget: float = 0.2,
budget_safety_ratio: float = 0.85,
) -> float:
if target_confidence is None or target_confidence <= 0.0 or target_confidence >= 1.0:
return 0.0
target_failure = min(budget_safety_ratio * failure_budget, 1.0 - target_confidence)
target_failure = max(target_failure, 1.0e-4)
target_conf = 1.0 - target_failure
logit_c = math.log(target_conf / (1.0 - target_conf))
logit_p = math.log(target_confidence / (1.0 - target_confidence))
return max(logit_c - logit_p, 0.0)
def fused_commit_confidence(
proposal_confidence: mx.array,
token_entropy: mx.array,
*,
eps: float = FUSED_EPS,
entropy_weight: float = 1.0,
confidence_power: float = 2.0,
top_k: int | None = None,
min_p: float | None = None,
target_confidence: float | None = None,
failure_budget: float = 0.2,
) -> mx.array:
"""Match the schema25 excess-entropy sigmoid and checkpoint power."""
p = mx.clip(mx.nan_to_num(proposal_confidence.astype(mx.float32), nan=0.5), eps, 1.0 - eps)
entropy = mx.maximum(mx.nan_to_num(token_entropy.astype(mx.float32), nan=0.0), 0.0)
binary_entropy = -p * mx.log(p) - (1.0 - p) * mx.log1p(-p)
excess = mx.maximum(entropy - binary_entropy, 0.0)
k_eff = None
if top_k is not None and top_k > 0:
k_eff = mx.full(p.shape, float(top_k), dtype=mx.float32)
if min_p is not None and min_p > 0:
thresh = mx.maximum(float(min_p) * p, 1.0e-6)
k_min_p = mx.maximum((1.0 - p) / thresh, 1.0)
k_eff = k_min_p if k_eff is None else mx.minimum(k_eff, k_min_p)
if k_eff is not None:
max_excess = (1.0 - p) * mx.log(k_eff)
excess = mx.minimum(excess, max_excess)
bias = commit_target_confidence_bias(target_confidence, failure_budget=failure_budget)
fused_logit = mx.log(p) - mx.log1p(-p) + bias - entropy_weight * excess
return mx.clip(mx.sigmoid(fused_logit) ** confidence_power, eps, 1.0 - eps)
def fused_commit_failure_rate(
proposal_confidence: mx.array, token_entropy: mx.array, **kwargs: object
) -> mx.array:
return 1.0 - fused_commit_confidence(proposal_confidence, token_entropy, **kwargs)
def _terminal_mask(tokens: mx.array, stop_token_id: int | Sequence[int]) -> mx.array:
ids = (stop_token_id,) if isinstance(stop_token_id, int) else tuple(dict.fromkeys(stop_token_id))
if not ids:
raise ValueError("At least one stop token ID is required.")
matches = tokens == int(ids[0])
for value in ids[1:]:
matches = matches | (tokens == int(value))
return matches
def prefix_failure_commit_lengths(
failure_rate: mx.array,
*,
failure_budget: float,
valid_mask: mx.array | None = None,
) -> mx.array:
"""Longest contiguous valid prefix whose cumulative failure stays below budget."""
if failure_rate.ndim != 2:
raise ValueError("Failure rate must have shape [batch, canvas].")
if failure_budget <= 0:
raise ValueError("Commit failure budget must be positive.")
if valid_mask is None:
valid_mask = mx.ones(failure_rate.shape, mx.bool_)
if valid_mask.shape != failure_rate.shape:
raise ValueError("Commit validity mask must match failure rate.")
risk = mx.clip(failure_rate.astype(mx.float32), 0.0, 1.0) * valid_mask.astype(mx.float32)
allowed = (mx.cumsum(risk, axis=-1) < failure_budget) & (
mx.cumprod(valid_mask.astype(mx.int32), axis=-1).astype(mx.bool_)
)
return mx.sum(mx.cumprod(allowed.astype(mx.int32), axis=-1), axis=-1)
def first_committed_token_lengths(
proposal: mx.array,
commit_lengths: mx.array,
token_id: int | Sequence[int],
*,
positions: mx.array | None = None,
) -> mx.array:
if proposal.ndim != 2 or commit_lengths.shape != proposal.shape[:1]:
raise ValueError("Proposal and commit lengths must share a batch dimension.")
if positions is None:
positions = mx.arange(proposal.shape[1])[None, :]
elif positions.shape != (1, proposal.shape[1]):
raise ValueError("Commit positions must have shape [1, canvas].")
matches = _terminal_mask(proposal, token_id) & (positions < commit_lengths[:, None])
first = mx.min(mx.where(matches, positions, proposal.shape[1]), axis=-1)
return mx.minimum(mx.where(first < proposal.shape[1], first + 1, commit_lengths), commit_lengths)
def bounded_prefix_failure_commit_lengths(
committed_token_ids: mx.array,
failure_rate: mx.array,
*,
failure_budget: float,
remaining_lengths: mx.array,
stop_token_id: int | Sequence[int],
valid_mask: mx.array | None = None,
positions: mx.array | None = None,
) -> mx.array:
if committed_token_ids.shape != failure_rate.shape:
raise ValueError("Committed token IDs and failure rate must share [batch, canvas].")
if remaining_lengths.shape != committed_token_ids.shape[:1]:
raise ValueError("Remaining lengths must have shape [batch].")
lengths = prefix_failure_commit_lengths(
failure_rate, failure_budget=failure_budget, valid_mask=valid_mask
)
lengths = mx.minimum(lengths, mx.maximum(remaining_lengths, 0))
return first_committed_token_lengths(
committed_token_ids, lengths, stop_token_id, positions=positions
)
@dataclass(frozen=True)
class MLXCommitPolicyDecision:
normal_lengths: mx.array
commit_lengths: mx.array
commit_token_ids: mx.array
jump_rows: mx.array
ponder_steps: mx.array
stagnation_steps: mx.array
def select_commit_lengths(
sampled_token_ids: mx.array,
normal_failure_rate: mx.array,
previous_failure_rate: mx.array,
greedy_token_ids: mx.array,
jump_failure_rate: mx.array,
*,
ponder_steps: mx.array,
stagnation_steps: mx.array,
active_rows: mx.array,
remaining_lengths: mx.array,
failure_budget: float,
stop_token_id: int | Sequence[int],
stagnation_threshold: int,
min_progress: float,
max_ponder_steps: int | None = None,
valid_mask: mx.array | None = None,
) -> MLXCommitPolicyDecision:
"""Select normal commits or bounded greedy JUMP independently per row."""
if not (sampled_token_ids.shape == normal_failure_rate.shape
== previous_failure_rate.shape == greedy_token_ids.shape
== jump_failure_rate.shape):
raise ValueError("Sampled and greedy statistics must share [batch, canvas].")
if not (ponder_steps.shape == stagnation_steps.shape == active_rows.shape
== remaining_lengths.shape == sampled_token_ids.shape[:1]):
raise ValueError("Commit row inputs must share [batch].")
if min_progress < 0:
raise ValueError("Minimum progress must be nonnegative.")
canvas = normal_failure_rate.shape[1]
positions = mx.arange(canvas)[None, :]
normal = bounded_prefix_failure_commit_lengths(
sampled_token_ids, normal_failure_rate,
failure_budget=failure_budget, remaining_lengths=remaining_lengths,
stop_token_id=stop_token_id, valid_mask=valid_mask, positions=positions,
)
previous = prefix_failure_commit_lengths(
previous_failure_rate, failure_budget=failure_budget, valid_mask=valid_mask
)
frontier = mx.maximum(previous, normal) + 1
valid_lengths = (mx.sum(valid_mask.astype(mx.int32), axis=-1) if valid_mask is not None
else mx.full(frontier.shape, canvas, mx.int32))
frontier = mx.minimum(frontier, valid_lengths)
progress_mask = (positions < frontier[:, None]) & active_rows[:, None]
if valid_mask is not None:
progress_mask = progress_mask & valid_mask
weights = progress_mask.astype(mx.float32)
progress = mx.sum((previous_failure_rate.astype(mx.float32)
- normal_failure_rate.astype(mx.float32)) * weights, axis=-1) / mx.maximum(
mx.sum(weights, axis=-1), 1.0)
waiting = active_rows & (normal == 0)
next_ponder = mx.where(normal > 0, 0, ponder_steps + waiting.astype(mx.int32))
next_stagnation = mx.where(normal > 0, 0, stagnation_steps + waiting.astype(mx.int32))
jump = (normal == 0) & active_rows & (next_stagnation >= stagnation_threshold)
jump = jump & (progress <= min_progress)
if max_ponder_steps is not None and max_ponder_steps > 0:
jump = jump | ((normal == 0) & active_rows & (next_ponder >= max_ponder_steps))
jump_commit = bounded_prefix_failure_commit_lengths(
greedy_token_ids, jump_failure_rate,
failure_budget=JUMP_FAILURE_BUDGET, remaining_lengths=remaining_lengths,
stop_token_id=stop_token_id, valid_mask=valid_mask, positions=positions,
)
commit_token_ids = mx.where(jump[:, None], greedy_token_ids, sampled_token_ids)
committed = mx.where(active_rows, mx.where(jump, jump_commit, normal), 0)
jump = jump & (committed > 0)
next_ponder = mx.where(committed > 0, 0, next_ponder).astype(mx.int32)
next_stagnation = mx.where(committed > 0, 0, next_stagnation).astype(mx.int32)
return MLXCommitPolicyDecision(
normal_lengths=normal,
commit_lengths=committed,
commit_token_ids=commit_token_ids,
jump_rows=jump,
ponder_steps=next_ponder,
stagnation_steps=next_stagnation,
)
def infer_commit_reason(
commit_lengths: mx.array,
*,
jump_rows: mx.array | None = None,
commit_token_ids: mx.array | None = None,
terminal_token_ids: Sequence[int] = (),
) -> mx.array:
"""Reason codes consumed by the persistent writer at commit only."""
committed = commit_lengths > 0
reasons = mx.where(committed, COMMIT_REASON_NORMAL, COMMIT_REASON_NONE)
if jump_rows is not None:
reasons = mx.where(committed & jump_rows, COMMIT_REASON_FORCED_JUMP, reasons)
if commit_token_ids is not None and terminal_token_ids:
positions = mx.arange(commit_token_ids.shape[1])[None, :]
terminal = mx.any(
_terminal_mask(commit_token_ids, terminal_token_ids)
& (positions < commit_lengths[:, None]), axis=-1,
)
reasons = mx.where(committed & terminal, COMMIT_REASON_TERMINAL, reasons)
return reasons.astype(mx.int32)