Add decode batch-grid stability experiment
This commit is contained in:
63
runs/frontier-decode-batch-grid-v0/materialize_decode_batch.py
Executable file
63
runs/frontier-decode-batch-grid-v0/materialize_decode_batch.py
Executable file
@@ -0,0 +1,63 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Create simultaneous, prefix-disjoint requests for a fixed decode batch."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--model", type=Path, required=True)
|
||||
parser.add_argument("--batch", type=int, required=True)
|
||||
parser.add_argument("--input-tokens", type=int, default=2048)
|
||||
parser.add_argument("--output-tokens", type=int, default=128)
|
||||
parser.add_argument("--output", type=Path, required=True)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
if min(args.batch, args.input_tokens, args.output_tokens) <= 0:
|
||||
raise ValueError("batch and token counts must be positive")
|
||||
if args.input_tokens + args.output_tokens > 40960:
|
||||
raise ValueError("request exceeds the server max model length")
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.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) < args.batch + 1:
|
||||
raise ValueError("tokenizer has too few non-special token IDs")
|
||||
|
||||
base = candidates[0]
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
with args.output.open("w") as output:
|
||||
for index in range(args.batch):
|
||||
row = {
|
||||
"source_index": index,
|
||||
"arrived_at": 0.0,
|
||||
"input_length": args.input_tokens,
|
||||
"output_length": args.output_tokens,
|
||||
"session_id": index,
|
||||
"runtime_block_ids": [],
|
||||
"body": {
|
||||
"prompt": [
|
||||
candidates[index + 1],
|
||||
*([base] * (args.input_tokens - 1)),
|
||||
],
|
||||
"min_tokens": args.output_tokens,
|
||||
"max_tokens": args.output_tokens,
|
||||
"ignore_eos": True,
|
||||
},
|
||||
}
|
||||
output.write(json.dumps(row, separators=(",", ":")) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user