| |
| """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) |
|
|
| |
| seed = cfg.get("training", {}).get("seed", 42) |
| set_seed(seed) |
|
|
| model_cfg = cfg["model"] |
| train_cfg = cfg.get("training", {}) |
|
|
| |
| 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 |
|
|
| |
| 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}") |
|
|
| |
| 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") |
|
|
| |
| trainer = create_trainer(model, tokenizer, cfg, train_ds, eval_ds) |
| trainer.train() |
|
|
| |
| 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() |
|
|