Pearl-Chat: Native Luganda Language Model in Pure JAX

Pearl-Chat is a decoder-only GPT-2 style transformer trained from scratch on Luganda text using pure JAX and Flax NNX. Engineered for the PyCon Africa 2026 workshop: Building Pearl-Chat: Engineering a Native Language Model from Scratch in Pure JAX.

Model architecture specifications

  • Vocabulary: 4,096 byte-level BPE tokens
  • Context length: 128 tokens
  • Embedding dimension: 128
  • Transformer decoder layers: 4
  • Attention heads: 4
  • Feed-forward dimension: 512
  • Output projection: Weight-tied with token embedding matrix
  • Parameter count: 1,331,986 parameters

Loading and inference with pure JAX and Flax NNX

from huggingface_hub import snapshot_download
from pearlchat.tokenizer import LugandaTokenizer
from pearlchat.model import PearlChatModel
from pearlchat.config import ModelConfig
from pearlchat.checkpointing import CheckpointManager
from flax import nnx

repo_dir = snapshot_download(repo_id="your-username/pearl-chat-luganda")
tokenizer = LugandaTokenizer.load(repo_dir)
config = ModelConfig(vocab_size=tokenizer.vocab_size, context_length=128, embed_dim=128, num_heads=4, num_layers=4, feed_forward_dim=512)
rngs = nnx.Rngs(params=0, dropout=0)
model = PearlChatModel(config, rngs)
manager = CheckpointManager(repo_dir)
model, _ = manager.restore_latest_checkpoint(model)
Downloads last month
-
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Datasets used to train kambale/pearl-chat