Spaces:
Running on Zero
Running on Zero
| #!/usr/bin/env python3 | |
| """ | |
| Frox AI Morph 1.1 — Training CLI | |
| Usage: | |
| python scripts/train.py --family nano --phase pretrain --steps 5000 | |
| python scripts/train.py --family classic --phase sft --steps 20000 | |
| python scripts/train.py --family classic --phase dpo --dataset-name <hf-hub-preference-dataset> | |
| python scripts/train.py --family code --phase sft --data ./my_code_data.jsonl | |
| Phases run in order: pretrain → sft → dpo. Each phase resumes from the | |
| previous phase's saved checkpoint automatically if --resume is set. | |
| --family loads exactly one tier from config/family/<name>.py (nano, | |
| mini, classic, pro, or code) — legacy --family 1.5b/3b/8b names still | |
| work too, routed to the closest matching tier. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import sys | |
| from pathlib import Path | |
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) | |
| import torch | |
| from model.architecture.morph_model import MorphForCausalLM | |
| from tokenizer.morph_tokenizer import build_morph_tokenizer | |
| from training.pipeline.trainer import ( | |
| MorphTrainer, MorphPretrainDataset, MorphSFTDataset, MorphDPODataset, | |
| apply_lora, | |
| ) | |
| from utils.common import ( | |
| set_seed, get_device, describe_device, print_banner, detect_environment, | |
| load_family_config, FAMILY_TIERS, | |
| ) | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Frox AI Morph 1.1 Trainer") | |
| parser.add_argument("--family", choices=list(FAMILY_TIERS), default="nano", | |
| help="Which Morph model-family tier to train (see config/family/*.py)") | |
| parser.add_argument("--phase", choices=["pretrain", "sft", "dpo"], default="sft") | |
| parser.add_argument("--steps", type=int, default=None, | |
| help="Override max steps for this phase") | |
| parser.add_argument("--data", type=str, default=None, | |
| help="Local JSONL/JSON data file (SFT/DPO). Omit to use HF Hub datasets.") | |
| parser.add_argument("--dataset-name", type=str, default=None, | |
| help="HF Hub dataset override for this phase") | |
| parser.add_argument("--resume", action="store_true", | |
| help="Resume from ./frox-morph-1-1-output checkpoint") | |
| parser.add_argument("--from-checkpoint", type=str, default=None, | |
| help="Load base weights from a specific checkpoint dir") | |
| parser.add_argument("--no-lora", action="store_true", help="Full fine-tune instead of LoRA") | |
| parser.add_argument("--seed", type=int, default=1337) | |
| parser.add_argument("--device", type=str, default=None) | |
| args = parser.parse_args() | |
| print_banner() | |
| set_seed(args.seed) | |
| device = get_device(args.device) | |
| print(f"Device: {describe_device(device)}") | |
| print(f"Environment: {detect_environment()}") | |
| config, family_module = load_family_config(args.family) | |
| model_name = getattr(family_module, "MODEL_NAME", args.family.title()) | |
| if args.steps: | |
| if args.phase == "pretrain": | |
| config.training.pretrain_max_steps = args.steps | |
| elif args.phase == "sft": | |
| config.training.sft_max_steps = args.steps | |
| elif args.phase == "dpo": | |
| config.training.dpo_max_steps = args.steps | |
| print(f"\n📐 {model_name} ({args.family})") | |
| print(f" hidden={config.text.hidden_size} layers={config.text.num_hidden_layers} " | |
| f"heads={config.text.num_attention_heads}/{config.text.num_key_value_heads} " | |
| f"context={config.text.max_position_embeddings}") | |
| # Tokenizer | |
| tokenizer = build_morph_tokenizer(save_path="./frox-morph-1-1-output/tokenizer") | |
| # Model | |
| if args.from_checkpoint: | |
| print(f"\n📂 Loading base weights from {args.from_checkpoint}") | |
| model = MorphForCausalLM.from_saved(args.from_checkpoint, device="cpu") | |
| else: | |
| model = MorphForCausalLM(config.text) | |
| params = model.param_count() | |
| print(f" Parameters: {params['total_billions']}B total, " | |
| f"{params['trainable_billions']}B trainable\n") | |
| if not args.no_lora and args.phase != "pretrain": | |
| model = apply_lora( | |
| model, | |
| rank=config.training.lora_rank, | |
| alpha=config.training.lora_alpha, | |
| dropout=config.training.lora_dropout, | |
| target_modules=config.training.lora_target_modules, | |
| ) | |
| # Dataset | |
| if args.phase == "pretrain": | |
| dataset = MorphPretrainDataset( | |
| tokenizer, seq_len=config.training.pretrain_seq_len, | |
| data_mix=config.training.data_mix, | |
| ) | |
| elif args.phase == "sft": | |
| dataset = MorphSFTDataset( | |
| tokenizer, data_path=args.data, | |
| dataset_name=args.dataset_name or "teknium/OpenHermes-2.5", | |
| seq_len=config.training.sft_seq_len, | |
| ) | |
| else: # dpo | |
| if not args.data and not args.dataset_name: | |
| parser.error( | |
| "--phase dpo requires --data (local JSONL) or --dataset-name " | |
| "(any HF Hub preference dataset with chosen/rejected fields)" | |
| ) | |
| dataset = MorphDPODataset( | |
| tokenizer, data_path=args.data, | |
| dataset_name=args.dataset_name, | |
| seq_len=config.training.dpo_seq_len, | |
| ) | |
| trainer = MorphTrainer( | |
| model=model, tokenizer=tokenizer, config=config, | |
| train_dataset=dataset if args.phase != "dpo" else MorphSFTDataset( | |
| tokenizer, max_samples=1 # dummy — DPO uses train_dpo() instead | |
| ), | |
| device=device, | |
| ) | |
| if args.phase == "dpo": | |
| trainer.train_dpo(dataset) | |
| else: | |
| trainer.train(max_steps=args.steps, phase=args.phase) | |
| final_path = f"./frox-morph-1-1-output/{args.family}_{args.phase}_final" | |
| if hasattr(trainer.model, "save_pretrained"): | |
| trainer.model.save_pretrained(final_path) | |
| elif hasattr(trainer.model, "save"): | |
| trainer.model.save(final_path) | |
| print(f"\n✅ {model_name} — {args.phase.upper()} complete. Saved to {final_path}") | |
| if __name__ == "__main__": | |
| main() | |