ZeLM-118M-ID

Model bahasa Indonesia 118M parameter β€” dilatih dari nol di TPU v5e-8 (Kaggle, gratis).

License Language Platform Budget Parameters Status


πŸ“‹ Daftar Isi


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

Training 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 .npz mentah (JAX/Flax params + Muon/AdamW opt_state), bukan format safetensors/transformers. Belum ada wrapper AutoModelForCausalLM untuk 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/transformers untuk 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

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support