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