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)