Generalize trace remapping for code workloads

This commit is contained in:
2026-07-24 00:09:50 +08:00
parent a9e88f14de
commit 029c8991b6
3 changed files with 976 additions and 0 deletions

View File

@@ -0,0 +1,269 @@
#!/usr/bin/env python3
"""Shared, dependency-free helpers for the S3-real preflight."""
from __future__ import annotations
import hashlib
import json
import math
from pathlib import Path
from typing import Any, Iterable, Iterator, Sequence
SOURCE_BLOCK_SIZE = 64
TARGET_BLOCK_SIZE = 16
SUBBLOCKS_PER_SOURCE = SOURCE_BLOCK_SIZE // TARGET_BLOCK_SIZE
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as stream:
for chunk in iter(lambda: stream.read(1 << 20), b""):
digest.update(chunk)
return digest.hexdigest()
def write_json(path: Path, payload: Any) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n")
def iter_jsonl(path: Path) -> Iterator[dict[str, Any]]:
with path.open() as stream:
for line_number, line in enumerate(stream, 1):
if not line.strip():
continue
try:
row = json.loads(line)
except json.JSONDecodeError as error:
raise ValueError(f"{path}:{line_number}: invalid JSON") from error
if not isinstance(row, dict):
raise ValueError(f"{path}:{line_number}: row must be an object")
yield row
def timestamp(row: dict[str, Any]) -> float:
return float(row["timestamp"])
def input_length(row: dict[str, Any]) -> int:
return int(row["input_length"])
def output_length(row: dict[str, Any]) -> int:
return max(1, int(row["output_length"]))
def parse_hash_ids(value: Any) -> list[int]:
if isinstance(value, list):
raw = value
elif isinstance(value, str):
stripped = value.strip()
if not stripped:
return []
if stripped.startswith("["):
raw = json.loads(stripped)
else:
delimiter = "|" if "|" in stripped else ","
raw = [part.strip() for part in stripped.split(delimiter) if part.strip()]
elif value is None:
return []
else:
raw = [value]
try:
return [int(item) for item in raw]
except (TypeError, ValueError) as error:
raise ValueError(f"hash_ids must contain integers, got {value!r}") from error
def percentile(values: Sequence[float | int], fraction: float) -> float | None:
if not values:
return None
if not 0 <= fraction <= 1:
raise ValueError("percentile fraction must be in [0, 1]")
ordered = sorted(float(value) for value in values)
if len(ordered) == 1:
return ordered[0]
position = (len(ordered) - 1) * fraction
lower = math.floor(position)
upper = math.ceil(position)
if lower == upper:
return ordered[lower]
return ordered[lower] * (upper - position) + ordered[upper] * (position - lower)
def distribution(values: Sequence[float | int]) -> dict[str, float | int | None]:
return {
"count": len(values),
"min": min(values) if values else None,
"p50": percentile(values, 0.50),
"p90": percentile(values, 0.90),
"p95": percentile(values, 0.95),
"p99": percentile(values, 0.99),
"max": max(values) if values else None,
"mean": sum(values) / len(values) if values else None,
}
def common_prefix_length(left: Sequence[Any], right: Sequence[Any]) -> int:
count = 0
for a, b in zip(left, right):
if a != b:
break
count += 1
return count
def zigzag(value: int) -> int:
"""Map every signed integer injectively to a non-negative integer."""
return 2 * value if value >= 0 else -2 * value - 1
def expanded_block_id(
source_hash: int,
subblock_index: int,
*,
subblocks_per_source: int = SUBBLOCKS_PER_SOURCE,
) -> int:
if not 0 <= subblock_index < subblocks_per_source:
raise ValueError(f"subblock index out of range: {subblock_index}")
return zigzag(int(source_hash)) * subblocks_per_source + subblock_index
def expand_hash_ids(
source_hashes: Sequence[int],
num_tokens: int,
*,
source_block_size: int = SOURCE_BLOCK_SIZE,
target_block_size: int = TARGET_BLOCK_SIZE,
) -> list[int]:
if num_tokens <= 0:
raise ValueError(f"input_length must be positive, got {num_tokens}")
if source_block_size <= 0 or target_block_size <= 0:
raise ValueError("block sizes must be positive")
if source_block_size % target_block_size:
raise ValueError("source block size must be divisible by target block size")
subblocks_per_source = source_block_size // target_block_size
required_source_blocks = math.ceil(num_tokens / source_block_size)
if len(source_hashes) != required_source_blocks:
raise ValueError(
f"strict {source_block_size}-block contract failed: "
f"input_length={num_tokens} requires {required_source_blocks} source blocks, "
f"got {len(source_hashes)}"
)
target_blocks = math.ceil(num_tokens / target_block_size)
expanded = [
expanded_block_id(
source_hash,
subblock,
subblocks_per_source=subblocks_per_source,
)
for source_hash in source_hashes
for subblock in range(subblocks_per_source)
]
return expanded[:target_blocks]
def token_payload(tokens: Sequence[int]) -> bytes:
return len(tokens).to_bytes(2, "little") + b"".join(
int(token).to_bytes(4, "little", signed=False) for token in tokens
)
def token_block_identity_records(
token_ids: Sequence[int], block_size: int = TARGET_BLOCK_SIZE
) -> list[tuple[int, bytes]]:
"""Return parent-sensitive block IDs and independent collision witnesses."""
if block_size <= 0:
raise ValueError("block_size must be positive")
parent = b"FRONTIER_EXACT_TRACE_ROOT"
records: list[tuple[int, bytes]] = []
for start in range(0, len(token_ids), block_size):
payload = token_payload(token_ids[start : start + block_size])
identity_input = parent + b"\0" + payload
parent = hashlib.blake2b(identity_input, digest_size=16).digest()
records.append(
(
int.from_bytes(parent, "big", signed=False),
hashlib.sha256(identity_input).digest(),
)
)
return records
def token_block_identities(
token_ids: Sequence[int], block_size: int = TARGET_BLOCK_SIZE
) -> list[int]:
return [identity for identity, _ in token_block_identity_records(token_ids, block_size)]
def token_block(block_id: int, *, vocab_size: int, token_offset: int = 1024) -> list[int]:
"""Encode one block identity injectively as exactly sixteen valid token IDs."""
base = vocab_size - token_offset
if base < 256:
raise ValueError("vocab has too few non-reserved token IDs")
value = int(block_id)
if value < 0:
raise ValueError("expanded block IDs must be non-negative")
digits = []
for _ in range(TARGET_BLOCK_SIZE):
digits.append(token_offset + value % base)
value //= base
if value:
raise ValueError("block identity exceeds the injective 16-token code space")
return digits
def synthetic_tokens(
block_ids: Sequence[int], num_tokens: int, *, vocab_size: int, token_offset: int = 1024
) -> list[int]:
tokens = [
token
for block_id in block_ids
for token in token_block(block_id, vocab_size=vocab_size, token_offset=token_offset)
][:num_tokens]
if len(tokens) != num_tokens:
raise ValueError(f"synthetic token length mismatch: {len(tokens)} != {num_tokens}")
return tokens
def stable_int(value: Any) -> int:
raw = json.dumps(value, sort_keys=True, separators=(",", ":")).encode()
return int.from_bytes(hashlib.sha256(raw).digest()[:8], "big") & ((1 << 63) - 1)
def resolve_session_root(row: dict[str, Any], root_of: dict[Any, Any]) -> Any:
chat = row.get("chat_id")
parent = row.get("parent_chat_id")
has_parent = parent not in (None, "", -1, "-1")
root = root_of.get(parent, parent) if has_parent else chat
if chat is not None:
root_of[chat] = root
return root
def session_uniform(*, seed: int, window_id: str, session_root: Any) -> float:
payload = json.dumps(
{"seed": seed, "window_id": window_id, "session_root": session_root},
sort_keys=True,
separators=(",", ":"),
).encode()
digest = hashlib.blake2b(payload, digest_size=8).digest()
return int.from_bytes(digest, "big") / float(1 << 64)
def session_id(row: dict[str, Any]) -> int:
return stable_int(
row.get("session_root", row.get("chat_id", row.get("parent_chat_id", "missing-chat")))
)
def batched(iterable: Iterable[Any], size: int) -> Iterator[list[Any]]:
batch: list[Any] = []
for item in iterable:
batch.append(item)
if len(batch) == size:
yield batch
batch = []
if batch:
yield batch