File size: 17,282 Bytes
9a82835
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6326055
 
c7c5a39
6326055
c7c5a39
 
 
045f351
6326055
c7c5a39
 
 
 
 
 
 
6326055
 
 
 
045f351
9a82835
 
 
 
 
 
 
 
 
 
 
 
045f351
 
 
 
 
9a82835
045f351
 
 
 
 
 
 
c7c5a39
9a82835
045f351
c7c5a39
 
 
 
 
 
 
045f351
 
9a82835
 
f208133
 
 
6326055
9a82835
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6326055
f208133
 
 
 
 
 
 
6326055
 
 
 
 
9a82835
 
 
 
 
 
 
 
 
 
 
 
 
 
6326055
 
 
 
 
 
 
 
 
 
00c2a02
9a82835
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6326055
9a82835
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
00c2a02
6326055
9a82835
6326055
045f351
 
 
 
6326055
 
9a82835
 
045f351
 
 
 
9a82835
 
 
045f351
 
 
 
 
 
 
 
9a82835
045f351
9a82835
045f351
 
 
c7c5a39
 
00c2a02
9a82835
045f351
9a82835
c7c5a39
 
 
045f351
c7c5a39
 
045f351
9a82835
c7c5a39
 
 
 
 
 
00c2a02
c7c5a39
 
00c2a02
c7c5a39
 
 
 
9a82835
c7c5a39
 
 
 
 
 
 
 
 
9a82835
 
045f351
 
c7c5a39
045f351
c7c5a39
9a82835
c7c5a39
 
 
 
 
 
 
 
9a82835
045f351
 
 
 
 
 
 
6326055
 
045f351
6326055
 
9a82835
6326055
 
 
045f351
 
 
 
 
 
 
 
9a82835
045f351
 
c7c5a39
 
 
 
 
 
 
 
 
6326055
 
9a82835
 
 
 
 
 
 
 
 
6326055
9a82835
045f351
6326055
 
9a82835
045f351
 
6326055
 
045f351
9a82835
6326055
 
 
9a82835
 
045f351
6326055
9a82835
 
6326055
 
9a82835
6326055
 
9a82835
 
045f351
6326055
045f351
 
 
 
 
 
 
6326055
 
 
045f351
6326055
 
 
9a82835
045f351
 
 
 
6326055
9a82835
00c2a02
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9a82835
 
00c2a02
 
 
 
 
 
 
 
 
 
 
 
6326055
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
"""
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()