Skip to main content

Author a model in Python

A model authored in Python over clika_runtime runs at the speed of the runtime's operators: Python decides which operator runs next, the kernels do the work, and the loop never reads a value it does not need. This guide walks the wheel's own chapter, examples/python/howto/author_a_model_in_python/llama/llama_from_scratch.py, a Llama-3.2-class decoder in about six hundred lines that loads a Hugging Face checkpoint, generates greedily and matches the clika-modelverse command line token for token. Every sample below is an excerpt of that file; the chapter's README holds the command lines and the measured table.

The walk is linear: the configuration, the module tree, the load, the fusions, the cache, the attention step, the decode loop, the measurement. Use ClikaRT from Python covers the tensor and module basics this page builds on; Tokenize text and apply a chat template covers the tokenizer the prompt goes through.

The configuration comes from config.json​

The architecture is a dataclass whose fields carry the checkpoint's own names, so config.json fills it directly. from_dict keeps the keys the dataclass declares and drops the rest, and __post_init__ derives what the file leaves implicit (the key/value head count, the head size). The rotary tables come from one operator, with Llama 3's frequency scaling when the checkpoint declares it.

llama_from_scratch.py (excerpt)
@dataclasses.dataclass
class LlamaConfig:
hidden_size: int
intermediate_size: int
num_hidden_layers: int
num_attention_heads: int
vocab_size: int
num_key_value_heads: int | None = None
head_dim: int | None = None
rms_norm_eps: float = 1e-5
rope_theta: float = 10000.0
rope_scaling: dict | None = None
tie_word_embeddings: bool = False
max_position_embeddings: int = 8192
eos_token_id: int | list[int] | None = None

def __post_init__(self) -> None:
if self.num_key_value_heads is None:
self.num_key_value_heads = self.num_attention_heads
if self.head_dim is None:
self.head_dim = self.hidden_size // self.num_attention_heads

@classmethod
def from_dict(cls, fields: dict) -> "LlamaConfig":
known = inspect.signature(cls).parameters
return cls(**{name: value for name, value in fields.items() if name in known})

def rotary_tables(self, max_positions: int, device: crt.Device) -> tuple[crt.Tensor, crt.Tensor]:
scaling = self.rope_scaling or {}
if self.rope_llama3:
return crt.generate_rotary_cache(
self.head_dim, max_positions, theta=self.rope_theta, scaling="llama3",
scale=float(scaling["factor"]),
low_freq_factor=float(scaling["low_freq_factor"]),
high_freq_factor=float(scaling["high_freq_factor"]),
original_max_pos=int(scaling["original_max_position_embeddings"]),
device=device,
)
return crt.generate_rotary_cache(self.head_dim, max_positions, theta=self.rope_theta,
device=device)

A module tree with the checkpoint's names​

load_state_dict binds by dotted name, so the tree spells the checkpoint's names: model.layers.3.self_attn.q_proj.weight is the attribute path model.layers[3].self_attn.q_proj and its weight. Every constructor takes the model dtype and declares each layer at it (the chapter resolves it from --dtype, else from the checkpoint's embedding table), so a bind adopts a matching checkpoint tensor as it is and no layer sits at the default dtype beside a bfloat16 checkpoint. The attention module owns the four projections and, after the load, one fused group for the first three.

llama_from_scratch.py (excerpt)
class LlamaAttention(nn.Module):
def __init__(self, config: LlamaConfig, dtype: crt.dtype) -> None:
super().__init__()
hidden, heads, kv_heads, head_dim = (config.hidden_size, config.num_attention_heads,
config.num_key_value_heads, config.head_dim)
self.num_heads, self.num_kv_heads, self.head_dim = heads, kv_heads, head_dim
self.q_proj = nn.Linear(hidden, heads * head_dim, bias=False, dtype=dtype)
self.k_proj = nn.Linear(hidden, kv_heads * head_dim, bias=False, dtype=dtype)
self.v_proj = nn.Linear(hidden, kv_heads * head_dim, bias=False, dtype=dtype)
self.o_proj = nn.Linear(heads * head_dim, hidden, bias=False, dtype=dtype)


class LlamaMLP(nn.Module):
def __init__(self, config: LlamaConfig, dtype: crt.dtype) -> None:
super().__init__()
self.gate_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False, dtype=dtype)
self.up_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False, dtype=dtype)
self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=False, dtype=dtype)

def forward(self, x: crt.Tensor) -> crt.Tensor:
return self.down_proj(crt.swiglu(self.gate_up(x)))

LlamaDecoderLayer holds one attention, one MLP and the two nn.RMSNorm weights; LlamaModel holds the nn.Embedding, an nn.ModuleList of layers and the final norm; LlamaForCausalLM adds lm_head and keeps the model dtype as self.dtype. The names are the checkpoint's at every level.

Build storage-free, then adopt the checkpoint​

The checkpoint comes first: crt.load reads the safetensors file into tensors on the device, and a tied checkpoint, which carries no lm_head.weight, gets the embedding table bound under the head's name too, so one load adopts one object onto both slots. The model dtype is the given --dtype, else the embedding table's own. The tree is then built under crt.device("meta"), where every parameter is a shape and a dtype with no bytes; to_empty gives the slots their placement, and load_state_dict(state, strict=True, assign=True) adopts the checkpoint tensors as the model's parameters: no copy, one resident copy of every weight.

llama_from_scratch.py (excerpt)
def build_model(config, snapshot, device, max_positions, dtype=None, state=None):
if state is None:
state = load_state(snapshot, device)
state = dict(state)
if config.tie_word_embeddings and LM_HEAD_KEY not in state:
state[LM_HEAD_KEY] = state[EMBED_KEY]
model_dtype = dtype if dtype is not None else state[EMBED_KEY].dtype
if dtype is not None:
state = cast_state(state, dtype)
with crt.device("meta"):
model = LlamaForCausalLM(config, model_dtype)
model.to_empty(device=device)
model.load_state_dict(state, strict=True, assign=True)
del state
model.tie_weights()
model.fuse()
model.model.build_rotary(max_positions, device)
model.eval()
return model

tie_weights is a check on a tied checkpoint: the head and the embedding must hold one object, and it raises when they do not. After the load, model.state_dict() lists the checkpoint's names in the checkpoint's order. A sharded checkpoint loads the same way: load_state reads the shards the index file names into one dict.

check_load.py
model = build_model(config, snapshot, crt.Device.cpu(), max_positions=4096)
print(model.lm_head.weight is model.model.embed_tokens.weight) # True
print(model.lm_head.weight.dtype) # clika_runtime.bfloat16
print(len(model.state_dict())) # 147

Fuse the projections after the load​

Three projections over the same input are one matmul over the concatenated weights. nn.fuse_linears returns the group as an nn.Linear whose forward emits the parts' outputs side by side; the parts keep their state-dict names and read their rows back from the fused weight, and the group itself lists no parameters, so holding it as an attribute adds nothing to a save. Fuse once the weights are bound and before the first forward.

llama_from_scratch.py (excerpt)
def fuse(self) -> None:
self.qkv = nn.fuse_linears(self.q_proj, self.k_proj, self.v_proj, name="qkv")

# in LlamaMLP
def fuse(self) -> None:
self.gate_up = nn.fuse_linears(self.gate_proj, self.up_proj, name="gate_up")

The forward splits the fused output where the parts meet:

llama_from_scratch.py (excerpt)
qkv = self.qkv(x) # [tokens, (heads + 2 * kv_heads) * head_dim]
q, k, v = crt.split_with_sizes(qkv, [heads * head_dim, kv_heads * head_dim, kv_heads * head_dim], -1)

The KV cache and the step indices​

nn.KVCache owns one key row and one value row per layer, sized once from the configuration. prepare_step takes the cumulative query lengths, the past length of each sequence and its slot, and returns the nn.StepIndices the attention operator reads: cu_seqlens_q, cu_seqlens_k, kvcache_start, slot_ids and max_seqlen_k, all device tensors except the last.

llama_from_scratch.py (excerpt)
def make_cache(config, model, max_positions, device):
kv_config = nn.KVCacheConfig(
num_layers=config.num_hidden_layers, num_kv_heads=config.num_key_value_heads,
head_dim=config.head_dim, kv_dtype=model.model.embed_tokens.weight.dtype,
max_seqs=1, max_tokens_per_seq=max_positions, device=device,
)
return nn.KVCache.make(kv_config, [nn.KVLayerSpec(preallocate=True) for _ in range(config.num_hidden_layers)])

cache = make_cache(config, model, 4096, device)
step = cache.prepare_step([0, len(prompt_ids)], [0], [0]) # the prefill of one sequence

Layer i reads its rows through cache.keys(i) and cache.values(i); the same handles receive the appended keys and values.

The attention step is one operator​

Rotation at the positions the cache start implies, the append of the new keys and values into the cache row, and the causal grouped-query attention over the whole row are one call, for the prefill and for every decode step alike. The rows are passed as past_key / past_value and again as out_present_key / out_present_value, so the append lands in place.

llama_from_scratch.py (excerpt)
attention, _, _ = crt.group_query_attention_varlen(
q, k, v, step.cu_seqlens_q, step.cu_seqlens_k, max_seqlen_k=step.max_seqlen_k,
past_key=past_key, past_value=past_value, kvcache_start=step.kvcache_start,
rope_cos=rope_cos, rope_sin=rope_sin, is_causal=True, rotary_mode="neox",
num_heads=heads, kv_num_heads=kv_heads,
out_present_key=past_key, out_present_value=past_value, slot_ids=step.slot_ids,
)
return self.o_proj(attention)

A residual add and the norm after it, fused​

Each block adds a residual and normalizes the sum twice. crt.add_rms_norm does both in one operator and returns the pair the next block needs: the normalized stream and the raw residual. The closing call takes the norm weight that follows the block, the next layer's input_layernorm or the model's final norm, so the first crt.rms_norm on the embedding is the only standalone norm in the model.

llama_from_scratch.py (excerpt)
o = self.self_attn(x, step, past_key, past_value, rope_cos, rope_sin)
x1, h1 = crt.add_rms_norm(o, h, None, None, [self.hidden_size], None,
self.post_attention_layernorm.weight, None, self.eps)
m = self.mlp(x1)
x2, h2 = crt.add_rms_norm(m, h1, None, None, [self.hidden_size], None, next_gain, None, self.eps)
return x2, h2

The head reads each sequence's last row from the device, crt.index_select(x, -2, step.cu_seqlens_q[1:] - 1), so no shape is read on the host inside the forward.

Greedy decoding with a one-token lookahead​

Operators return before their work runs, and a host read waits for the value it needs. The loop uses that: after the prefill, each iteration queues the next step on the device first, with the argmax tensor itself as its input, and reads the current token afterwards. The device is never idle while Python reads a token, and every timing is the host clock at a token read.

llama_from_scratch.py (excerpt)
cache.reset()
prompt = crt.tensor(list(prompt_ids), dtype=crt.int32, device=device)
length = len(prompt_ids)
logits = model(prompt, cache.prepare_step([0, length], [0], [0]), cache)
token = crt.argmax(logits, dims=[-1], index_dtype=crt.int32) # [1]
crt.async_eval(token)
past = length
ids: list[int] = []
for produced in range(max_new_tokens):
queued = None
if produced + 1 < max_new_tokens:
next_logits = model(token, cache.prepare_step([0, 1], [past], [0]), cache)
queued = crt.argmax(next_logits, dims=[-1], index_dtype=crt.int32)
crt.async_eval(queued)
value = int(token.item())
ids.append(value)
if value in stops or queued is None:
break
token = queued
past += 1

crt.async_eval submits the queued step without waiting; token.item() is the one read per token. The stop ids come from generation_config.json, with config.json as the fallback, the same source the command line reads.

Run it and measure it​

The prompt goes through the checkpoint's chat template as one user turn with the assistant priming, rendered at the wall clock (now_epoch_seconds=int(time.time())), which is how clika-modelverse <snapshot> prompt renders it, so the two replies compare byte for byte.

python3 llama_from_scratch.py --model meta-llama/Llama-3.2-1B-Instruct \
--prompt 'Write a long story about a lighthouse keeper who finds a map.' \
--max-new-tokens 64 --device cpu --json reply.json

--json writes the timings, the memory readings from clika_runtime.memory_stats(device) at four points of the run, the rendered prompt and the token ids. The chapter's measure_decode.py --full runs the same cell on both sides, alternating, and writes the report:

python3 measure_decode.py --full \
--cli /path/to/clika-modelverse --snapshot /path/to/checkpoint --device cpu \
--isl 128 --osl 64 --iters 5 --warmup 1 --repeats 3 \
--prompt 'Write a long story about a lighthouse keeper who finds a map.' --max-new-tokens 64 \
--out ./llama_cpu

The chapter README carries the measured table for this cell, the machine it ran on, how each column is defined, and how to read the spread against the difference between the two rows. A number is a claim about one command on one machine: quote it with its command lines and its cell.

From here: Structure inputs and outputs as pytrees for the containers the compile and save boundaries take, Use ClikaRT with PyTorch for a module tree that starts as torch.nn.Module, and Trace eager code to graphs for capturing a step as a graph.