Projects/ESA/Documentation

[ ESA / MLBRICKS KIT 1.0.0B1 ]

ESA Quick Start

Start with the ESA layer, recurrent state and the ready-made ESA language model.

DISTRIBUTIONmlbricks-kitDOCS VERSION1.0.0b1

QUICK START

Use ESA as a sequence layer.

PYTHON
import torch
from mlbricks import ESA

x = torch.randn(2, 128, 384)
layer = ESA(embd=384, head=6, backend="auto")
y = layer(x)
print(y.shape)

RECURRENT INFERENCE

Prefill once, then advance the state.

PYTHON
prompt = torch.randn(1, 64, 384)
out, state = layer.prefill(prompt)

next_token = torch.randn(1, 1, 384)
out, state = layer.decode_step(next_token, state)

READY-MADE MODEL

PYTHON
from mlbricks import ESAModel, ESAModelConfig

cfg = ESAModelConfig(
    vocab_size=50_257,
    block=2048,
    n_layer=6,
    head=6,
    embd=384,
)
model = ESAModel(cfg, device="auto")

UNIFIED LIFECYCLE

PYTHON
import mlbricks as mlb

mlb.save(model, "esa_model")
model = mlb.load("esa_model", device="auto")
info = mlb.inspect(model)