indigo.tf / train.py
adyoi's picture
Upload train.py with huggingface_hub
382bae9 verified
Raw
History Blame Contribute Delete
10.1 kB
import os
import time
import math
import random
import argparse
import numpy as np
import tensorflow as tf
from safetensors.numpy import load_file, save_file
from indigotf.bpe import BPETokenizer
from indigotf.common import (
build_tokenizer,
collect_text_files,
load_meta,
read_clean,
save_meta,
)
from indigotf.model import build_gpt
from indigotf.tokenizer import CharTokenizer
CONFIG_KEYS = ("vocab_size", "block_size", "n_layer", "n_head", "n_embd", "dropout")
def get_batch(np_data, block_size, batch_size):
ix = np.random.randint(0, len(np_data) - block_size - 1, size=batch_size)
idx = ix[:, None] + np.arange(block_size)
x = np_data[idx]
y = np_data[idx + 1]
return tf.constant(x), tf.constant(y)
@tf.function(reduce_retracing=True)
def train_step(model, optimizer, loss_fn, x, y):
with tf.GradientTape() as tape:
logits = model(x, training=True)
loss = loss_fn(y, logits)
grads = tape.gradient(loss, model.trainable_variables)
optimizer.apply_gradients(zip(grads, model.trainable_variables))
return loss
@tf.function(reduce_retracing=True)
def eval_step(model, loss_fn, x, y):
logits = model(x, training=False)
return loss_fn(y, logits)
def save_weights_tf(model, base_path, config, vocab, step, val_loss, tinfo):
tensors = {v.path: np.asarray(v) for v in model.weights}
save_file(tensors, base_path)
save_meta(base_path, config, vocab, step, val_loss, backend="tensorflow", tokenizer=tinfo)
def main(argv=None):
parser = argparse.ArgumentParser(description="Latih model Indigo-TF (backend TensorFlow/Keras)")
parser.add_argument("--data", nargs="+", default=["data/sample.txt"])
parser.add_argument("--out", default="out")
parser.add_argument("--steps", type=int, default=2000)
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--block-size", type=int, default=128)
parser.add_argument("--n-layer", type=int, default=4)
parser.add_argument("--n-head", type=int, default=4)
parser.add_argument("--n-embd", type=int, default=128)
parser.add_argument("--dropout", type=float, default=0.1)
parser.add_argument("--lr", type=float, default=3e-4)
parser.add_argument("--warmup", type=int, default=100)
parser.add_argument("--weight-decay", type=float, default=0.1)
parser.add_argument("--eval-interval", type=int, default=200)
parser.add_argument("--eval-iters", type=int, default=20)
parser.add_argument("--seed", type=int, default=1337)
parser.add_argument("--init-from", default=None, help="checkpoint safetensors sebelumnya")
parser.add_argument("--tokenizer", default="char", choices=["char", "bpe"])
parser.add_argument("--vocab-size", type=int, default=512)
parser.add_argument("--device", default="auto", choices=["auto", "cpu", "gpu"])
args = parser.parse_args(argv)
if args.device == "cpu":
tf.config.set_visible_devices([], "GPU")
gpus = tf.config.list_physical_devices("GPU")
device_label = f"gpu({len(gpus)})" if gpus and args.device != "cpu" else "cpu"
random.seed(args.seed)
np.random.seed(args.seed)
tf.random.set_seed(args.seed)
os.makedirs(args.out, exist_ok=True)
files = sorted(collect_text_files(args.data))
if not files:
raise SystemExit("tidak ada file teks ditemukan")
rng = random.Random(args.seed)
rng.shuffle(files)
n_val = max(1, round(len(files) * 0.1)) 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
start_step = 0
init_state = None
if args.init_from:
meta = load_meta(args.init_from)
if meta.get("backend") != "tensorflow":
raise SystemExit(f"{args.init_from} bukan checkpoint backend TensorFlow")
config_d = meta["config"]
start_step = meta.get("step", 0)
init_state = {k.replace("/", "_"): v for k, v in load_file(args.init_from).items()}
print(f"melanjutkan dari {args.init_from} (step {start_step})")
tokenizer = build_tokenizer(meta.get("tokenizer") or {"type": "char"}, meta.get("vocab"))
tinfo = meta.get("tokenizer") or {"type": "char"}
else:
if args.tokenizer == "bpe":
tokenizer = BPETokenizer.train(all_text, args.vocab_size)
tinfo = tokenizer.state()
else:
tokenizer = CharTokenizer.from_text(all_text)
tinfo = {"type": "char"}
config_d = {
"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,
"bias": False,
}
if config_d["vocab_size"] != tokenizer.vocab_size:
raise SystemExit(
f"vocab tidak cocok: checkpoint={config_d['vocab_size']}, tokenizer={tokenizer.vocab_size}"
)
train_np = np.array(tokenizer.encode(train_text), dtype=np.int64)
val_np = np.array(tokenizer.encode(val_text), dtype=np.int64)
if len(train_np) < config_d["block_size"] * 2:
raise SystemExit(f"data latih terlalu pendek ({len(train_np)} token)")
print(
f"tokenizer={tinfo['type']} | tokens latih={len(train_np):,} | "
f"tokens validasi={len(val_np):,} | vocab={tokenizer.vocab_size}"
)
total_steps = start_step + args.steps
model = build_gpt(
vocab_size=config_d["vocab_size"],
block_size=config_d["block_size"],
n_layer=config_d["n_layer"],
n_head=config_d["n_head"],
n_embd=config_d["n_embd"],
dropout=config_d["dropout"],
)
if init_state is not None:
by_path = {v.path.replace("/", "_"): v for v in model.weights}
missing = [p for p in by_path if p not in init_state]
if missing:
raise SystemExit(f"bobot tidak cocok dengan checkpoint: {missing[:5]}")
model.set_weights([init_state[v.path.replace("/", "_")] for v in model.weights])
n_params = int(sum(int(np.prod(v.shape)) for v in model.weights))
print(
f"device={device_label} | params={n_params / 1e6:.2f}M | "
f"vocab={config_d['vocab_size']} | total_steps={total_steps}"
)
optimizer = tf.keras.optimizers.AdamW(
learning_rate=args.lr, beta_1=0.9, beta_2=0.95, weight_decay=args.weight_decay, clipnorm=1.0
)
loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)
def lr_at(step):
if step < args.warmup:
return args.lr * (step + 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))
best_val = float("inf")
last_val = None
t0 = time.time()
for step in range(start_step, total_steps):
optimizer.learning_rate.assign(lr_at(step))
x, y = get_batch(train_np, config_d["block_size"], args.batch_size)
loss = train_step(model, optimizer, loss_fn, x, y)
if step % args.eval_interval == 0 or step == total_steps - 1:
if len(val_np) > config_d["block_size"] + 1:
losses = []
for _ in range(args.eval_iters):
vx, vy = get_batch(val_np, config_d["block_size"], args.batch_size)
losses.append(float(eval_step(model, loss_fn, vx, vy)))
val_loss = sum(losses) / len(losses)
marker = ""
if val_loss < best_val:
best_val = val_loss
save_weights_tf(
model,
os.path.join(args.out, "indigo_best.safetensors"),
config_d,
tokenizer.itos if hasattr(tokenizer, "itos") else None,
total_steps,
val_loss,
tinfo,
)
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} | "
f"loss {float(loss):.4f} | val {val_str} | {time.time() - t0:.1f}s"
)
final_path = os.path.join(args.out, "indigo.safetensors")
save_weights_tf(model, final_path, config_d, tokenizer.itos if hasattr(tokenizer, "itos") else None,
total_steps, last_val, tinfo)
print(f"model tersimpan di {final_path} (+_meta.json)")
comp_ratio = 1.0
if tinfo.get("type") == "bpe":
n_chars = len((train_text + val_text).encode("utf-8"))
comp_ratio = n_chars / max(1, len(train_np))
stats = {
"out": args.out,
"device": device_label,
"backend": "tensorflow",
"tokenizer": tinfo.get("type", "char"),
"vocab_size": tokenizer.vocab_size,
"compression_ratio": round(comp_ratio, 4),
"tokens_train": len(train_np),
"tokens_val": len(val_np),
"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(n_params / 1e6, 4),
"config": config_d,
"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()