#!/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()