Models / gpt2

GPT-2 (124M)

OpenAI's 124M-parameter GPT-2: learned positions, pre-norm blocks, a fused QKV projection, and a head tied to the token embedding.

124.4M parametersgpt2MITtext-generationdecoder-only
models/gpt27 files · 41.9 kB in the registry · GitHub
README.md1.9 kB
md
# GPT-2 (124M)

OpenAI's smallest GPT-2: 12 pre-norm blocks of width 768 with 12 heads,
learned absolute positions up to 1024, a fused QKV projection split by
slicing, the tanh GELU, and an output head tied to the token embedding. The
Linnet source keeps GPT-2's `[in, out]` projection layout (`Conv1D` in the
reference code), so the published checkpoint loads as it is.

## Loading

```python
from linnet import nest

model = nest.load("gpt2", backend="torch", numerics="equivalent")
logits = model(tokens)                      # Tensor[B, S; i32] -> Tensor[B, S, 50257; f32]
```

`bindings.json` maps Linnet parameter paths (`blocks.0.attn.qkv.weight`) to
the checkpoint's names (`h.0.attn.c_attn.weight`). Sequences longer than
1024 positions are rejected at compile time (`where S <= MaxPositions`).

## Entries

| Entry | |
| --- | --- |
| `forward<B, S>(tokens)` | logits for a whole sequence |
| `decode(token, pos)` | one token through the KV caches (`Batch = 1`, `MaxSeq = 1024`) |
| `prefill<S>(tokens, pos)` | a whole prompt through the KV caches in one pass; logits after its last token, `decode` continues at `pos + S` |
| `prefill_slots<M, S>(tokens, slots, lengths)` | `M` requests' prompts into rows `slots` of the caches in one pass (each padded to `S`, its first `lengths[m]` tokens real), as they join a batch being served |
| `decode_rows(tokens, positions)` | one token for every row of the caches, each at its own position: the step `linnet.serve` takes for continuous batching |

## Provenance

- Weights: [openai-community/gpt2](https://huggingface.co/openai-community/gpt2), MIT.
- Code: [openai/gpt-2](https://github.com/openai/gpt-2).
- Announcement: [Better language models and their implications](https://openai.com/research/better-language-models).

The Linnet forward pass matches `transformers` logits to 2e-2 absolute
(see `python/linnet/tests/torch/test_hf_checkpoints.py` in the Linnet
repository).

Open on GitHub

bench.json18.4 kB
json
{
  "date": "2026-09-29T14:54:02+00:00",
  "environment": {
    "python": "3.12.3",
    "machine": "x86_64",
    "torch": "2.14.0",
    "jax": "0.11.2",
    "transformers": "5.17.0",
    "diffusers": "0.40.0",
    "onnxruntime": "1.30.0",
    "linnet": "0.1.0",
    "gpu": "NVIDIA H100 80GB HBM3",
    "gpus": "1",
    "driver": "580.126.09"
  },
  "workload": {
    "device": "cuda",
    "prompt_tokens": 512,
    "new_tokens": 128,
    "batch": 1,
    "warmup": 3,
    "iters": 10,
    "seed": 0,
    "offload_gib": 8.0,
    "serve_requests": 256,
    "serve_concurrency": 64,
    "serve_prompt_min": 128,
    "serve_prompt_max": 512,
    "serve_new": 128
  },
  "methods": [
    {
      "method": "transformers (eager)",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 6.703074090182781,
        "decode_tok_s": 175.28543516774835,
        "load_s": 2.007332729175687,
        "peak_vram_mib": 1096.9
      },
      "max_abs_diff": 0.0,
      "notes": "",
      "error": null,
      "key": "transformers-eager"
    },
    {
      "method": "transformers (torch.compile, static cache)",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 4.066447727382183,
        "decode_tok_s": 419.1342739561902,
        "load_s": 1.467197461053729,
        "peak_vram_mib": 1280.0
      },
      "max_abs_diff": 1.25,
      "notes": "",
      "error": null,
      "key": "transformers-compile"
    },
    {
      "method": "vLLM",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 5.9449076652526855,
        "decode_tok_s": 696.2508439753117,
        "load_s": 46.960085187107325,
        "first_token": 2213.0,
        "peak_vram_mib": 69720.9
      },
      "max_abs_diff": null,
      "notes": "reserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "vllm"
    },
    {
      "method": "Linnet torch (generated source)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 3.094634972512722,
        "decode_tok_s": 290.2121965718771,
        "load_s": 1.0232124645262957,
        "peak_vram_mib": 1020.0
      },
      "max_abs_diff": 0.75,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-torch"
    },
    {
      "method": "Linnet torch (CUDA graphs)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 0.7712710648775101,
        "decode_tok_s": 1896.8173529286903,
        "load_s": 0.4614080637693405,
        "peak_vram_mib": 1144.9
      },
      "max_abs_diff": 1.25,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-cudagraphs"
    },
    {
      "method": "Linnet torch (torch.compile, inductor)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 1.5868023037910461,
        "decode_tok_s": 552.821169559486,
        "load_s": 0.4720837324857712,
        "peak_vram_mib": 1022.9
      },
      "max_abs_diff": 1.25,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-inductor"
    },
    {
      "method": "Linnet JAX (XLA, StableHLO)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 1.2286584824323654,
        "decode_tok_s": 1847.56881982313,
        "load_s": 0.5153531208634377,
        "peak_vram_mib": 514.0,
        "driver_vram_mib": 1140.9
      },
      "max_abs_diff": 0.75,
      "notes": "KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "linnet-jax"
    },
    {
      "method": "Linnet JAX (XLA, generated source)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 1.313304528594017,
        "decode_tok_s": 1714.8556513928784,
        "load_s": 0.47653499245643616,
        "peak_vram_mib": 514.0,
        "driver_vram_mib": 1140.9
      },
      "max_abs_diff": 1.25,
      "notes": "KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "linnet-jax-source"
    },
    {
      "method": "KerasHub (JAX)",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 6.838363595306873,
        "decode_tok_s": 1452.4523362313137,
        "load_s": 6.1567362900823355,
        "peak_vram_mib": 660.4,
        "driver_vram_mib": 1654.0
      },
      "max_abs_diff": 0.25,
      "notes": "peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "keras-hub"
    },
    {
      "method": "Linnet ONNX f32 -> ONNX Runtime (CUDA)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 4.214955493807793,
        "decode_tok_s": 427.3610844462876,
        "load_s": 0.2946742195636034,
        "peak_vram_mib": 2606.0
      },
      "max_abs_diff": 1.3825111389160156,
      "notes": "f32; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-f32"
    },
    {
      "method": "Linnet ONNX f32 -> ONNX Runtime (TensorRT)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 2.690473571419716,
        "decode_tok_s": 541.8647804834536,
        "load_s": 0.2855817433446646,
        "peak_vram_mib": 4098.0
      },
      "max_abs_diff": 1.3637847900390625,
      "notes": "f32; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-f32-trt"
    },
    {
      "method": "Linnet ONNX f16 -> ONNX Runtime (CUDA)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 3.41070257127285,
        "decode_tok_s": 468.2145868087278,
        "load_s": 0.2914519440382719,
        "peak_vram_mib": 1714.0
      },
      "max_abs_diff": 1.34375,
      "notes": "f16; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-f16"
    },
    {
      "method": "Linnet ONNX f16 -> ONNX Runtime (TensorRT)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 2.041681669652462,
        "decode_tok_s": 557.2825754893278,
        "load_s": 0.2913510575890541,
        "peak_vram_mib": 3334.0
      },
      "max_abs_diff": 1.78125,
      "notes": "f16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-f16-trt"
    },
    {
      "method": "Linnet ONNX bf16 -> ONNX Runtime (CUDA)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 3.7594325840473175,
        "decode_tok_s": 421.0260150417744,
        "load_s": 0.2960293795913458,
        "peak_vram_mib": 2216.0
      },
      "max_abs_diff": 0.5,
      "notes": "bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-bf16"
    },
    {
      "method": "Linnet ONNX bf16 -> ONNX Runtime (TensorRT)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 2.199768088757992,
        "decode_tok_s": 510.0053596687303,
        "load_s": 0.28357444144785404,
        "peak_vram_mib": 3166.0
      },
      "max_abs_diff": 0.25,
      "notes": "bf16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-bf16-trt"
    },
    {
      "method": "vLLM (offline, continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 25857.812874887193,
        "serve_s": 1.2672378811985254,
        "load_s": 22.20858423411846,
        "peak_vram_mib": 69406.8
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85",
      "error": null,
      "key": "serve-vllm"
    },
    {
      "method": "transformers (generate_batch, continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 2950.7114793370224,
        "serve_s": 11.105118284001946,
        "load_s": 1.2816545199602842,
        "peak_vram_mib": 80574.0
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; paged|sdpa attention; reserves a paged KV-cache pool up front, so its memory is a setting, not a need; no per-request timestamps",
      "error": null,
      "key": "serve-transformers"
    },
    {
      "method": "KerasHub (JAX, static batches)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 2604.558799059782,
        "serve_s": 12.581017564982176,
        "load_s": 6.146484849974513,
        "peak_vram_mib": 3191.7,
        "driver_vram_mib": 8818.0
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; every batch runs to the longest possible prompt plus the new tokens; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "serve-keras-hub"
    },
    {
      "method": "Triton Inference Server (vLLM backend)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 3856.919261513569,
        "serve_s": 8.49589990824461,
        "load_s": 43.0544276740402,
        "serve_ttft_ms": 4635.077881626785,
        "peak_vram_mib": 69970.4
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed, first tokens timed at the client",
      "error": null,
      "key": "serve-triton-vllm"
    },
    {
      "method": "Linnet torch (linnet.serve, CUDA graphs)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 38317.395044137775,
        "serve_s": 0.8551729563623667,
        "load_s": 4.3678231770172715,
        "serve_ttft_ms": 363.8608152978122,
        "peak_vram_mib": 2708.0
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions",
      "error": null,
      "key": "serve-linnet-torch"
    },
    {
      "method": "Linnet jax (linnet.serve, XLA)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 35104.33942548846,
        "serve_s": 0.9334458513185382,
        "load_s": 2.512056190520525,
        "serve_ttft_ms": 468.92378851771355,
        "peak_vram_mib": 2055.9,
        "driver_vram_mib": 4756.0
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "serve-linnet-jax"
    },
    {
      "method": "Linnet onnx (linnet.serve, ONNX Runtime, f16)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 22510.041395468208,
        "serve_s": 1.4557058969512582,
        "load_s": 0.242721576243639,
        "serve_ttft_ms": 639.9877788498998,
        "peak_vram_mib": 11512.0
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions",
      "error": null,
      "key": "serve-linnet-onnx"
    },
    {
      "method": "Triton Inference Server (Python backend, linnet.serve)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 14165.680965886502,
        "serve_s": 2.3131962437182665,
        "load_s": 19.046876681037247,
        "serve_ttft_ms": 1302.7412672527134,
        "peak_vram_mib": 3416.4
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed, first tokens timed at the client",
      "error": null,
      "key": "serve-triton-linnet"
    },
    {
      "method": "Linnet -> vLLM",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 6.37507438659668,
        "decode_tok_s": 821.3562907897797,
        "load_s": 20.473177609965205,
        "first_token": 2213.0,
        "peak_vram_mib": 69720.8,
        "driver_vram_mib": 69720.8
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (1 s), then vLLM on the exported checkpoint; reserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "linnet-vllm"
    },
    {
      "method": "Linnet -> vLLM (offline, continuous batching)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 27903.674382105426,
        "serve_s": 1.174325630068779,
        "load_s": 19.213800007477403,
        "peak_vram_mib": 69414.8,
        "driver_vram_mib": 69414.8
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (0 s), then vLLM on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85",
      "error": null,
      "key": "serve-linnet-vllm"
    },
    {
      "method": "Linnet -> llama.cpp (GGUF, f16)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 3.1551189770267327,
        "decode_tok_s": 1650.541843,
        "load_s": 8.608857814222574,
        "peak_vram_mib": 1056.8
      },
      "max_abs_diff": null,
      "notes": "linnet.gguf.export, then llama-bench with every layer on the GPU; the first token is the prompt at llama-bench's prompt rate",
      "error": null,
      "key": "linnet-llamacpp"
    },
    {
      "method": "SGLang",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 10.368083603680134,
        "decode_tok_s": 896.7476008183686,
        "load_s": 30.858465744182467,
        "first_token": 2213.0,
        "peak_vram_mib": 69351.6
      },
      "max_abs_diff": null,
      "notes": "reserves a KV-cache pool up front (mem_fraction_static 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "sglang"
    },
    {
      "method": "Linnet -> SGLang",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 9.253870695829391,
        "decode_tok_s": 922.6518512864118,
        "load_s": 24.528803979977965,
        "first_token": 2213.0,
        "peak_vram_mib": 69351.6,
        "driver_vram_mib": 69351.6
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (2 s), then SGLang on the exported checkpoint; reserves a KV-cache pool up front (mem_fraction_static 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "linnet-sglang"
    },
    {
      "method": "SGLang (offline, continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 28787.82069900377,
        "serve_s": 1.1382591389119625,
        "load_s": 31.076543428003788,
        "peak_vram_mib": 70021.6
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at mem_fraction_static 0.85",
      "error": null,
      "key": "serve-sglang"
    },
    {
      "method": "Linnet -> SGLang (offline, continuous batching)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 27129.13680828877,
        "serve_s": 1.2078526578843594,
        "load_s": 25.906673204153776,
        "peak_vram_mib": 70021.6,
        "driver_vram_mib": 70021.6
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (0 s), then SGLang (offline, continuous batching) on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at mem_fraction_static 0.85",
      "error": null,
      "key": "serve-linnet-sglang"
    },
    {
      "method": "Text Generation Inference",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 14.190178364515305,
        "decode_tok_s": 547.7135727036018,
        "load_s": 30.075674967840314,
        "peak_vram_mib": 61906.9
      },
      "max_abs_diff": null,
      "notes": "over HTTP, streamed; the prompt is the token ids decoded and tokenized again; reserves a KV-cache pool up front (cuda-memory-fraction 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "tgi"
    },
    {
      "method": "Linnet -> Text Generation Inference",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 14.092078432440758,
        "decode_tok_s": 580.3222178226044,
        "load_s": 26.069788286462426,
        "peak_vram_mib": 61906.9,
        "driver_vram_mib": 61906.9
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (2 s), then Text Generation Inference on the exported checkpoint; over HTTP, streamed; the prompt is the token ids decoded and tokenized again; reserves a KV-cache pool up front (cuda-memory-fraction 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "linnet-tgi"
    },
    {
      "method": "Text Generation Inference (continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 6189.046201277235,
        "serve_s": 5.130354333668947,
        "load_s": 32.07420591451228,
        "serve_ttft_ms": 1721.4535884559155,
        "peak_vram_mib": 62204.9
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed; stops at end-of-sequence (TGI cannot ignore it), so throughput counts the tokens produced; KV-cache pool at cuda-memory-fraction 0.85",
      "error": null,
      "key": "serve-tgi"
    },
    {
      "method": "Linnet -> Text Generation Inference (continuous batching)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 6066.486756425216,
        "serve_s": 5.253782177343965,
        "load_s": 28.069601520895958,
        "serve_ttft_ms": 1881.6367276012897,
        "peak_vram_mib": 62320.9,
        "driver_vram_mib": 62320.9
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (0 s), then Text Generation Inference (continuous batching) on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed; stops at end-of-sequence (TGI cannot ignore it), so throughput counts the tokens produced; KV-cache pool at cuda-memory-fraction 0.85",
      "error": null,
      "key": "serve-linnet-tgi"
    }
  ],
  "reference": "transformers-eager"
}

Open on GitHub

bindings.json7.0 kB
json
{
  "wte": "wte.weight",
  "wpe": "wpe.weight",
  "blocks.0.ln_1.weight": "h.0.ln_1.weight",
  "blocks.0.ln_1.bias": "h.0.ln_1.bias",
  "blocks.0.ln_2.weight": "h.0.ln_2.weight",
  "blocks.0.ln_2.bias": "h.0.ln_2.bias",
  "blocks.0.attn.qkv.weight": "h.0.attn.c_attn.weight",
  "blocks.0.attn.qkv.bias": "h.0.attn.c_attn.bias",
  "blocks.0.attn.out.weight": "h.0.attn.c_proj.weight",
  "blocks.0.attn.out.bias": "h.0.attn.c_proj.bias",
  "blocks.0.mlp.up.weight": "h.0.mlp.c_fc.weight",
  "blocks.0.mlp.up.bias": "h.0.mlp.c_fc.bias",
  "blocks.0.mlp.down.weight": "h.0.mlp.c_proj.weight",
  "blocks.0.mlp.down.bias": "h.0.mlp.c_proj.bias",
  "blocks.1.ln_1.weight": "h.1.ln_1.weight",
  "blocks.1.ln_1.bias": "h.1.ln_1.bias",
  "blocks.1.ln_2.weight": "h.1.ln_2.weight",
  "blocks.1.ln_2.bias": "h.1.ln_2.bias",
  "blocks.1.attn.qkv.weight": "h.1.attn.c_attn.weight",
  "blocks.1.attn.qkv.bias": "h.1.attn.c_attn.bias",
  "blocks.1.attn.out.weight": "h.1.attn.c_proj.weight",
  "blocks.1.attn.out.bias": "h.1.attn.c_proj.bias",
  "blocks.1.mlp.up.weight": "h.1.mlp.c_fc.weight",
  "blocks.1.mlp.up.bias": "h.1.mlp.c_fc.bias",
  "blocks.1.mlp.down.weight": "h.1.mlp.c_proj.weight",
  "blocks.1.mlp.down.bias": "h.1.mlp.c_proj.bias",
  "blocks.2.ln_1.weight": "h.2.ln_1.weight",
  "blocks.2.ln_1.bias": "h.2.ln_1.bias",
  "blocks.2.ln_2.weight": "h.2.ln_2.weight",
  "blocks.2.ln_2.bias": "h.2.ln_2.bias",
  "blocks.2.attn.qkv.weight": "h.2.attn.c_attn.weight",
  "blocks.2.attn.qkv.bias": "h.2.attn.c_attn.bias",
  "blocks.2.attn.out.weight": "h.2.attn.c_proj.weight",
  "blocks.2.attn.out.bias": "h.2.attn.c_proj.bias",
  "blocks.2.mlp.up.weight": "h.2.mlp.c_fc.weight",
  "blocks.2.mlp.up.bias": "h.2.mlp.c_fc.bias",
  "blocks.2.mlp.down.weight": "h.2.mlp.c_proj.weight",
  "blocks.2.mlp.down.bias": "h.2.mlp.c_proj.bias",
  "blocks.3.ln_1.weight": "h.3.ln_1.weight",
  "blocks.3.ln_1.bias": "h.3.ln_1.bias",
  "blocks.3.ln_2.weight": "h.3.ln_2.weight",
  "blocks.3.ln_2.bias": "h.3.ln_2.bias",
  "blocks.3.attn.qkv.weight": "h.3.attn.c_attn.weight",
  "blocks.3.attn.qkv.bias": "h.3.attn.c_attn.bias",
  "blocks.3.attn.out.weight": "h.3.attn.c_proj.weight",
  "blocks.3.attn.out.bias": "h.3.attn.c_proj.bias",
  "blocks.3.mlp.up.weight": "h.3.mlp.c_fc.weight",
  "blocks.3.mlp.up.bias": "h.3.mlp.c_fc.bias",
  "blocks.3.mlp.down.weight": "h.3.mlp.c_proj.weight",
  "blocks.3.mlp.down.bias": "h.3.mlp.c_proj.bias",
  "blocks.4.ln_1.weight": "h.4.ln_1.weight",
  "blocks.4.ln_1.bias": "h.4.ln_1.bias",
  "blocks.4.ln_2.weight": "h.4.ln_2.weight",
  "blocks.4.ln_2.bias": "h.4.ln_2.bias",
  "blocks.4.attn.qkv.weight": "h.4.attn.c_attn.weight",
  "blocks.4.attn.qkv.bias": "h.4.attn.c_attn.bias",
  "blocks.4.attn.out.weight": "h.4.attn.c_proj.weight",
  "blocks.4.attn.out.bias": "h.4.attn.c_proj.bias",
  "blocks.4.mlp.up.weight": "h.4.mlp.c_fc.weight",
  "blocks.4.mlp.up.bias": "h.4.mlp.c_fc.bias",
  "blocks.4.mlp.down.weight": "h.4.mlp.c_proj.weight",
  "blocks.4.mlp.down.bias": "h.4.mlp.c_proj.bias",
  "blocks.5.ln_1.weight": "h.5.ln_1.weight",
  "blocks.5.ln_1.bias": "h.5.ln_1.bias",
  "blocks.5.ln_2.weight": "h.5.ln_2.weight",
  "blocks.5.ln_2.bias": "h.5.ln_2.bias",
  "blocks.5.attn.qkv.weight": "h.5.attn.c_attn.weight",
  "blocks.5.attn.qkv.bias": "h.5.attn.c_attn.bias",
  "blocks.5.attn.out.weight": "h.5.attn.c_proj.weight",
  "blocks.5.attn.out.bias": "h.5.attn.c_proj.bias",
  "blocks.5.mlp.up.weight": "h.5.mlp.c_fc.weight",
  "blocks.5.mlp.up.bias": "h.5.mlp.c_fc.bias",
  "blocks.5.mlp.down.weight": "h.5.mlp.c_proj.weight",
  "blocks.5.mlp.down.bias": "h.5.mlp.c_proj.bias",
  "blocks.6.ln_1.weight": "h.6.ln_1.weight",
  "blocks.6.ln_1.bias": "h.6.ln_1.bias",
  "blocks.6.ln_2.weight": "h.6.ln_2.weight",
  "blocks.6.ln_2.bias": "h.6.ln_2.bias",
  "blocks.6.attn.qkv.weight": "h.6.attn.c_attn.weight",
  "blocks.6.attn.qkv.bias": "h.6.attn.c_attn.bias",
  "blocks.6.attn.out.weight": "h.6.attn.c_proj.weight",
  "blocks.6.attn.out.bias": "h.6.attn.c_proj.bias",
  "blocks.6.mlp.up.weight": "h.6.mlp.c_fc.weight",
  "blocks.6.mlp.up.bias": "h.6.mlp.c_fc.bias",
  "blocks.6.mlp.down.weight": "h.6.mlp.c_proj.weight",
  "blocks.6.mlp.down.bias": "h.6.mlp.c_proj.bias",
  "blocks.7.ln_1.weight": "h.7.ln_1.weight",
  "blocks.7.ln_1.bias": "h.7.ln_1.bias",
  "blocks.7.ln_2.weight": "h.7.ln_2.weight",
  "blocks.7.ln_2.bias": "h.7.ln_2.bias",
  "blocks.7.attn.qkv.weight": "h.7.attn.c_attn.weight",
  "blocks.7.attn.qkv.bias": "h.7.attn.c_attn.bias",
  "blocks.7.attn.out.weight": "h.7.attn.c_proj.weight",
  "blocks.7.attn.out.bias": "h.7.attn.c_proj.bias",
  "blocks.7.mlp.up.weight": "h.7.mlp.c_fc.weight",
  "blocks.7.mlp.up.bias": "h.7.mlp.c_fc.bias",
  "blocks.7.mlp.down.weight": "h.7.mlp.c_proj.weight",
  "blocks.7.mlp.down.bias": "h.7.mlp.c_proj.bias",
  "blocks.8.ln_1.weight": "h.8.ln_1.weight",
  "blocks.8.ln_1.bias": "h.8.ln_1.bias",
  "blocks.8.ln_2.weight": "h.8.ln_2.weight",
  "blocks.8.ln_2.bias": "h.8.ln_2.bias",
  "blocks.8.attn.qkv.weight": "h.8.attn.c_attn.weight",
  "blocks.8.attn.qkv.bias": "h.8.attn.c_attn.bias",
  "blocks.8.attn.out.weight": "h.8.attn.c_proj.weight",
  "blocks.8.attn.out.bias": "h.8.attn.c_proj.bias",
  "blocks.8.mlp.up.weight": "h.8.mlp.c_fc.weight",
  "blocks.8.mlp.up.bias": "h.8.mlp.c_fc.bias",
  "blocks.8.mlp.down.weight": "h.8.mlp.c_proj.weight",
  "blocks.8.mlp.down.bias": "h.8.mlp.c_proj.bias",
  "blocks.9.ln_1.weight": "h.9.ln_1.weight",
  "blocks.9.ln_1.bias": "h.9.ln_1.bias",
  "blocks.9.ln_2.weight": "h.9.ln_2.weight",
  "blocks.9.ln_2.bias": "h.9.ln_2.bias",
  "blocks.9.attn.qkv.weight": "h.9.attn.c_attn.weight",
  "blocks.9.attn.qkv.bias": "h.9.attn.c_attn.bias",
  "blocks.9.attn.out.weight": "h.9.attn.c_proj.weight",
  "blocks.9.attn.out.bias": "h.9.attn.c_proj.bias",
  "blocks.9.mlp.up.weight": "h.9.mlp.c_fc.weight",
  "blocks.9.mlp.up.bias": "h.9.mlp.c_fc.bias",
  "blocks.9.mlp.down.weight": "h.9.mlp.c_proj.weight",
  "blocks.9.mlp.down.bias": "h.9.mlp.c_proj.bias",
  "blocks.10.ln_1.weight": "h.10.ln_1.weight",
  "blocks.10.ln_1.bias": "h.10.ln_1.bias",
  "blocks.10.ln_2.weight": "h.10.ln_2.weight",
  "blocks.10.ln_2.bias": "h.10.ln_2.bias",
  "blocks.10.attn.qkv.weight": "h.10.attn.c_attn.weight",
  "blocks.10.attn.qkv.bias": "h.10.attn.c_attn.bias",
  "blocks.10.attn.out.weight": "h.10.attn.c_proj.weight",
  "blocks.10.attn.out.bias": "h.10.attn.c_proj.bias",
  "blocks.10.mlp.up.weight": "h.10.mlp.c_fc.weight",
  "blocks.10.mlp.up.bias": "h.10.mlp.c_fc.bias",
  "blocks.10.mlp.down.weight": "h.10.mlp.c_proj.weight",
  "blocks.10.mlp.down.bias": "h.10.mlp.c_proj.bias",
  "blocks.11.ln_1.weight": "h.11.ln_1.weight",
  "blocks.11.ln_1.bias": "h.11.ln_1.bias",
  "blocks.11.ln_2.weight": "h.11.ln_2.weight",
  "blocks.11.ln_2.bias": "h.11.ln_2.bias",
  "blocks.11.attn.qkv.weight": "h.11.attn.c_attn.weight",
  "blocks.11.attn.qkv.bias": "h.11.attn.c_attn.bias",
  "blocks.11.attn.out.weight": "h.11.attn.c_proj.weight",
  "blocks.11.attn.out.bias": "h.11.attn.c_proj.bias",
  "blocks.11.mlp.up.weight": "h.11.mlp.c_fc.weight",
  "blocks.11.mlp.up.bias": "h.11.mlp.c_fc.bias",
  "blocks.11.mlp.down.weight": "h.11.mlp.c_proj.weight",
  "blocks.11.mlp.down.bias": "h.11.mlp.c_proj.bias"
}

Open on GitHub

gpt2.linnet12.5 kB
linnet
// GPT-2: learned absolute positions, pre-norm blocks with biased LayerNorm,
// a fused QKV projection split by slicing, the tanh GELU, and an output head
// tied to the token embedding through index notation.
//
// `forward` attends over a whole sequence. `decode` takes one position and
// keeps the keys and values seen so far in two `state` members, sized by
// the block's `Batch` and `MaxSeq`: the block assigns them, the runtime
// keeps them between calls, and graph exports thread them in and out.
// `prefill` is `decode` for a whole prompt: every position in one pass.
module examples.gpt2

use std.nn.activations::{gelu}
use std.nn.attention::{attention, attention_rows, causal_mask}
use std.nn.cache::{write_at, write_rows, write_slots, write_span}
use std.nn.embedding::{embedding}
use std.nn.linear::{linear}
use std.nn.norm::{layer_norm}

// LayerNorm with the bias GPT-2 uses; `layer_norm` takes it as an optional.
pub block LayerNorm<H: Dim, T: Float> {
    param weight: Tensor[H; T]
    param bias: Tensor[H; T]

    pub fn forward<*S: Shape>(x: Tensor[*S, H; T]) -> Tensor[*S, H; T] {
        return layer_norm(x, weight, some(bias))
    }
}

// GPT-2 stores its projections as [In, Out] (`Conv1D` in the reference
// code). Keeping that layout lets the published checkpoints load as they
// are; the transpose is a view in every backend.
pub block Conv1D<In: Dim, Out: Dim, T: Float> {
    param weight: Tensor[In, Out; T]
    param bias: Tensor[Out; T]

    pub fn forward<*S: Shape>(x: Tensor[*S, In; T]) -> Tensor[*S, Out; T] {
        return linear(x, permute(weight, [1, 0]), some(bias))
    }
}

pub block Attention<H: Dim, Heads: Dim, Batch: Dim, MaxSeq: Dim, T: Float>
where
    Heads > 0,
    H % Heads == 0,
    MaxSeq > 0
{
    // One projection produces queries, keys, and values side by side.
    sub qkv: Conv1D<H, 3 * H, T>
    sub out: Conv1D<H, H, T>

    state cache_k: Tensor[Batch, Heads, MaxSeq, H / Heads; T]
    state cache_v: Tensor[Batch, Heads, MaxSeq, H / Heads; T]

    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
        let projected = qkv.forward(x)
        let q = heads<B, S, Heads, H / Heads, T>(projected[:, :, 0:H])
        let k = heads<B, S, Heads, H / Heads, T>(projected[:, :, H:2 * H])
        let v = heads<B, S, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
        let mixed = attention(q, k, v, rsqrt(cast<f32>(H / Heads)), some(causal_mask<S, S>()))
        return out.forward(reshape(permute(mixed, [0, 2, 1, 3]), [B, S, H]))
    }

    // One new position `pos`: its key and value are written into the caches
    // at that position (a masked write, no scatter), and the query attends
    // over every cached position up to and including it.
    pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32) -> Tensor[Batch, 1, H; T] {
        let projected = qkv.forward(x)
        let q = heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, 0:H])
        let k = heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, H:2 * H])
        let v = heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
        cache_k = write_at(cache_k, k, pos)
        cache_v = write_at(cache_v, v, pos)
        let positions = iota<i32>(MaxSeq)
        let seen[s] = positions[s] <= pos
        let mixed = attention(
            q,
            cache_k,
            cache_v,
            rsqrt(cast<f32>(H / Heads)),
            some(reshape(seen, [1, MaxSeq])),
        )
        return out.forward(reshape(permute(mixed, [0, 2, 1, 3]), [Batch, 1, H]))
    }

    // `S` new positions starting at `pos`, as a prompt arrives: their keys
    // and values go into the caches as one span, and each query attends over
    // every cached position up to and including its own. A prompt costs one
    // pass instead of one per token, which is most of the time to the first
    // generated token.
    pub fn prefill<S: Dim>(x: Tensor[Batch, S, H; T], pos: i32) -> Tensor[Batch, S, H; T]
    where S > 0 {
        let projected = qkv.forward(x)
        let q = heads<Batch, S, Heads, H / Heads, T>(projected[:, :, 0:H])
        let k = heads<Batch, S, Heads, H / Heads, T>(projected[:, :, H:2 * H])
        let v = heads<Batch, S, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
        cache_k = write_span(cache_k, k, pos)
        cache_v = write_span(cache_v, v, pos)
        let positions = iota<i32>(MaxSeq)
        let queries = iota<i32>(S)
        let seen[i, s] = positions[s] <= pos + queries[i]
        let mixed = attention(q, cache_k, cache_v, rsqrt(cast<f32>(H / Heads)), some(seen))
        return out.forward(reshape(permute(mixed, [0, 2, 1, 3]), [Batch, S, H]))
    }

    // One new position per sequence, each at its own `positions[b]`: the
    // step a server takes for every request it is decoding together.
    pub fn decode_rows(
        x: Tensor[Batch, 1, H; T],
        positions: Tensor[Batch; i32],
    ) -> Tensor[Batch, 1, H; T] {
        let projected = qkv.forward(x)
        let q = heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, 0:H])
        let k = heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, H:2 * H])
        let v = heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
        cache_k = write_rows(cache_k, k, positions)
        cache_v = write_rows(cache_v, v, positions)
        let slots = iota<i32>(MaxSeq)
        let seen[b, s] = slots[s] <= positions[b]
        let mixed = attention_rows(
            q,
            cache_k,
            cache_v,
            rsqrt(cast<f32>(H / Heads)),
            reshape(seen, [Batch, 1, MaxSeq]),
        )
        return out.forward(reshape(permute(mixed, [0, 2, 1, 3]), [Batch, 1, H]))
    }

    // `M` requests' prompts of `S` tokens, from their first positions, into
    // rows `slots` of the caches; each prompt's queries attend causally among
    // themselves.
    pub fn prefill_slots<M: Dim, S: Dim>(
        x: Tensor[M, S, H; T],
        slots: Tensor[M; i32],
    ) -> Tensor[M, S, H; T]
    where
        M > 0,
        S > 0
    {
        let projected = qkv.forward(x)
        let q = heads<M, S, Heads, H / Heads, T>(projected[:, :, 0:H])
        let k = heads<M, S, Heads, H / Heads, T>(projected[:, :, H:2 * H])
        let v = heads<M, S, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
        cache_k = write_slots(cache_k, k, slots, 0)
        cache_v = write_slots(cache_v, v, slots, 0)
        let mixed = attention(q, k, v, rsqrt(cast<f32>(H / Heads)), some(causal_mask<S, S>()))
        return out.forward(reshape(permute(mixed, [0, 2, 1, 3]), [M, S, H]))
    }
}

fn heads<B: Dim, S: Dim, N: Dim, D: Dim, T: Float>(
    x: Tensor[B, S, N * D; T],
) -> Tensor[B, N, S, D; T] {
    return permute(reshape(x, [B, S, N, D]), [0, 2, 1, 3])
}

pub block Mlp<H: Dim, T: Float> {
    sub up: Conv1D<H, 4 * H, T>
    sub down: Conv1D<4 * H, H, T>

    pub fn forward<*S: Shape>(x: Tensor[*S, H; T]) -> Tensor[*S, H; T] {
        return down.forward(gelu(up.forward(x)))
    }
}

pub block Block<H: Dim, Heads: Dim, Batch: Dim, MaxSeq: Dim, T: Float>
where
    Heads > 0,
    H % Heads == 0,
    MaxSeq > 0
{
    sub ln_1: LayerNorm<H, T>
    sub attn: Attention<H, Heads, Batch, MaxSeq, T>
    sub ln_2: LayerNorm<H, T>
    sub mlp: Mlp<H, T>

    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
        let attended = x + attn.forward(ln_1.forward(x))
        return attended + mlp.forward(ln_2.forward(attended))
    }

    pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32) -> Tensor[Batch, 1, H; T] {
        let attended = x + attn.decode(ln_1.forward(x), pos)
        return attended + mlp.forward(ln_2.forward(attended))
    }

    pub fn prefill<S: Dim>(x: Tensor[Batch, S, H; T], pos: i32) -> Tensor[Batch, S, H; T]
    where S > 0 {
        let attended = x + attn.prefill(ln_1.forward(x), pos)
        return attended + mlp.forward(ln_2.forward(attended))
    }

    pub fn decode_rows(
        x: Tensor[Batch, 1, H; T],
        positions: Tensor[Batch; i32],
    ) -> Tensor[Batch, 1, H; T] {
        let attended = x + attn.decode_rows(ln_1.forward(x), positions)
        return attended + mlp.forward(ln_2.forward(attended))
    }

    pub fn prefill_slots<M: Dim, S: Dim>(
        x: Tensor[M, S, H; T],
        slots: Tensor[M; i32],
    ) -> Tensor[M, S, H; T]
    where
        M > 0,
        S > 0
    {
        let attended = x + attn.prefill_slots(ln_1.forward(x), slots)
        return attended + mlp.forward(ln_2.forward(attended))
    }
}

pub block Model<
    Vocab: Dim,
    MaxPositions: Dim,
    H: Dim,
    Heads: Dim,
    Layers: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float = f32,
>
where
    Heads > 0,
    H % Heads == 0,
    MaxSeq > 0
{
    param wte: Tensor[Vocab, H; T]
    param wpe: Tensor[MaxPositions, H; T]
    sub blocks: [Block<H, Heads, Batch, MaxSeq, T>; Layers]
    sub ln_f: LayerNorm<H, T>

    // Sequences longer than the position table are rejected at compile time.
    pub entry forward<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, Vocab; T]
    where S <= MaxPositions {
        // Positions 0..S come from `iota`; both lookups are the same operation.
        let x0 = embedding(tokens, wte) + embedding(iota<i32>(S), wpe)
        var x = x0
        static for layer in blocks {
            x = layer.forward(x)
        }
        let h = ln_f.forward(x)
        // The head shares the embedding: a contraction against `wte`.
        let logits[b, s, v] = sum[i] h[b, s, i] * wte[v, i]
        return logits
    }

    // Logits for one new token per sequence at position `pos`, attending
    // over the positions decoded before it through the blocks' KV caches.
    // Feed a prompt token by token, then the tokens the model produces.
    pub entry decode(token: Tensor[Batch, 1; i32], pos: i32) -> Tensor[Batch, Vocab; T] {
        let pos_at[d] = wpe[cast<i64>(pos), d]
        let pos_row = reshape(pos_at, [1, H])
        var x = embedding(token, wte) + pos_row
        static for layer in blocks {
            x = layer.decode(x, pos)
        }
        let h = ln_f.forward(x[:, 0, :])
        // The head shares the embedding: a contraction against `wte`.
        let logits[b, v] = sum[i] h[b, i] * wte[v, i]
        return logits
    }

    // A prompt of `S` tokens at `pos`, written into the blocks' KV caches in
    // one pass. Returns the logits after its last token, so `decode`
    // continues at `pos + S`: the time to the first token is this one call.
    pub entry prefill<S: Dim>(tokens: Tensor[Batch, S; i32], pos: i32) -> Tensor[Batch, Vocab; T]
    where S > 0 {
        let rows[i] = cast<i64>(pos) + iota<i64>(S)[i]
        let pos_rows[i, d] = wpe[rows[i], d]
        var x = embedding(tokens, wte) + pos_rows
        static for layer in blocks {
            x = layer.prefill(x, pos)
        }
        let h = ln_f.forward(x[:, S - 1, :])
        // The head shares the embedding: a contraction against `wte`.
        let logits[b, v] = sum[i] h[b, i] * wte[v, i]
        return logits
    }

    // Logits for one new token per sequence, each at its own position
    // `positions[b]`: the step a server takes for every request it is
    // decoding together (continuous batching). A row with no request in it
    // computes along and is ignored.
    pub entry decode_rows(
        tokens: Tensor[Batch, 1; i32],
        positions: Tensor[Batch; i32],
    ) -> Tensor[Batch, Vocab; T] {
        let pos_at[b, d] = wpe[cast<i64>(positions[b]), d]
        var x = embedding(tokens, wte) + reshape(pos_at, [Batch, 1, H])
        static for layer in blocks {
            x = layer.decode_rows(x, positions)
        }
        let h = ln_f.forward(x[:, 0, :])
        let logits[b, v] = sum[i] h[b, i] * wte[v, i]
        return logits
    }

    // `M` requests' prompts into rows `slots` of the caches, as they join a
    // batch being decoded: one pass for all of them. Row `m` of `tokens` holds
    // `S` tokens of which the first `lengths[m]` are its prompt; the rest pad
    // it to one of a few compiled lengths and are never read. Returns the
    // logits after each prompt's last token; `decode_rows` continues row
    // `slots[m]` at position `lengths[m]`.
    pub entry prefill_slots<M: Dim, S: Dim>(
        tokens: Tensor[M, S; i32],
        slots: Tensor[M; i32],
        lengths: Tensor[M; i32],
    ) -> Tensor[M, Vocab; T]
    where
        M > 0,
        S > 0,
        S <= MaxPositions
    {
        var x = embedding(tokens, wte) + embedding(iota<i32>(S), wpe)
        static for layer in blocks {
            x = layer.prefill_slots(x, slots)
        }
        let last[b, i] = x[b, cast<i64>(lengths[b]) - 1, i]
        let h = ln_f.forward(last)
        let logits[b, v] = sum[i] h[b, i] * wte[v, i]
        return logits
    }
}

Open on GitHub

nest.toml789 B
toml
[model]
name = "gpt2"
title = "GPT-2 (124M)"
summary = "OpenAI's 124M-parameter GPT-2: learned positions, pre-norm blocks, a fused QKV projection, and a head tied to the token embedding."
license = "MIT"
family = "gpt2"
tags = ["text-generation", "decoder-only"]

[links]
huggingface = "https://huggingface.co/openai-community/gpt2"
github = "https://github.com/openai/gpt-2"
homepage = "https://openai.com/research/better-language-models"

[source]
path = "gpt2.linnet"
root = "Model"
entry = "forward"

[generics]
Vocab = 50257
MaxPositions = 1024
H = 768
Heads = 12
Layers = 12
Batch = 1
MaxSeq = 1024
T = "f32"

[check]
B = 1
S = 8

[weights]
repo = "openai-community/gpt2"
revision = "607a30d783dfa663caf39e06633721c8d4cfcd7e"
files = ["model.safetensors"]
bindings = "bindings.json"

Open on GitHub

samples.json1.3 kB
json
{
  "kind": "text",
  "prompt": "The printing press changed Europe because",
  "chat": false,
  "output": "of the printing press.\n\nThe printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The",
  "reference": "of the printing press.\n\nThe printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The printing press changed Europe because of the printing press. The",
  "reference_stack": "transformers (KV cache, argmax, bf16)",
  "tokens": 96,
  "agreeing_prefix": 96
}

Open on GitHub

model.safetensorsHugging Face ↗