Generalize trace remapping for code workloads
This commit is contained in:
269
runs/frontier-s3-real-v0/trace_utils.py
Normal file
269
runs/frontier-s3-real-v0/trace_utils.py
Normal 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
|
||||
Reference in New Issue
Block a user