The forward entry with one level of blocks expanded. Every edge carries the tensor type the compiler inferred at that point, in the model's own generics.
Entries
| Entry | Signature |
|---|---|
forward | forward<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, Vocab; T] |
decode | decode(token: Tensor[Batch, 1; i32], pos: i32) -> Tensor[Batch, Vocab; T] |
prefill | prefill<S: Dim>(tokens: Tensor[Batch, S; i32], pos: i32) -> Tensor[Batch, Vocab; T] |
decode_rows | decode_rows(tokens: Tensor[Batch, 1; i32], positions: Tensor[Batch; i32]) -> Tensor[Batch, Vocab; T] |
prefill_slots | prefill_slots<M: Dim, S: Dim>(tokens: Tensor[M, S; i32], slots: Tensor[M; i32], lengths: Tensor[M; i32]) -> Tensor[M, Vocab; T] |
Generics
The root block's generics as this checkpoint binds them.
Vocab | 50257 |
MaxPositions | 1024 |
H | 768 |
Heads | 12 |
Layers | 12 |
Batch | 1 |
MaxSeq | 1024 |
T | f32 |
Blocks
Every block of the program with its members and functions, as linnet inspect prints them.
Attention
text
examples.gpt2::Attention<H: Dim, Heads: Dim, Batch: Dim, MaxSeq: Dim, T: Float>
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]
pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32) -> Tensor[Batch, 1, H; T]
pub fn prefill<S: Dim>(x: Tensor[Batch, S, H; T], pos: i32) -> Tensor[Batch, S, H; T]
pub fn decode_rows(x: Tensor[Batch, 1, H; T], positions: Tensor[Batch; i32]) -> Tensor[Batch, 1, H; T]
pub fn prefill_slots<M: Dim, S: Dim>(x: Tensor[M, S, H; T], slots: Tensor[M; i32]) -> Tensor[M, S, H; T]Block
text
examples.gpt2::Block<H: Dim, Heads: Dim, Batch: Dim, MaxSeq: Dim, T: Float>
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]
pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32) -> Tensor[Batch, 1, H; T]
pub fn prefill<S: Dim>(x: Tensor[Batch, S, H; T], pos: i32) -> Tensor[Batch, S, H; T]
pub fn decode_rows(x: Tensor[Batch, 1, H; T], positions: Tensor[Batch; i32]) -> Tensor[Batch, 1, H; T]
pub fn prefill_slots<M: Dim, S: Dim>(x: Tensor[M, S, H; T], slots: Tensor[M; i32]) -> Tensor[M, S, H; T]Conv1D
text
examples.gpt2::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]Embedding
text
std.nn.embedding::Embedding<Vocab: Dim, H: Dim, T: Float = bf16>
param weight: Tensor[Vocab, H; T]
pub fn forward<*S: Shape>(ids: Tensor[*S; i32]) -> Tensor[*S, H; T]LayerNorm
text
examples.gpt2::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]Linear
text
std.nn.linear::Linear<In: Dim, Out: Dim, T: Float = bf16>
param weight: Tensor[Out, In; T]
param bias: Tensor[Out; T]?
pub fn forward<*S: Shape>(x: Tensor[*S, In; T]) -> Tensor[*S, Out; T]Mlp
text
examples.gpt2::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]Model
text
examples.gpt2::Model<Vocab: Dim, MaxPositions: Dim, H: Dim, Heads: Dim, Layers: Dim, Batch: Dim, MaxSeq: Dim, T: Float = f32>
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>
pub entry forward<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, Vocab; T]
pub entry decode(token: Tensor[Batch, 1; i32], pos: i32) -> Tensor[Batch, Vocab; T]
pub entry prefill<S: Dim>(tokens: Tensor[Batch, S; i32], pos: i32) -> Tensor[Batch, Vocab; T]
pub entry decode_rows(tokens: Tensor[Batch, 1; i32], positions: Tensor[Batch; i32]) -> Tensor[Batch, Vocab; T]
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]RmsNorm
text
std.nn.norm::RmsNorm<H: Dim, T: Float = bf16>
param weight: Tensor[H; T]
pub fn forward<*S: Shape>(x: Tensor[*S, H; T]) -> Tensor[*S, H; T]