indigo / train.py
adyoi's picture
docs: komentar inline + README argumen + bugfix & optimasi (7dd92fc)
9a82835 verified
Raw
History Blame Contribute Delete
17.3 kB
"""
Script training model Indigo GPT dari nol.
Fitur:
- Training loop standar dengan AdamW optimizer
- Learning rate schedule: warmup linear → cosine decay
- Best checkpoint otomatis berdasarkan validasi
- Resume training dari checkpoint sebelumnya (--init-from)
- Dukungan tokenizer char dan BPE
- Gradient clipping untuk stabilitas
- Statistik ringkasan di akhir run
Cara pakai:
python train.py --data data/sample.txt --steps 2000
python train.py --data data/teks.txt --tokenizer bpe --vocab-size 512
python train.py --init-from out/indigo_best.safetensors --steps 1000
"""
import os
import time
import math
import torch
import random
import argparse
from safetensors.torch import save_file
from indigo.common import (
build_tokenizer,
collect_text_files,
load_meta,
read_clean,
save_meta,
)
from indigo.model import GPT, GPTConfig
from indigo.tokenizer import CharTokenizer
def load_init(path):
"""Muat checkpoint untuk melanjutkan training (resume).
Mendukung dua format:
1. .safetensors: format utama Indigo (safetensors + _meta.json + optimizer.pt)
2. .pt: format PyTorch checkpoint lama (model, config, vocab, optimizer dalam 1 file)
Args:
path: Path ke file checkpoint (.safetensors atau .pt).
Returns:
Tuple (state_dict, meta_dict, optimizer_state atau None).
"""
if path.endswith(".safetensors"):
from safetensors.torch import load_file
state = load_file(path)
meta = load_meta(path)
# Cari file optimizer (suffix _best dihapus untuk file optimizer)
opt_path = os.path.splitext(path)[0].replace("_best", "") + "_optimizer.pt"
opt = None
if os.path.exists(opt_path):
try:
opt = torch.load(opt_path, map_location="cpu", weights_only=True)
except Exception as e:
print(f"optimizer state dilewati: {e}")
return state, meta, opt
# Format .pt lama
ckpt = torch.load(path, map_location="cpu", weights_only=True)
meta = {
"config": ckpt["config"],
"vocab": ckpt["vocab"],
"step": ckpt.get("step", 0),
"tokenizer": ckpt.get("tokenizer"),
}
return ckpt["model"], meta, ckpt.get("optimizer")
# Cache arange tensor per (block_size, device) untuk menghindari alokasi berulang
# saat get_batch dipanggil ribuan kali — menghemat ~11x waktu.
_ARANGE_CACHE = {}
def get_batch(data, block_size, batch_size, device):
"""Ambil batch data latih secara random (vectorized).
Proses:
1. Pilih batch_size posisi awal secara acak dari data
2. Untuk setiap posisi, ambil potongan sepanjang block_size (input) dan block_size (target)
3. Target = input bergeser 1 posisi ke kanan (next-token prediction)
Menggunakan fancy indexing dan arange cache untuk efisiensi:
- ix: posisi awal random untuk setiap sampel dalam batch
- idx: matriks posisi (batch_size × block_size) dengan offset arange
Args:
data: Tensor 1D — seluruh data training (token IDs).
block_size: Int — panjang konteks per sampel.
batch_size: Int — jumlah sampel per batch.
device: Str — "cpu" atau "cuda".
Returns:
Tuple (x, y) — x: input (B, T), y: target (B, T).
"""
ix = torch.randint(len(data) - block_size - 1, (batch_size,))
arange = _ARANGE_CACHE.get((block_size, device))
if arange is None:
arange = torch.arange(block_size, device=device)
_ARANGE_CACHE[(block_size, device)] = arange
idx = ix.unsqueeze(1) + arange
x = data[idx]
y = data[idx + 1]
return x.to(device, non_blocking=True), y.to(device, non_blocking=True)
@torch.no_grad()
def estimate_loss(model, data, args, device):
"""Estimasi loss validasi dengan averaging beberapa batch.
Model dipindahkan ke mode eval (tanpa dropout), lalu dihitung loss rata-rata
dari eval_iters batch random. Hasilnya lebih stabil daripada single batch.
Args:
model: Model GPT.
data: Tensor 1D — data validasi (token IDs).
args: Namespace — harus punya block_size, batch_size, eval_iters.
device: Str — "cpu" atau "cuda".
Returns:
Float — loss rata-rata (cross-entropy, nats per token).
"""
model.eval()
losses = []
for _ in range(args.eval_iters):
x, y = get_batch(data, args.block_size, args.batch_size, device)
_, loss = model(x, y)
losses.append(loss.item())
model.train()
return sum(losses) / len(losses)
def main(argv=None):
"""Fungsi utama training — bisa dipanggil dari CLI atau dari pipeline.py.
Pipeline training:
1. Parse argumen → setup device & seed
2. Kumpulkan file teks → split train/val
3. Bangun atau muat tokenizer → encode teks ke token IDs
4. Bangun atau muat model GPT
5. Setup optimizer (AdamW) + learning rate schedule
6. Loop training: forward → loss → backward → clip grad → step optimizer
7. Setiap eval_interval langkah: hitung val loss → save best checkpoint
8. Simpan checkpoint final + optimizer state + statistik
Args:
argv: List argumen CLI (atau None untuk pakai sys.argv).
Returns:
Dict statistik training (dipakai oleh pipeline.py untuk manifest.json).
"""
parser = argparse.ArgumentParser(description="Latih model Indigo dari scratch")
# --- Data ---
parser.add_argument("--data", nargs="+", default=["data/sample.txt"],
help="path file/folder teks untuk training (bisa banyak, spasi-separated)")
# --- Output ---
parser.add_argument("--out", default="out",
help="folder output checkpoint (.safetensors + _meta.json + _optimizer.pt)")
# --- Hyperparameter Training ---
parser.add_argument("--steps", type=int, default=2000,
help="jumlah total langkah training (default: 2000)")
parser.add_argument("--batch-size", type=int, default=32,
help="jumlah sampel per batch (default: 32)")
parser.add_argument("--block-size", type=int, default=128,
help="panjang konteks token per sampel (default: 128)")
parser.add_argument("--lr", type=float, default=3e-4,
help="learning rate maksimum (default: 3e-4)")
parser.add_argument("--warmup", type=int, default=100,
help="jumlah langkah warmup linear sebelum cosine decay (default: 100)")
parser.add_argument("--weight-decay", type=float, default=0.1,
help="L2 regularization / weight decay (default: 0.1)")
parser.add_argument("--dropout", type=float, default=0.1,
help="dropout rate (0.0 = nonaktif, default: 0.1)")
# --- Arsitektur Model ---
parser.add_argument("--n-layer", type=int, default=4,
help="jumlah blok transformer (default: 4)")
parser.add_argument("--n-head", type=int, default=4,
help="jumlah head per attention layer (default: 4)")
parser.add_argument("--n-embd", type=int, default=128,
help="dimensi embedding / hidden size (default: 128)")
# --- Tokenizer ---
parser.add_argument("--tokenizer", default="char", choices=["char", "bpe"],
help="jenis tokenizer: 'char' (karakter) atau 'bpe' (subword, default: char)")
parser.add_argument("--vocab-size", type=int, default=512,
help="ukuran vocab untuk BPE (diabaikan jika --tokenizer char, default: 512)")
# --- Evaluasi & Seed ---
parser.add_argument("--eval-interval", type=int, default=200,
help="evaluasi validasi setiap N langkah (0 = tidak ada validasi, default: 200)")
parser.add_argument("--eval-iters", type=int, default=20,
help="jumlah batch untuk estimasi loss validasi (default: 20)")
parser.add_argument("--seed", type=int, default=1337,
help="seed random untuk reproduktibilitas (default: 1337)")
# --- Validasi & Resume ---
parser.add_argument("--val-fraction", type=float, default=0.1,
help="proporsi file untuk validasi (default: 0.1 = 10%%)")
parser.add_argument("--init-from", default=None,
help="path checkpoint untuk melanjutkan training (resume)")
parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda"],
help="device training: auto/cpu/cuda (default: auto)")
args = parser.parse_args(argv)
# --- Setup seed & device ---
torch.manual_seed(args.seed)
if args.device == "auto":
device = "cuda" if torch.cuda.is_available() else "cpu"
else:
device = args.device
os.makedirs(args.out, exist_ok=True)
# --- Kumpulkan & split data ---
# collect_text_files: jika path adalah direktori, cari .txt rekursif
paths = collect_text_files(args.data)
if not paths:
raise SystemExit("tidak ada file teks ditemukan")
# Acak urutan file → split: n_val file untuk validasi, sisanya untuk training
# Split dilakukan per-file (bukan per-karakter), sehingga satu file kecil
# bisa menghabiskan seluruh kuota validasi
files = sorted(paths)
rng = random.Random(args.seed)
rng.shuffle(files)
n_val = max(1, round(len(files) * args.val_fraction)) if len(files) > 1 else 0
print(f"file latih={len(files) - n_val} | file validasi={n_val}")
train_text = "".join(read_clean(p) for p in files[n_val:])
val_text = "".join(read_clean(p) for p in files[:n_val])
all_text = train_text + val_text # dibutuhkan untuk training tokenizer BPE
# --- Setup model & tokenizer ---
init_state = None
init_opt = None
start_step = 0
init_meta = None
config = None
comp_ratio = 1.0
if args.init_from:
# Resume dari checkpoint: muat model, tokenizer, dan optimizer
init_state, init_meta, init_opt = load_init(args.init_from)
config = GPTConfig(**init_meta["config"])
start_step = init_meta.get("step", 0)
print(f"melanjutkan dari {args.init_from} (step {start_step})")
tokenizer = build_tokenizer(init_meta.get("tokenizer") or {"type": "char"}, init_meta["vocab"])
tinfo = init_meta.get("tokenizer") or {"type": "char"}
else:
# Training dari nol: bangun tokenizer baru
if args.tokenizer == "bpe":
from indigo.bpe import BPETokenizer
tokenizer = BPETokenizer.train(all_text, args.vocab_size)
tinfo = tokenizer.state()
n_chars = len(all_text.encode("utf-8"))
comp_ratio = n_chars / max(1, len(tokenizer.encode(all_text)))
print(
f"tokenizer=bpe | vocab={tokenizer.vocab_size} | "
f"kompresi {n_chars:,} karakter -> rasio {comp_ratio:.2f}x"
)
else:
tokenizer = CharTokenizer.from_text(all_text)
tinfo = {"type": "char"}
# Bangun config model baru dari argumen CLI
if config is None:
config = GPTConfig(
vocab_size=tokenizer.vocab_size,
block_size=args.block_size,
n_layer=args.n_layer,
n_head=args.n_head,
n_embd=args.n_embd,
dropout=args.dropout,
)
# Validasi: vocab size model harus cocok dengan tokenizer
if config.vocab_size != tokenizer.vocab_size:
raise SystemExit(
f"vocab tidak cocok: checkpoint={config.vocab_size}, tokenizer={tokenizer.vocab_size}"
)
# --- Encode teks ke token IDs ---
train_data = torch.tensor(tokenizer.encode(train_text), dtype=torch.long)
val_data = torch.tensor(tokenizer.encode(val_text), dtype=torch.long)
if len(train_data) < args.block_size * 2:
raise SystemExit(f"data latih terlalu pendek ({len(train_data)} token), minimal {args.block_size * 2}")
print(
f"tokens latih={len(train_data):,} | tokens validasi={len(val_data):,}"
)
# --- Inisialisasi model ---
model = GPT(config)
if init_state is not None:
missing, unexpected = model.load_state_dict(init_state, strict=False)
if missing or unexpected:
print(f"state_dict: missing={missing} unexpected={unexpected}")
model = model.to(device)
total_steps = start_step + args.steps
print(
f"device={device} | params={model.num_params() / 1e6:.2f}M | "
f"vocab={tokenizer.vocab_size} | total_steps={total_steps}"
)
# --- Setup optimizer: AdamW dengan betas=(0.9, 0.95) ---
optimizer = torch.optim.AdamW(
model.parameters(), lr=args.lr, betas=(0.9, 0.95), weight_decay=args.weight_decay
)
if init_opt is not None:
try:
optimizer.load_state_dict(init_opt)
print("state optimizer dipulihkan")
except Exception as e:
print(f"optimizer state dilewati: {e}")
def save_model(base_path, val_loss):
"""Simpan checkpoint model + metadata ke file .safetensors + _meta.json."""
tensors = {k: v.detach().clone().contiguous() for k, v in model.state_dict().items()}
save_file(tensors, base_path)
save_meta(
base_path,
config.__dict__,
tokenizer.itos if hasattr(tokenizer, "itos") else None,
total_steps,
val_loss,
backend="pytorch",
tokenizer=tinfo,
)
def lr_at(step):
"""Hitung learning rate pada step tertentu.
Schedule:
- Warmup (step < warmup): linear naik dari 0 ke lr maks
- Setelah warmup: cosine decay dari lr maks ke 10% lr maks
- Formula cosine: 0.1*lr + 0.45*lr * (1 + cos(pi * progress))
Jika warmup=0, langsung masuk cosine decay dari step 0.
"""
if step < args.warmup:
return args.lr * (step + 1) / max(1, args.warmup)
progress = (step - args.warmup) / max(1, total_steps - args.warmup)
return 0.1 * args.lr + 0.45 * args.lr * (1 + math.cos(math.pi * progress))
# --- Training loop ---
best_val = float("inf")
last_val = None
model.train()
t0 = time.time()
for step in range(start_step, total_steps):
# Update learning rate sesuai schedule
lr = lr_at(step)
for g in optimizer.param_groups:
g["lr"] = lr
# Forward pass: ambil batch → hitung loss
x, y = get_batch(train_data, config.block_size, args.batch_size, device)
_, loss = model(x, y)
# Backward pass: zero grad → backward → clip grad → step optimizer
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # mencegah gradient explosion
optimizer.step()
# Evaluasi validasi + simpan best checkpoint
if args.eval_interval > 0 and (step % args.eval_interval == 0 or step == total_steps - 1):
if len(val_data) > config.block_size + 1:
val_loss = estimate_loss(model, val_data, args, device)
marker = ""
if val_loss < best_val:
best_val = val_loss
save_model(os.path.join(args.out, "indigo_best.safetensors"), val_loss)
marker = " <- best"
last_val = val_loss
val_str = f"{val_loss:.4f}{marker}"
else:
val_str = "n/a"
print(
f"step {step + 1:5d}/{total_steps} | lr {lr:.2e} | "
f"loss {loss.item():.4f} | val {val_str} | {time.time() - t0:.1f}s"
)
# --- Simpan checkpoint final (bukan best) ---
final_path = os.path.join(args.out, "indigo.safetensors")
save_model(final_path, last_val)
torch.save(optimizer.state_dict(), os.path.join(args.out, "indigo_optimizer.pt"))
print(f"model tersimpan di {final_path} (+_meta.json, indigo_optimizer.pt)")
# --- Ringkasan statistik ---
stats = {
"out": args.out,
"device": device,
"backend": "pytorch",
"tokenizer": tinfo.get("type", "char"),
"vocab_size": tokenizer.vocab_size,
"compression_ratio": round(comp_ratio, 4),
"tokens_train": len(train_data),
"tokens_val": len(val_data),
"files_train": max(0, len(files) - n_val),
"files_val": n_val,
"steps_trained": args.steps,
"total_steps": total_steps,
"best_val": best_val if best_val != float("inf") else None,
"last_val": last_val,
"nats_per_char_best": (
round(best_val / comp_ratio, 4)
if best_val != float("inf") and comp_ratio else None
),
"params_million": round(model.num_params() / 1e6, 4),
"config": config.__dict__,
"args": {k: v for k, v in vars(args).items() if k != "data"},
"elapsed_sec": round(time.time() - t0, 1),
}
print(
f"ringkasan: best_val={stats['best_val']} | "
f"nats/karakter={stats['nats_per_char_best']} | params={stats['params_million']}M"
)
return stats
if __name__ == "__main__":
main()