Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 67 additions & 0 deletions src/kvcache_upper_bound/cli/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
load_bucket_analysis_config,
write_bucket_outputs,
)
from kvcache_upper_bound.synthetic import SyntheticTraceConfig, generate_synthetic_trace
from kvcache_upper_bound.verification import (
build_bucket_audit_report,
write_bucket_audit_outputs,
Expand Down Expand Up @@ -185,6 +186,38 @@ def main() -> int:
help="Allow replay-only synthetic hash_ids when benchmark records do not provide them",
)

generate_trace_parser = subparsers.add_parser(
"generate-trace",
help="Generate a parametric synthetic trace for KVCache analysis",
)
generate_trace_parser.add_argument(
"--sessions", type=int, required=True, help="Number of sessions"
)
generate_trace_parser.add_argument(
"--turns", type=int, required=True, help="Turns per session"
)
generate_trace_parser.add_argument(
"--shared-prefix-blocks", type=int, required=True,
help="Number of shared prefix blocks",
)
generate_trace_parser.add_argument(
"--new-blocks-per-turn", type=int, required=True,
help="Average new blocks per turn",
)
generate_trace_parser.add_argument(
"--block-size", type=int, default=16, help="Tokens per block (default: 16)"
)
generate_trace_parser.add_argument(
"--prefix-diversity", type=float, default=0.3,
help="Prefix diversity 0.0-1.0 (default: 0.3)",
)
generate_trace_parser.add_argument(
"--seed", type=int, default=None, help="Random seed for reproducibility"
)
generate_trace_parser.add_argument(
"--output", required=True, help="Output JSONL file path"
)

args = parser.parse_args()
if args.command == "list-datasets":
return _run_list_datasets()
Expand All @@ -200,6 +233,8 @@ def main() -> int:
return _run_convert_conversation_dataset(args)
if args.command == "convert-benchmark-results":
return _run_convert_benchmark_results(args)
if args.command == "generate-trace":
return _run_generate_trace(args)
raise ValueError(f"unsupported command: {args.command}")


Expand Down Expand Up @@ -487,6 +522,38 @@ def _run_convert_benchmark_results(args: argparse.Namespace) -> int:
return 0


def _run_generate_trace(args: argparse.Namespace) -> int:
config = SyntheticTraceConfig(
num_sessions=args.sessions,
turns_per_session=args.turns,
shared_prefix_blocks=args.shared_prefix_blocks,
avg_new_blocks_per_turn=args.new_blocks_per_turn,
block_size=args.block_size,
prefix_diversity=args.prefix_diversity,
seed=args.seed,
)
records = generate_synthetic_trace(config)
output_path = Path(args.output)
output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, "w", encoding="utf-8") as f:
for record in records:
f.write(json.dumps(record, ensure_ascii=False) + "\n")
payload = {
"mode": "generate_trace",
"output": str(output_path.resolve()),
"num_sessions": config.num_sessions,
"turns_per_session": config.turns_per_session,
"shared_prefix_blocks": config.shared_prefix_blocks,
"avg_new_blocks_per_turn": config.avg_new_blocks_per_turn,
"block_size": config.block_size,
"prefix_diversity": config.prefix_diversity,
"seed": config.seed,
"total_records": len(records),
}
print(json.dumps(payload, ensure_ascii=False, indent=2))
return 0


def _build_analysis_metadata_payload(
trace: str,
config_path: str,
Expand Down
5 changes: 5 additions & 0 deletions src/kvcache_upper_bound/synthetic/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
"""Parametric synthetic trace generator for KVCache analysis."""

from .generator import SyntheticTraceConfig, generate_synthetic_trace

__all__ = ["SyntheticTraceConfig", "generate_synthetic_trace"]
119 changes: 119 additions & 0 deletions src/kvcache_upper_bound/synthetic/generator.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,119 @@
"""Parametric synthetic trace generator.

Generates JSONL-compatible trace records that model multi-session,
multi-turn conversations with configurable prefix sharing.
"""

from __future__ import annotations

import hashlib
import math
import random
from dataclasses import dataclass, field
from typing import List


@dataclass
class SyntheticTraceConfig:
"""Configuration for synthetic trace generation."""

num_sessions: int
turns_per_session: int
shared_prefix_blocks: int
avg_new_blocks_per_turn: int
block_size: int = 16
prefix_diversity: float = 0.3
session_interleave: bool = True
seed: int | None = None


def _block_id(label: str, idx: int) -> str:
"""Generate a deterministic block ID using SHA-256."""
return hashlib.sha256(f"{label}-{idx}".encode()).hexdigest()[:16]


def generate_synthetic_trace(config: SyntheticTraceConfig) -> List[dict]:
"""Generate a synthetic trace from the given configuration.

Returns a list of JSONL-compatible records, each representing
a single request in a multi-turn conversation.
"""
rng = random.Random(config.seed)

# Determine prefix groups: groups = max(1, floor(diversity * num_sessions))
num_groups = max(1, math.floor(config.prefix_diversity * config.num_sessions))

# Assign sessions to prefix groups
session_groups: List[int] = []
for i in range(config.num_sessions):
session_groups.append(i % num_groups)

# Pre-generate shared prefix block IDs for each group
group_prefixes: List[List[str]] = []
for g in range(num_groups):
prefix_ids = [
_block_id(f"prefix-{g}", idx)
for idx in range(config.shared_prefix_blocks)
]
group_prefixes.append(prefix_ids)

# Build request schedule
# If session_interleave is True, interleave turns from different sessions
# Otherwise, complete each session sequentially
if config.session_interleave:
# Round-robin across sessions for each turn
schedule: List[tuple] = [] # (session_idx, turn)
for turn in range(config.turns_per_session):
session_order = list(range(config.num_sessions))
rng.shuffle(session_order)
for sid in session_order:
schedule.append((sid, turn))
else:
schedule = []
for sid in range(config.num_sessions):
for turn in range(config.turns_per_session):
schedule.append((sid, turn))

# Track per-session accumulated private blocks
session_private_blocks: List[List[str]] = [[] for _ in range(config.num_sessions)]

records: List[dict] = []
timestamp = 1000 # Start timestamp in ms

for request_idx, (session_id, turn) in enumerate(schedule):
group_id = session_groups[session_id]
prefix_ids = group_prefixes[group_id]

# Generate new unique blocks for this turn
new_blocks = [
_block_id(f"session-{session_id}-turn-{turn}", idx)
for idx in range(config.avg_new_blocks_per_turn)
]

# The full hash_ids for this request:
# prefix blocks + accumulated private blocks + new blocks
accumulated_private = list(session_private_blocks[session_id])
hash_ids = prefix_ids + accumulated_private + new_blocks

# Update accumulated private blocks for next turn
session_private_blocks[session_id].extend(new_blocks)

# Compute token lengths
input_length = len(hash_ids) * config.block_size
output_length = config.block_size # Minimal output per turn

record = {
"request_id": f"req-{request_idx:06d}",
"chat_id": f"session-{session_id}",
"parent_chat_id": f"session-{session_id}" if turn > 0 else None,
"turn": turn + 1,
"type": "text",
"timestamp": timestamp,
"input_length": input_length,
"output_length": output_length,
"hash_ids": hash_ids,
}
records.append(record)
timestamp += rng.randint(10, 100)

return records
158 changes: 158 additions & 0 deletions tests/test_cli_generate_trace.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,158 @@
"""Tests for the generate-trace CLI command."""

from __future__ import annotations

import json
import tempfile
import unittest
from pathlib import Path

from tests import _bootstrap # noqa: F401

from kvcache_upper_bound.cli.main import main


class CLIGenerateTraceTest(unittest.TestCase):
"""Tests for the generate-trace CLI subcommand."""

def test_generates_jsonl_output(self) -> None:
with tempfile.TemporaryDirectory() as tmpdir:
output_path = Path(tmpdir) / "trace.jsonl"
import sys
original_argv = sys.argv
try:
sys.argv = [
"kvcache",
"generate-trace",
"--sessions", "3",
"--turns", "2",
"--shared-prefix-blocks", "4",
"--new-blocks-per-turn", "2",
"--seed", "42",
"--output", str(output_path),
]
result = main()
finally:
sys.argv = original_argv

self.assertEqual(result, 0)
self.assertTrue(output_path.exists())

# Verify JSONL content
lines = output_path.read_text().strip().split("\n")
self.assertEqual(len(lines), 6) # 3 sessions * 2 turns

for line in lines:
record = json.loads(line)
self.assertIn("request_id", record)
self.assertIn("hash_ids", record)
self.assertIn("chat_id", record)
self.assertIn("turn", record)
self.assertIn("input_length", record)

def test_correct_record_count(self) -> None:
with tempfile.TemporaryDirectory() as tmpdir:
output_path = Path(tmpdir) / "trace.jsonl"
import sys
original_argv = sys.argv
try:
sys.argv = [
"kvcache",
"generate-trace",
"--sessions", "5",
"--turns", "4",
"--shared-prefix-blocks", "2",
"--new-blocks-per-turn", "3",
"--output", str(output_path),
]
result = main()
finally:
sys.argv = original_argv

self.assertEqual(result, 0)
lines = output_path.read_text().strip().split("\n")
self.assertEqual(len(lines), 20) # 5 * 4

def test_prefix_diversity_parameter(self) -> None:
with tempfile.TemporaryDirectory() as tmpdir:
output_path = Path(tmpdir) / "trace.jsonl"
import sys
original_argv = sys.argv
try:
sys.argv = [
"kvcache",
"generate-trace",
"--sessions", "4",
"--turns", "1",
"--shared-prefix-blocks", "3",
"--new-blocks-per-turn", "1",
"--prefix-diversity", "0.0",
"--seed", "10",
"--output", str(output_path),
]
result = main()
finally:
sys.argv = original_argv

self.assertEqual(result, 0)
lines = output_path.read_text().strip().split("\n")
records = [json.loads(line) for line in lines]

# All sessions should share the same prefix
prefixes = set()
for r in records:
prefixes.add(tuple(r["hash_ids"][:3]))
self.assertEqual(len(prefixes), 1)

def test_block_size_parameter(self) -> None:
with tempfile.TemporaryDirectory() as tmpdir:
output_path = Path(tmpdir) / "trace.jsonl"
import sys
original_argv = sys.argv
try:
sys.argv = [
"kvcache",
"generate-trace",
"--sessions", "2",
"--turns", "1",
"--shared-prefix-blocks", "2",
"--new-blocks-per-turn", "1",
"--block-size", "32",
"--seed", "42",
"--output", str(output_path),
]
result = main()
finally:
sys.argv = original_argv

self.assertEqual(result, 0)
lines = output_path.read_text().strip().split("\n")
record = json.loads(lines[0])
# 2 prefix + 1 new = 3 blocks * 32 = 96
self.assertEqual(record["input_length"], 3 * 32)

def test_creates_parent_directories(self) -> None:
with tempfile.TemporaryDirectory() as tmpdir:
output_path = Path(tmpdir) / "subdir" / "nested" / "trace.jsonl"
import sys
original_argv = sys.argv
try:
sys.argv = [
"kvcache",
"generate-trace",
"--sessions", "1",
"--turns", "1",
"--shared-prefix-blocks", "1",
"--new-blocks-per-turn", "1",
"--output", str(output_path),
]
result = main()
finally:
sys.argv = original_argv

self.assertEqual(result, 0)
self.assertTrue(output_path.exists())


if __name__ == "__main__":
unittest.main()
Loading