"""Offline random GPT-2 fixture. Measures request latency, not language quality."""
import argparse
import json
import statistics
import time
import torch
import transformers
from transformers import GPT2Config, GPT2LMHeadModel


def timed_call(generate, synchronize, clock=time.perf_counter):
    synchronize()
    start = clock()
    output = generate()
    synchronize()
    elapsed = clock() - start
    if elapsed <= 0:
        raise ValueError("nonpositive elapsed time")
    return output, elapsed


def fixture(device):
    torch.manual_seed(7)
    config = GPT2Config(n_layer=1, n_head=2, n_embd=32, vocab_size=64,
                        n_positions=64, bos_token_id=1, eos_token_id=None,
                        pad_token_id=0)
    model = GPT2LMHeadModel(config).eval().to(device)
    inputs = torch.tensor([[1, 3, 5, 7]], device=device)
    mask = torch.ones_like(inputs)
    return model, inputs, mask


def run(device_name="cpu", repeats=5, new_tokens=8):
    if device_name not in ("cpu", "cuda:0"):
        raise ValueError("select cpu or cuda:0 explicitly")
    if type(repeats) is not int or repeats < 1:
        raise ValueError("repeats must be positive")
    if type(new_tokens) is not int or not 1 <= new_tokens <= 60:
        raise ValueError("new_tokens must be 1..60")
    if device_name == "cuda:0" and not torch.cuda.is_available():
        raise ValueError("CUDA unavailable; select cpu for the fixture")
    device = torch.device(device_name)
    model, inputs, mask = fixture(device)

    def synchronize():
        if device.type == "cuda":
            torch.cuda.synchronize(device)

    def generate():
        with torch.inference_mode():
            return model.generate(inputs, attention_mask=mask,
                                  max_new_tokens=new_tokens,
                                  do_sample=False, use_cache=True)

    for _ in range(2):
        generate()
    synchronize()
    samples = []
    for _ in range(repeats):
        output, elapsed = timed_call(generate, synchronize)
        actual_tokens = output.shape[1] - inputs.shape[1]
        samples.append({"new_tokens": actual_tokens, "seconds": elapsed,
                        "tokens_per_second": actual_tokens / elapsed})
    return {"fixture": "random GPT-2; not a trained LLM", "device": str(device),
            "dtype": str(next(model.parameters()).dtype),
            "parameters": sum(p.numel() for p in model.parameters()),
            "torch": torch.__version__, "transformers": transformers.__version__,
            "prompt_tokens": inputs.shape[1], "warmup_requests": 2,
            "median_request_seconds": statistics.median(s["seconds"] for s in samples),
            "samples": samples}


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--device", choices=["cpu", "cuda:0"], default="cpu")
    args = parser.parse_args()
    torch.set_num_threads(1)
    print(json.dumps(run(args.device), indent=2))
