frox-nano-v2 / src /scripts /train.py
Hritik045678's picture
Upload folder using huggingface_hub
296a506 verified
Raw
History Blame Contribute Delete
6.03 kB
#!/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()