Activationiol / train.py
w-ahmad's picture
Upload entire model folder
a3d0c31 verified
Raw
History Blame Contribute Delete
1.85 kB
#!/usr/bin/env python3
"""Train one TinyLlama variant from a YAML config."""
import argparse
import yaml
from transformers import AutoTokenizer, set_seed
from exp import TinyLlamaConfig, TinyLlamaForCausalLM, build_dataset, create_trainer
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True, help="Path to YAML config")
parser.add_argument("--push", action="store_true", help="Push final model to HF Hub")
args = parser.parse_args()
with open(args.config) as f:
cfg = yaml.safe_load(f)
# Explicit seed before any randomness (dataset, model init, dropout, sampler)
seed = cfg.get("training", {}).get("seed", 42)
set_seed(seed)
model_cfg = cfg["model"]
train_cfg = cfg.get("training", {})
# Tokenizer
tok_name = model_cfg.pop("tokenizer_name", "meta-llama/Llama-2-7b-hf")
tokenizer = AutoTokenizer.from_pretrained(tok_name)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
# Model
tiny_config = TinyLlamaConfig(**model_cfg)
model = TinyLlamaForCausalLM(tiny_config)
n_params = sum(p.numel() for p in model.parameters()) / 1e6
print(f"Model: {n_params:.2f}M params | GLU: {tiny_config.glu_activation}")
# Data
msl = model_cfg.get("max_position_embeddings", 512)
train_ds = build_dataset(tokenizer, max_seq_len=msl, split="train")
eval_ds = build_dataset(tokenizer, max_seq_len=msl, split="validation")
# Train
trainer = create_trainer(model, tokenizer, cfg, train_ds, eval_ds)
trainer.train()
# Save & push
out = train_cfg.get("output_dir", "./out")
trainer.save_model(out)
if args.push or train_cfg.get("push_to_hub", False):
trainer.push_to_hub()
print(f"Done. Artifacts in {out}")
if __name__ == "__main__":
main()