use topos::checkpoint::named_restore;
use topos::{
Element, LayerNorm, Linear, Module, Parameters, Path, Segment, Symbol, Tape, Tensor, Value,
Visitor, concat, named_parameters, scaled_dot_product,
};
use crate::weights::Weights;
pub const CONTEXT_LEN: usize = 256;
pub const EMBED_DIM: usize = 768;
const HEAD_COUNT: usize = 12;
const HEAD_DIM: usize = EMBED_DIM / HEAD_COUNT;
const LAYER_COUNT: usize = 12;
pub const VOCABULARY_LEN: usize = 50257;
const POSITION_COUNT: usize = 1024;
#[derive(Clone)]
struct Gelu {
half: Symbol,
one: Symbol,
root: Symbol,
coefficient: Symbol,
}
impl Gelu {
fn new<E: Element + From<f32>>(tape: &Tape<E>) -> Self {
let scalar = |value: f32| tape.leaf(Tensor::filled([], E::from(value))).symbol();
Self {
half: scalar(0.5),
one: scalar(1.0),
root: scalar(0.797_884_6),
coefficient: scalar(0.044_715),
}
}
}
impl<E: Element> Module<E> for Gelu {
fn express<'tape>(&self, input: Value<'tape, E>) -> Value<'tape, E> {
let tape = input.tape();
let half = tape.resolve(self.half);
let one = tape.resolve(self.one);
let root = tape.resolve(self.root);
let coefficient = tape.resolve(self.coefficient);
let cubic = input * input * input * coefficient.broadcast_like(input);
let inner = ((input + cubic) * root.broadcast_like(input)).tanh();
input * (inner + one.broadcast_like(inner)) * half.broadcast_like(input)
}
}
struct Decoded<'tape, E> {
output: Value<'tape, E>,
keys: Value<'tape, E>,
values: Value<'tape, E>,
}
struct Attention<E> {
fused: Linear<E>,
projection: Linear<E>,
mask: Symbol,
scale: Symbol,
}
impl<E: Element + From<f32>> Attention<E> {
fn new(tape: &Tape<E>, mask: Symbol, scale: Symbol) -> Self {
let zeros = |shape: [usize; 2]| Tensor::filled(shape, E::from(0.0));
let bias = |extent: usize| Tensor::filled([extent], E::from(0.0));
Self {
fused: Linear::new(tape, zeros([EMBED_DIM, 3 * EMBED_DIM]), bias(3 * EMBED_DIM)),
projection: Linear::new(tape, zeros([EMBED_DIM, EMBED_DIM]), bias(EMBED_DIM)),
mask,
scale,
}
}
}
impl<E: Element> Attention<E> {
fn express_decode<'tape>(
&self,
tape: &'tape Tape<E>,
input: Value<'tape, E>,
keys: Value<'tape, E>,
values: Value<'tape, E>,
position: Value<'tape, E>,
mask: Value<'tape, E>,
) -> Decoded<'tape, E> {
let scale = tape.resolve(self.scale);
let fused = self.fused.express(input);
let keys = keys + fused.narrow(1, EMBED_DIM, EMBED_DIM).scatter(position);
let values = values + fused.narrow(1, 2 * EMBED_DIM, EMBED_DIM).scatter(position);
let heads: Vec<Value<'tape, E>> = (0..HEAD_COUNT)
.map(|head| {
let query = fused.narrow(1, head * HEAD_DIM, HEAD_DIM);
let key = keys.narrow(1, head * HEAD_DIM, HEAD_DIM);
let value = values.narrow(1, head * HEAD_DIM, HEAD_DIM);
scaled_dot_product(query, key, value, mask, scale)
})
.collect();
Decoded {
output: self.projection.express(concat(&heads, 1)),
keys,
values,
}
}
}
impl<E: Element> Module<E> for Attention<E> {
fn express<'tape>(&self, input: Value<'tape, E>) -> Value<'tape, E> {
let tape = input.tape();
let mask = tape.resolve(self.mask);
let scale = tape.resolve(self.scale);
let fused = self.fused.express(input);
let heads: Vec<Value<'tape, E>> = (0..HEAD_COUNT)
.map(|head| {
let query = fused.narrow(1, head * HEAD_DIM, HEAD_DIM);
let key = fused.narrow(1, EMBED_DIM + head * HEAD_DIM, HEAD_DIM);
let value = fused.narrow(1, 2 * EMBED_DIM + head * HEAD_DIM, HEAD_DIM);
scaled_dot_product(query, key, value, mask, scale)
})
.collect();
self.projection.express(concat(&heads, 1))
}
fn visit(&self, visitor: &mut dyn Visitor) {
visitor.enter(Segment::Name("c_attn"));
self.fused.visit(visitor);
visitor.leave();
visitor.enter(Segment::Name("c_proj"));
self.projection.visit(visitor);
visitor.leave();
}
}
struct FeedForward<E> {
up: Linear<E>,
activation: Gelu,
down: Linear<E>,
}
impl<E: Element + From<f32>> FeedForward<E> {
fn new(tape: &Tape<E>, activation: Gelu) -> Self {
let zeros = |shape: [usize; 2]| Tensor::filled(shape, E::from(0.0));
let bias = |extent: usize| Tensor::filled([extent], E::from(0.0));
Self {
up: Linear::new(tape, zeros([EMBED_DIM, 4 * EMBED_DIM]), bias(4 * EMBED_DIM)),
activation,
down: Linear::new(tape, zeros([4 * EMBED_DIM, EMBED_DIM]), bias(EMBED_DIM)),
}
}
}
impl<E: Element> Module<E> for FeedForward<E> {
fn express<'tape>(&self, input: Value<'tape, E>) -> Value<'tape, E> {
let lifted = self.up.express(input);
let hidden = self.activation.express(lifted);
self.down.express(hidden)
}
fn visit(&self, visitor: &mut dyn Visitor) {
visitor.enter(Segment::Name("c_fc"));
self.up.visit(visitor);
visitor.leave();
visitor.enter(Segment::Name("c_proj"));
self.down.visit(visitor);
visitor.leave();
}
}
struct Block<E> {
attention_norm: LayerNorm<E>,
attention: Attention<E>,
hidden_norm: LayerNorm<E>,
feed_forward: FeedForward<E>,
}
impl<E: Element + From<f32>> Block<E> {
fn new(tape: &Tape<E>, mask: Symbol, scale: Symbol, activation: Gelu) -> Self {
Self {
attention_norm: layer_norm(tape),
attention: Attention::new(tape, mask, scale),
hidden_norm: layer_norm(tape),
feed_forward: FeedForward::new(tape, activation),
}
}
}
impl<E: Element> Block<E> {
fn express_decode<'tape>(
&self,
tape: &'tape Tape<E>,
input: Value<'tape, E>,
keys: Value<'tape, E>,
values: Value<'tape, E>,
position: Value<'tape, E>,
mask: Value<'tape, E>,
) -> Decoded<'tape, E> {
let attended = self.attention.express_decode(
tape,
self.attention_norm.express(input),
keys,
values,
position,
mask,
);
let stream = input + attended.output;
let lifted = self.feed_forward.express(self.hidden_norm.express(stream));
Decoded {
output: stream + lifted,
keys: attended.keys,
values: attended.values,
}
}
}
impl<E: Element> Module<E> for Block<E> {
fn express<'tape>(&self, input: Value<'tape, E>) -> Value<'tape, E> {
let attended = self.attention.express(self.attention_norm.express(input));
let stream = input + attended;
let lifted = self.feed_forward.express(self.hidden_norm.express(stream));
stream + lifted
}
fn visit(&self, visitor: &mut dyn Visitor) {
visitor.enter(Segment::Name("ln_1"));
self.attention_norm.visit(visitor);
visitor.leave();
visitor.enter(Segment::Name("attn"));
self.attention.visit(visitor);
visitor.leave();
visitor.enter(Segment::Name("ln_2"));
self.hidden_norm.visit(visitor);
visitor.leave();
visitor.enter(Segment::Name("mlp"));
self.feed_forward.visit(visitor);
visitor.leave();
}
}
fn layer_norm<E: Element + From<f32>>(tape: &Tape<E>) -> LayerNorm<E> {
LayerNorm::new(
tape,
Tensor::filled([EMBED_DIM], E::from(1.0)),
Tensor::filled([EMBED_DIM], E::from(0.0)),
Tensor::filled([], E::from(1e-5)),
)
}
pub struct Gpt2<E> {
embeddings: Symbol,
positions: Symbol,
blocks: Vec<Block<E>>,
final_norm: LayerNorm<E>,
}
impl<E: Element + From<f32> + 'static> Gpt2<E> {
pub fn new(tape: &Tape<E>) -> Self {
let embeddings = tape
.parameter(Tensor::filled([VOCABULARY_LEN, EMBED_DIM], E::from(0.0)))
.symbol();
let positions = tape
.parameter(Tensor::filled([POSITION_COUNT, EMBED_DIM], E::from(0.0)))
.symbol();
let mask_elements: Vec<E> = (0..CONTEXT_LEN * CONTEXT_LEN)
.map(|at| {
if at % CONTEXT_LEN <= at / CONTEXT_LEN {
E::from(0.0)
} else {
E::from(f32::NEG_INFINITY)
}
})
.collect();
let mask = tape
.leaf(Tensor::new([CONTEXT_LEN, CONTEXT_LEN], mask_elements))
.symbol();
let scale = tape
.leaf(Tensor::filled([], E::from(1.0 / (HEAD_DIM as f32).sqrt())))
.symbol();
let activation = Gelu::new(tape);
let blocks = (0..LAYER_COUNT)
.map(|_| Block::new(tape, mask, scale, activation.clone()))
.collect();
Self {
embeddings,
positions,
blocks,
final_norm: layer_norm(tape),
}
}
pub fn embeddings(&self) -> Symbol {
self.embeddings
}
}
impl<E: Element> Gpt2<E> {
pub fn layers(&self) -> usize {
self.blocks.len()
}
pub fn express_decode<'tape>(
&self,
tape: &'tape Tape<E>,
embedded: Value<'tape, E>,
caches: &[(Value<'tape, E>, Value<'tape, E>)],
position: Value<'tape, E>,
mask: Value<'tape, E>,
) -> (Value<'tape, E>, Vec<(Value<'tape, E>, Value<'tape, E>)>) {
assert_eq!(caches.len(), self.blocks.len(), "one cache pair per block");
let positions = tape.resolve(self.positions);
let row = positions.narrow(0, 0, CONTEXT_LEN).gather(position);
let mut stream = embedded + row;
let mut updated = Vec::with_capacity(caches.len());
for (block, &(keys, values)) in self.blocks.iter().zip(caches) {
let decoded = block.express_decode(tape, stream, keys, values, position, mask);
stream = decoded.output;
updated.push((decoded.keys, decoded.values));
}
(self.final_norm.express(stream), updated)
}
}
impl<E: Element> Module<E> for Gpt2<E> {
fn express<'tape>(&self, input: Value<'tape, E>) -> Value<'tape, E> {
let tape = input.tape();
let positions = tape.resolve(self.positions);
let context = input.shape().axes()[0];
let stream = input + positions.narrow(0, 0, context);
let stream = self
.blocks
.iter()
.fold(stream, |value, block| block.express(value));
self.final_norm.express(stream)
}
fn visit(&self, visitor: &mut dyn Visitor) {
visitor.enter(Segment::Name("wte"));
visitor.parameter("weights", self.embeddings);
visitor.leave();
visitor.enter(Segment::Name("wpe"));
visitor.parameter("weights", self.positions);
visitor.leave();
visitor.enter(Segment::Name("h"));
for (index, block) in self.blocks.iter().enumerate() {
visitor.enter(Segment::Index(index));
block.visit(visitor);
visitor.leave();
}
visitor.leave();
visitor.enter(Segment::Name("ln_f"));
self.final_norm.visit(visitor);
visitor.leave();
}
}
fn foreign_name(path: &Path) -> String {
let segments = path.segments();
let mut name = String::new();
for (position, segment) in segments.iter().enumerate() {
if position > 0 {
name.push('.');
}
if position + 1 < segments.len() {
name.push_str(&segment.to_string());
continue;
}
let leaf = match segment {
Segment::Name("weights") | Segment::Name("scale") => "weight",
Segment::Name("shift") | Segment::Name("bias") => "bias",
other => panic!("no checkpoint spelling for the leaf `{other}`"),
};
name.push_str(leaf);
}
name
}
pub fn load<E: Element + From<f32>>(
parameters: &Parameters<E>,
model: &Gpt2<E>,
weights: &Weights,
) -> Parameters<E> {
let entries: Vec<(Path, Tensor<E>)> = named_parameters(model)
.into_iter()
.map(|(path, _)| {
let payload = weights.tensor(&foreign_name(&path)).convert::<E>();
(path, payload)
})
.collect();
named_restore(parameters, model, entries)
}