| |
| """ |
| 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}") |
|
|
| |
| _tok_path = "./frox-morph-1-1-output/tokenizer" |
| tokenizer = build_morph_tokenizer( |
| tokenizer_path=_tok_path if Path(_tok_path).exists() else None, |
| save_path=_tok_path, |
| ) |
|
|
| |
| 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, |
| ) |
|
|
| |
| 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: |
| 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 |
| ), |
| 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() |
|
|