ZeLM-118M-ID
Model bahasa Indonesia 118M parameter β dilatih dari nol di TPU v5e-8 (Kaggle, gratis).
π Daftar Isi
- Deskripsi
- Arsitektur
- Tokenizer
- Training
- Status Checkpoint
- Cara Penggunaan
- Intended Use & Limitasi
- Infrastruktur
- Roadmap
Deskripsi
ZeLM-118M-ID adalah causal language model untuk Bahasa Indonesia, dilatih dari nol (from scratch) di TPU v5e-8 lewat Kaggle Notebook gratis.
Ini proyek pretraining independen β dirancang dan dijalankan dengan pipeline yang lebih terstruktur dibanding eksperimen-eksperimen sebelumnya: sharding data-parallel 8-core sejak awal, gradient accumulation, evaluasi berkala di val set, logging metrik lengkap, dan checkpoint yang membawa state optimizer penuh (bukan cuma bobot model) supaya training bisa dilanjut persis dari titik terakhir.
Repo ini berisi checkpoint pretraining periodik, bukan model final. Training masih berjalan β lihat Status Checkpoint untuk step terakhir yang tersedia.
β οΈ Ini bukan model final. Checkpoint di repo ini adalah snapshot training yang sedang berlangsung, diupload berkala tiap 10.000 step. Jangan jadikan acuan performa akhir model.
Arsitektur
| Komponen | Nilai |
|---|---|
| Arsitektur | ZeLM (Causal LM, Transformer decoder custom) |
| Hidden Size (d_model) | 512 |
| Jumlah Layer | 36 |
| Attention Heads | 8 Query / 2 Key-Value (GQA) |
| Head Dimension | 64 |
| FFN Dimension | 1,408 (SwiGLU) |
| Vocab Size | 32,768 |
| Max Sequence Length | 512 token |
| Positional Encoding | RoPE (base=10,000) |
| Normalization | RMSNorm (eps=1e-6) |
| Tied Embeddings | Ya |
| Total Parameter | 118.27M |
36 layer dijalankan lewat nn.scan (satu compiled graph untuk semua layer, bukan 36 blok terpisah) supaya waktu kompilasi JAX tidak meledak β desainnya condong "narrow-deep" (banyak layer, hidden size sedang) ketimbang "wide-shallow".
Tokenizer
HELIX v2 β tokenizer BPE custom dengan skema token khusus \HELIXβnamaβHELIX/, dipakai untuk menandai giliran percakapan (user/assistant/system), blok reasoning, tool-call, dokumen, konteks, retrieval, kode, memori, dan slot multimodal β disiapkan sejak tokenizer meski model belum tentu dilatih untuk semua fungsi tersebut.
- Basis: BPE (Byte-Pair Encoding) β bukan Unigram
- Vocab size: 32,768
- Training data tokenizer mencakup porsi khusus kode sumber (GitHub, StackOverflow) selain korpus Bahasa Indonesia, supaya vocab-nya lebih efisien untuk teks yang bercampur kode
- Token separator antar-dokumen: ID 22
Ada dua varian tokenizer yang sempat dibandingkan β Unigram dan BPE. Dari beberapa uji singkat, BPE menunjukkan hasil yang lebih baik untuk kasus ini, jadi BPE yang dipakai untuk training model ini.
Ini versi kedua dari skema HELIX β ada kemungkinan sedikit perbedaan ID token dari versi sebelumnya karena retraining ulang, belum divalidasi silang 1:1.
Training
Dataset
Dilatih dari korpus web Bahasa Indonesia yang sudah dikurasi/difilter kualitasnya, dikombinasikan dengan sumber teks formal (ensiklopedis). Detail komposisi dan sumber spesifik tidak dipublikasikan.
| Item | Detail |
|---|---|
| Total token | ~2.4B token (target), dikemas jadi train.bin (5GB, uint16) + val.bin |
Token dikemas rapat tanpa padding (packing) β dokumen dirangkai jadi satu stream panjang dipisah token separator, lalu dipotong per seq_len=512. Disimpan sebagai uint16 numpy array mentah, dibaca lewat iterator custom saat training.
Konfigurasi Training
| Parameter | Nilai |
|---|---|
| Sequence Length | 512 token |
| Batch Size Efektif | 128 |
| Micro-batch Size | 64 (8 per core, 8 TPU core) |
| Gradient Accumulation | 2 micro-step per update |
| Learning Rate | 3e-4 puncak, cosine decay ke 3% (min_lr_ratio) |
| Warmup Steps | 500 |
| Weight Decay | 0.1 |
| Gradient Clipping | max_norm=1.0 |
| Optimizer | Muon (parameter matriks) + AdamW Ξ²1=0.9, Ξ²2=0.95 (parameter non-matriks) |
| z-loss coefficient | 1e-4 |
| Precision | bfloat16 komputasi, fp32 master params |
| Target Total Steps | 114,705 (3 epoch dari dataset) |
Optimizer Muon menangani parameter berbentuk matriks (attention, FFN), sedangkan AdamW menangani parameter non-matriks (embedding, bias, norm).
Infrastruktur Training
Dilatih di TPU v5e-8 (8 chip) lewat Kaggle Notebook, memakai data-parallel murni (jax.jit + Mesh/NamedSharding) β bukan pmap. Params direplikasi penuh ke semua 8 core (model cukup kecil untuk muat di 1 chip), batch di-shard rata ke 8 core.
Throughput terukur di kisaran ~430,000 token/detik setelah compile warmup selesai.
Status Checkpoint
β οΈ Training masih berjalan. Checkpoint saat ini berhenti di step 41,000 dari target 114,705 (~36%). Training akan dilanjutkan dan checkpoint akan diupdate berkala β cek riwayat commit repo ini untuk versi terbaru.
Checkpoint diupload ke repo ini setiap 10,000 step training. Checkpoint lokal (Kaggle) disimpan tiap 1,000 step dengan retensi 3 file terbaru; checkpoint yang di-upload ke sini mengikuti retensi yang sama.
Setiap file checkpoint (step_XXXXX.npz) berisi parameter model, state optimizer (Muon + AdamW momentum, termasuk state gradient accumulation), dan riwayat log training (loss, perplexity, accuracy, learning rate) sejak awal training β bukan cuma dari sesi terakhir.
Kurva Loss
Loss train (biru) dan val (oranye putus-putus) dari step 0 sampai checkpoint terakhir. Digenerate langsung dari riwayat log yang tersimpan di dalam file checkpoint, bukan dihitung ulang β jadi akan otomatis lebih panjang tiap kali checkpoint dan grafik ini diupdate.
Cara Penggunaan
β οΈ Penting: Checkpoint ini dalam format
.npzmentah (JAX/Flax params + Muon/AdamW opt_state), bukan formatsafetensors/transformers. Belum ada wrapperAutoModelForCausalLMuntuk model ini β load manual lewat JAX/Flax.
Install
pip install jax flax optax sentencepiece huggingface_hub
Load Checkpoint
from huggingface_hub import hf_hub_download
import numpy as np
import jax.numpy as jnp
# ganti nama file sesuai checkpoint terbaru di repo ini
ckpt_path = hf_hub_download(
repo_id="Veenn/zelm-118m-id",
filename="step_41000.npz",
)
ckpt = np.load(ckpt_path)
params = {}
for key in ckpt.files:
if key.startswith("params/"):
path = key[len("params/"):].split(".")
d = params
for p in path[:-1]:
d = d.setdefault(p, {})
d[path[-1]] = jnp.asarray(ckpt[key])
Generate
import jax
import sentencepiece as spm
sp = spm.SentencePieceProcessor(model_file="magnetar_v2_tokenizer.model")
SEP = 22
def generate_text(model, params, prompt_text, max_new_tokens=100, temperature=0.8, top_k=40, seed=0):
rng = jax.random.PRNGKey(seed)
ids = sp.encode(prompt_text, out_type=int, add_bos=False, add_eos=False)
for _ in range(max_new_tokens):
context = ids[-512:]
x = jnp.asarray([context])
logits = model.apply({"params": params}, x, deterministic=True)
next_logits = logits[0, -1, :] / temperature
top_vals, top_idx = jax.lax.top_k(next_logits, top_k)
probs = jax.nn.softmax(top_vals)
rng, subkey = jax.random.split(rng)
choice = jax.random.categorical(subkey, jnp.log(probs))
next_id = int(top_idx[choice])
if next_id == SEP:
break
ids.append(next_id)
return sp.decode(ids)
Definisi model (ZeLM, ZeLMConfig) ada di model.py pada repo ini.
Intended Use & Limitasi
β Cocok untuk
- Eksperimen/riset pipeline training LLM skala kecil di TPU
- Basis fine-tuning downstream task Bahasa Indonesia (setelah training selesai)
- Belajar arsitektur GQA + RoPE + SwiGLU + Muon optimizer
β Tidak cocok untuk
- Penggunaan produksi apa pun β ini checkpoint pretraining yang belum selesai
- Chatbot/asisten langsung β belum instruction-tuned
- Factual QA β base model, belum ada fine-tuning
- Aplikasi yang butuh model stabil/final β checkpoint ini akan berubah tiap update
β οΈ Limitasi yang Diketahui
- Training belum selesai β checkpoint ini snapshot sementara di step 41,000/114,705
- Belum ada evaluasi internal mendalam (probing, head specialization, dll) β akan ditambahkan setelah training selesai
- Tokenizer HELIX v2 kemungkinan punya sedikit perbedaan ID dari versi sebelumnya β belum divalidasi silang
- Belum ada format
safetensors/transformersβ perlu load manual JAX/Flax
Infrastruktur
| Item | Detail |
|---|---|
| Platform | Kaggle Notebook (Free Tier) |
| Accelerator | TPU v5e-8 (8 chip) |
| Paralelisasi | Data-parallel (jax.jit + Mesh/NamedSharding) |
| Throughput | ~430,000 token/detik |
| Framework | JAX + Flax + Optax |
| Precision | bfloat16 (komputasi), fp32 (master params) |
| Budget | $0 |
Roadmap
- Lanjutkan training dari step 41,000 sampai target 114,705 (3 epoch)
- Eval internal β probing classifier, head specialization, kalibrasi setelah training selesai
- Konversi ke
safetensors/transformersuntuk kemudahan penggunaan - ZeLM-118M-ID-Instruct β instruction tuning setelah base model selesai
Dilatih dari nol. Gratis. Untuk Bahasa Indonesia.
ZeLM-118M-ID Β· Apache 2.0 Β· @Veenn
