| |
| """Train one TinyLlama variant from a YAML config.""" |
| import argparse |
| import yaml |
| import torch |
|
|
| 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) |
| model = model.to(torch.bfloat16) |
|
|
| n_params = sum(p.numel() for p in model.parameters()) / 1e6 |
| print(f"Model: {n_params:.2f}M params | MLP type: {tiny_config.mlp_type} | Activation: {tiny_config.activation}") |
|
|
| |
| msl = model_cfg.get("max_position_embeddings", 512) |
| train_ds = build_dataset( |
| tokenizer, |
| max_seq_len=msl, |
| split="train", |
| max_samples=None |
| ) |
| eval_ds = build_dataset( |
| tokenizer, |
| max_seq_len=msl, |
| split="validation", |
| max_samples=None |
| ) |
|
|
| |
| 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() |
|
|