410 lines
15 KiB
Python
410 lines
15 KiB
Python
#!/usr/bin/env python3
|
|
"""Materialize controlled Qwen30 workload families from one exact trace cohort."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import copy
|
|
import csv
|
|
import hashlib
|
|
import json
|
|
import math
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
from typing import Any, Iterable
|
|
|
|
|
|
FIELDS = (
|
|
"arrived_at",
|
|
"num_prefill_tokens",
|
|
"num_decode_tokens",
|
|
"session_id",
|
|
"block_hash_ids",
|
|
)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Family:
|
|
name: str
|
|
shape: str
|
|
arrival: str
|
|
prefix: bool
|
|
|
|
|
|
FAMILIES = (
|
|
Family("w0-short-fixed-uniform-none", "short-fixed", "uniform", False),
|
|
Family("w1-mean-fixed-uniform-none", "mean-fixed", "uniform", False),
|
|
Family("w2-mean-fixed-trace-none", "mean-fixed", "trace", False),
|
|
Family("w3-heterogeneous-uniform-none", "heterogeneous", "uniform", False),
|
|
Family("w4-heterogeneous-trace-none", "heterogeneous", "trace", False),
|
|
Family("w5-heterogeneous-uniform-prefix", "heterogeneous", "uniform", True),
|
|
Family("w6-heterogeneous-trace-prefix", "heterogeneous", "trace", True),
|
|
)
|
|
|
|
|
|
def parse_args() -> argparse.Namespace:
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument(
|
|
"--source-public",
|
|
type=Path,
|
|
help="Optional Frontier CSV; otherwise project complete blocks from private rows.",
|
|
)
|
|
parser.add_argument("--source-private", type=Path, required=True)
|
|
parser.add_argument("--model", type=Path, required=True)
|
|
parser.add_argument("--output-root", type=Path, required=True)
|
|
parser.add_argument(
|
|
"--reference-decode-tokens-per-second", type=float, default=3064.0
|
|
)
|
|
parser.add_argument(
|
|
"--rho", type=float, action="append", default=None,
|
|
help="Normalized decode offered load; repeat for multiple levels.",
|
|
)
|
|
parser.add_argument("--requests", type=int)
|
|
return parser.parse_args()
|
|
|
|
|
|
def sha256(path: Path) -> str:
|
|
digest = hashlib.sha256()
|
|
with path.open("rb") as source:
|
|
for chunk in iter(lambda: source.read(1 << 20), b""):
|
|
digest.update(chunk)
|
|
return digest.hexdigest()
|
|
|
|
|
|
def vector_sha256(rows: Iterable[dict[str, Any]]) -> str:
|
|
digest = hashlib.sha256()
|
|
for row in rows:
|
|
digest.update(
|
|
json.dumps(
|
|
[
|
|
row["source_index"],
|
|
row["arrived_at"],
|
|
row["input_length"],
|
|
row["output_length"],
|
|
row["session_id"],
|
|
row["runtime_block_ids"],
|
|
],
|
|
separators=(",", ":"),
|
|
).encode()
|
|
)
|
|
digest.update(b"\n")
|
|
return digest.hexdigest()
|
|
|
|
|
|
def load_source(
|
|
public_path: Path | None, private_path: Path
|
|
) -> list[dict[str, Any]]:
|
|
private_rows = [json.loads(line) for line in private_path.open() if line.strip()]
|
|
if not private_rows:
|
|
raise ValueError("empty private source")
|
|
if public_path is None:
|
|
public_rows = [
|
|
{
|
|
"arrived_at": row["arrived_at"],
|
|
"num_prefill_tokens": row["input_length"],
|
|
"num_decode_tokens": row["output_length"],
|
|
"block_hash_ids": "|".join(
|
|
str(value)
|
|
for value in row["runtime_block_ids"][
|
|
: int(row["input_length"]) // 16
|
|
]
|
|
),
|
|
}
|
|
for row in private_rows
|
|
]
|
|
else:
|
|
with public_path.open(newline="") as source:
|
|
public_rows = list(csv.DictReader(source))
|
|
if not public_rows or len(public_rows) != len(private_rows):
|
|
raise ValueError("source public/private request count mismatch")
|
|
|
|
rows = []
|
|
for index, (public, private) in enumerate(
|
|
zip(public_rows, private_rows, strict=True)
|
|
):
|
|
input_length = int(public["num_prefill_tokens"])
|
|
output_length = int(public["num_decode_tokens"])
|
|
if input_length != int(private["input_length"]):
|
|
raise ValueError(f"source input-length drift at row {index}")
|
|
if output_length != int(private["output_length"]):
|
|
raise ValueError(f"source output-length drift at row {index}")
|
|
runtime_ids = [
|
|
int(value) for value in public["block_hash_ids"].split("|") if value
|
|
]
|
|
if len(runtime_ids) != input_length // 16:
|
|
raise ValueError(f"incomplete prefix-block projection at row {index}")
|
|
rows.append(
|
|
{
|
|
"source_index": int(private["source_index"]),
|
|
"source_arrived_at": float(public["arrived_at"]),
|
|
"input_length": input_length,
|
|
"output_length": output_length,
|
|
"session_id": int(private["session_id"]),
|
|
"runtime_block_ids": runtime_ids,
|
|
"body": copy.deepcopy(private["body"]),
|
|
}
|
|
)
|
|
if any(
|
|
right["source_arrived_at"] < left["source_arrived_at"]
|
|
for left, right in zip(rows, rows[1:])
|
|
):
|
|
raise ValueError("source arrival order drift")
|
|
return rows
|
|
|
|
|
|
def arrivals_for(
|
|
source_rows: list[dict[str, Any]], arrival: str, target_rate: float
|
|
) -> list[float]:
|
|
if target_rate <= 0 or not math.isfinite(target_rate):
|
|
raise ValueError("target request rate must be finite and positive")
|
|
if arrival == "uniform":
|
|
return [index / target_rate for index in range(len(source_rows))]
|
|
if arrival != "trace":
|
|
raise ValueError(f"unknown arrival family: {arrival}")
|
|
if len(source_rows) < 2:
|
|
raise ValueError("trace arrivals require at least two requests")
|
|
first = source_rows[0]["source_arrived_at"]
|
|
relative = [row["source_arrived_at"] - first for row in source_rows]
|
|
if relative[-1] <= 0:
|
|
raise ValueError("trace arrival window must be positive")
|
|
source_rate = (len(relative) - 1) / relative[-1]
|
|
scale = source_rate / target_rate
|
|
return [value * scale for value in relative]
|
|
|
|
|
|
def non_special_token_ids(model: Path, requests: int) -> list[int]:
|
|
from transformers import AutoTokenizer
|
|
|
|
tokenizer = AutoTokenizer.from_pretrained(model, trust_remote_code=True)
|
|
special = set(tokenizer.all_special_ids)
|
|
candidates = [token for token in range(tokenizer.vocab_size) if token not in special]
|
|
if len(candidates) < requests + 1:
|
|
raise ValueError("tokenizer lacks enough non-special token IDs")
|
|
return candidates[: requests + 1]
|
|
|
|
|
|
def build_family_rows(
|
|
source_rows: list[dict[str, Any]],
|
|
family: Family,
|
|
*,
|
|
target_rate: float,
|
|
fixed_token_ids: list[int],
|
|
) -> list[dict[str, Any]]:
|
|
arrivals = arrivals_for(source_rows, family.arrival, target_rate)
|
|
mean_input = round(sum(row["input_length"] for row in source_rows) / len(source_rows))
|
|
mean_output = round(sum(row["output_length"] for row in source_rows) / len(source_rows))
|
|
if family.shape == "short-fixed":
|
|
fixed_shape = (2048, 128)
|
|
elif family.shape == "mean-fixed":
|
|
fixed_shape = (mean_input, mean_output)
|
|
elif family.shape == "heterogeneous":
|
|
fixed_shape = None
|
|
else:
|
|
raise ValueError(f"unknown shape family: {family.shape}")
|
|
|
|
rows = []
|
|
for index, (source, arrived_at) in enumerate(
|
|
zip(source_rows, arrivals, strict=True)
|
|
):
|
|
if fixed_shape is None:
|
|
input_length = source["input_length"]
|
|
output_length = source["output_length"]
|
|
body = copy.deepcopy(source["body"])
|
|
body.update(
|
|
{
|
|
"min_tokens": output_length,
|
|
"max_tokens": output_length,
|
|
"ignore_eos": True,
|
|
}
|
|
)
|
|
else:
|
|
input_length, output_length = fixed_shape
|
|
body = copy.deepcopy(source["body"])
|
|
body.update(
|
|
{
|
|
"prompt": [
|
|
fixed_token_ids[index + 1],
|
|
*([fixed_token_ids[0]] * (input_length - 1)),
|
|
],
|
|
"min_tokens": output_length,
|
|
"max_tokens": output_length,
|
|
"ignore_eos": True,
|
|
}
|
|
)
|
|
if input_length + output_length > 40960:
|
|
raise ValueError(f"shape exceeds model limit at row {index}")
|
|
rows.append(
|
|
{
|
|
"source_index": source["source_index"],
|
|
"arrived_at": arrived_at,
|
|
"input_length": input_length,
|
|
"output_length": output_length,
|
|
"session_id": source["session_id"] if family.prefix else index,
|
|
"runtime_block_ids": (
|
|
list(source["runtime_block_ids"]) if family.prefix else []
|
|
),
|
|
"body": body,
|
|
}
|
|
)
|
|
return rows
|
|
|
|
|
|
def write_case(
|
|
root: Path,
|
|
rows: list[dict[str, Any]],
|
|
*,
|
|
family: Family,
|
|
rho: float,
|
|
target_rate: float,
|
|
reference_capacity: float,
|
|
source_public: Path | None,
|
|
source_private: Path,
|
|
) -> dict[str, Any]:
|
|
public_root = root / "public"
|
|
private_root = root / "private"
|
|
public_root.mkdir(parents=True, exist_ok=True)
|
|
private_root.mkdir(parents=True, exist_ok=True)
|
|
public_path = public_root / "frontier.csv"
|
|
private_path = private_root / "real_requests.jsonl"
|
|
|
|
with public_path.open("w", newline="") as output:
|
|
writer = csv.DictWriter(output, fieldnames=FIELDS, lineterminator="\n")
|
|
writer.writeheader()
|
|
for row in rows:
|
|
writer.writerow(
|
|
{
|
|
"arrived_at": f"{row['arrived_at']:.12f}",
|
|
"num_prefill_tokens": row["input_length"],
|
|
"num_decode_tokens": row["output_length"],
|
|
"session_id": row["session_id"],
|
|
"block_hash_ids": "|".join(
|
|
str(value) for value in row["runtime_block_ids"]
|
|
),
|
|
}
|
|
)
|
|
with private_path.open("w") as output:
|
|
for row in rows:
|
|
output.write(json.dumps(row, separators=(",", ":")) + "\n")
|
|
|
|
input_lengths = [row["input_length"] for row in rows]
|
|
output_lengths = [row["output_length"] for row in rows]
|
|
arrivals = [row["arrived_at"] for row in rows]
|
|
empirical_rate = (len(rows) - 1) / (arrivals[-1] - arrivals[0])
|
|
manifest = {
|
|
"schema": "frontier-workload-regime-v1",
|
|
"family": family.name,
|
|
"shape_contract": family.shape,
|
|
"arrival_contract": family.arrival,
|
|
"prefix_caching": family.prefix,
|
|
"requests": len(rows),
|
|
"rho": rho,
|
|
"reference_decode_tokens_per_second": reference_capacity,
|
|
"global_offered_request_rate": target_rate,
|
|
"empirical_interarrival_rate": empirical_rate,
|
|
"decode_offered_tokens_per_second": target_rate
|
|
* sum(output_lengths)
|
|
/ len(output_lengths),
|
|
"input_tokens": {
|
|
"mean": sum(input_lengths) / len(input_lengths),
|
|
"min": min(input_lengths),
|
|
"max": max(input_lengths),
|
|
},
|
|
"output_tokens": {
|
|
"mean": sum(output_lengths) / len(output_lengths),
|
|
"min": min(output_lengths),
|
|
"max": max(output_lengths),
|
|
},
|
|
"first_arrival_s": arrivals[0],
|
|
"last_arrival_s": arrivals[-1],
|
|
"source_public": (
|
|
str(source_public.resolve()) if source_public is not None else None
|
|
),
|
|
"source_public_sha256": (
|
|
sha256(source_public) if source_public is not None else None
|
|
),
|
|
"frontier_block_projection": (
|
|
"source_public_complete_blocks"
|
|
if source_public is not None
|
|
else "private_runtime_ids_first_floor_isl_div_16_blocks"
|
|
),
|
|
"source_private": str(source_private.resolve()),
|
|
"source_private_sha256": sha256(source_private),
|
|
"public_csv": str(public_path.resolve()),
|
|
"public_csv_sha256": sha256(public_path),
|
|
"private_jsonl": str(private_path.resolve()),
|
|
"private_jsonl_sha256": sha256(private_path),
|
|
"row_vector_sha256": vector_sha256(rows),
|
|
}
|
|
manifest_path = public_root / "manifest.json"
|
|
manifest_path.write_text(json.dumps(manifest, indent=2, sort_keys=True) + "\n")
|
|
return manifest
|
|
|
|
|
|
def main() -> None:
|
|
args = parse_args()
|
|
rhos = args.rho or [0.05, 0.25, 0.50, 0.90, 1.20]
|
|
if any(not math.isfinite(rho) or rho <= 0 for rho in rhos):
|
|
raise ValueError("rho values must be finite and positive")
|
|
if args.reference_decode_tokens_per_second <= 0:
|
|
raise ValueError("reference decode capacity must be positive")
|
|
|
|
source_rows = load_source(args.source_public, args.source_private)
|
|
if args.requests is not None:
|
|
if args.requests < 2 or args.requests > len(source_rows):
|
|
raise ValueError("requests must be between 2 and the source cohort size")
|
|
source_rows = source_rows[: args.requests]
|
|
fixed_token_ids = non_special_token_ids(args.model, len(source_rows))
|
|
|
|
cases = []
|
|
for family in FAMILIES:
|
|
if family.shape == "short-fixed":
|
|
mean_output = 128.0
|
|
elif family.shape == "mean-fixed":
|
|
mean_output = round(
|
|
sum(row["output_length"] for row in source_rows) / len(source_rows)
|
|
)
|
|
else:
|
|
mean_output = sum(row["output_length"] for row in source_rows) / len(
|
|
source_rows
|
|
)
|
|
for rho in rhos:
|
|
target_rate = (
|
|
rho * args.reference_decode_tokens_per_second / mean_output
|
|
)
|
|
rows = build_family_rows(
|
|
source_rows,
|
|
family,
|
|
target_rate=target_rate,
|
|
fixed_token_ids=fixed_token_ids,
|
|
)
|
|
rho_label = f"rho{rho:.2f}".replace(".", "p")
|
|
root = args.output_root / family.name / rho_label
|
|
cases.append(
|
|
write_case(
|
|
root,
|
|
rows,
|
|
family=family,
|
|
rho=rho,
|
|
target_rate=target_rate,
|
|
reference_capacity=args.reference_decode_tokens_per_second,
|
|
source_public=args.source_public,
|
|
source_private=args.source_private,
|
|
)
|
|
)
|
|
|
|
experiment_manifest = {
|
|
"schema": "frontier-workload-regime-suite-v1",
|
|
"reference_decode_tokens_per_second": args.reference_decode_tokens_per_second,
|
|
"rhos": rhos,
|
|
"families": [family.__dict__ for family in FAMILIES],
|
|
"cases": cases,
|
|
}
|
|
args.output_root.mkdir(parents=True, exist_ok=True)
|
|
output = args.output_root / "manifest.json"
|
|
output.write_text(json.dumps(experiment_manifest, indent=2, sort_keys=True) + "\n")
|
|
print(output)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|