"""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)