File size: 6,142 Bytes
5997967
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
#!/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
    _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()