Cesium2 / src /train.py
MORPH-AI
feat: extend context window to 8192 with RoPE scaling
9a21993
Raw
History Blame Contribute Delete
10.5 kB
"""
MORPH-AI Training Script v6
Trains the full MorphModel v6 (novel components + LoRA on base).
Memory-efficient training with 4-bit QLoRA, gradient checkpointing, 8-bit optimizer,
MoE expert pruning, and MoD sparsity regularization.
Runs on free Colab T4 or local GPU/CPU.
"""
import gc
import json
import os
import sys
import torch
from pathlib import Path
# ensure both `python src/train.py` (root on path) and
# `python -m src.train` / Colab (src package) imports work
_ROOT = Path(__file__).resolve().parent.parent
if str(_ROOT) not in sys.path:
sys.path.insert(0, str(_ROOT))
if str(_ROOT / "src") not in sys.path:
sys.path.insert(0, str(_ROOT / "src"))
try:
from src.architecture import MorphConfig, MorphModel
except ImportError:
from architecture import MorphConfig, MorphModel
def build_chat_text(ex) -> str:
if isinstance(ex, str):
return ex # already formatted ChatML text
prompt = ex.get("prompt", ex.get("input", ""))
response = ex.get("response", ex.get("output", ""))
skill = ex.get("skill_token", "")
system_prompt = ex.get("system_prompt", "")
if system_prompt:
system_msg = system_prompt
elif skill:
system_msg = f"You are a specialized module. Use the skill marker {skill}."
else:
system_msg = "You are a helpful assistant."
return (
"<|im_start|>system\n"
+ system_msg
+ "<|im_end|>\n"
+ "<|im_start|>user\n"
+ prompt
+ "<|im_end|>\n"
+ "<|im_start|>assistant\n"
+ response
+ "<|im_end|>\n"
)
def load_jsonl(dir_path: str) -> list:
"""Load dataset lines. Lines may be raw ChatML strings OR JSON objects
with prompt/response keys; both are handled downstream by build_chat_text."""
data_dir = Path(dir_path)
files = sorted(data_dir.glob("*.jsonl"))
if not files:
raise FileNotFoundError(f"No .jsonl files found in {dir_path}")
examples = []
for f in files:
with open(f, encoding="utf-8") as fh:
for line in fh:
line = line.strip()
if not line:
continue
try:
examples.append(json.loads(line))
except json.JSONDecodeError:
examples.append(line) # raw text line
print(f"Loaded {len(examples)} examples from {len(files)} files")
return examples
def train(
output_dir: str = "./output/morph-model",
base_model: str = "Qwen/Qwen2.5-1.5B-Instruct",
data_dir: str = "./datasets",
num_epochs: int = 3,
per_device_batch_size: int = 4,
gradient_accumulation_steps: int = 8,
learning_rate: float = 2e-4,
max_seq_len: int = 8192,
use_4bit: bool = True,
train_components: bool = True,
use_8bit_optimizer: bool = True,
prune_experts_every: int = 500,
mod_sparsity_weight: float = 0.01,
):
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
BitsAndBytesConfig,
DataCollatorForLanguageModeling,
Trainer,
TrainingArguments,
)
from datasets import Dataset
if torch.cuda.is_available():
print(f"GPU: {torch.cuda.get_device_name(0)}")
else:
print("No GPU detected - running on CPU (slow). Set use_4bit=False.")
use_4bit = False
use_8bit_optimizer = False
config = MorphConfig(base_model=base_model, max_seq_len=max_seq_len)
model = MorphModel(config)
if use_4bit:
del model.base_model_raw
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
model.base_model_raw = AutoModelForCausalLM.from_pretrained(
base_model,
quantization_config=BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4",
),
device_map="auto",
trust_remote_code=True,
)
print("Base model loaded in 4-bit (nf4)")
from peft import prepare_model_for_kbit_training
if use_4bit:
model.base_model_raw = prepare_model_for_kbit_training(
model.base_model_raw, use_gradient_checkpointing=True
)
print("Base model prepared for k-bit training")
model.apply_lora()
V6_COMPONENTS = (
"coordinator",
"reasoner",
"code_bias",
"scratchpad",
"verifier",
"skill_module",
"depth_module",
"moe_block",
"memory",
"mod",
"multimodal_fusion",
"tool_use",
"document_module",
"video_module",
)
if train_components:
for name, param in model.named_parameters():
if name.startswith(V6_COMPONENTS):
param.requires_grad = True
print(f"Trainable params: {model.get_trainable_params():,}")
else:
print(f"Trainable params (LoRA only): {model.get_trainable_params():,}")
try:
model.base_model.gradient_checkpointing_enable()
except Exception as e:
print(f"Note: gradient checkpointing skipped ({e})")
tokenizer = AutoTokenizer.from_pretrained(base_model, trust_remote_code=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
examples = load_jsonl(data_dir)
texts = [build_chat_text(ex) for ex in examples]
def tokenize(text: str):
return tokenizer(
text,
truncation=True,
max_length=max_seq_len,
padding=False,
return_tensors=None,
)
raw = [tokenize(t) for t in texts]
token_ids = [torch.tensor(r["input_ids"]).unsqueeze(0) for r in raw]
try:
from src.architecture import build_code_features
except ImportError:
from architecture import build_code_features
code_feats = [build_code_features(tokenizer, ids) for ids in token_ids]
records = []
for r, cf in zip(raw, code_feats):
records.append(
{
"input_ids": r["input_ids"],
"attention_mask": r["attention_mask"],
"code_feat": cf.squeeze(0),
}
)
dataset = Dataset.from_list(records)
def collate(batch):
import torch.nn.functional as Fn
from transformers import DataCollatorForLanguageModeling
max_len = max(len(b["input_ids"]) for b in batch)
input_ids = torch.full((len(batch), max_len), tokenizer.pad_token_id, dtype=torch.long)
attention_mask = torch.zeros((len(batch), max_len), dtype=torch.long)
for i, b in enumerate(batch):
n = len(b["input_ids"])
input_ids[i, :n] = torch.tensor(b["input_ids"])
attention_mask[i, :n] = torch.tensor(b["attention_mask"])
feat0 = torch.tensor(batch[0]["code_feat"])
feat_dim = feat0.shape[-1]
code_feat = torch.zeros((len(batch), max_len, feat_dim), dtype=torch.float32)
for i, b in enumerate(batch):
cf = torch.tensor(b["code_feat"])
n = cf.shape[0]
code_feat[i, :n] = cf
lm_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False)
lm_batch = lm_collator([{"input_ids": b["input_ids"]} for b in batch])
return {
"input_ids": input_ids,
"attention_mask": attention_mask,
"labels": lm_batch["labels"],
"code_feat": code_feat,
}
optim = "paged_adamw_8bit" if use_8bit_optimizer else "adamw_torch"
training_args = TrainingArguments(
output_dir="./checkpoints",
per_device_train_batch_size=per_device_batch_size,
gradient_accumulation_steps=gradient_accumulation_steps,
num_train_epochs=num_epochs,
learning_rate=learning_rate,
warmup_steps=50,
logging_steps=10,
save_steps=500,
save_total_limit=2,
report_to=[],
optim=optim,
gradient_checkpointing=True,
bf16=torch.cuda.is_available(),
)
class MorphTrainer(Trainer):
def on_step_begin(self, args, state, control, **kwargs):
if state.global_step > 0 and prune_experts_every > 0:
if state.global_step % prune_experts_every == 0:
model.prune_experts()
trainer = MorphTrainer(
model=model,
args=training_args,
train_dataset=dataset,
data_collator=collate,
)
trainer.train()
out = Path(output_dir)
out.mkdir(parents=True, exist_ok=True)
model.save_checkpoint(str(out))
tokenizer.save_pretrained(str(out))
print(f"Model saved to {out}")
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="MORPH-AI training")
parser.add_argument("--output", default="./output/morph-model")
parser.add_argument("--base", default="Qwen/Qwen2.5-1.5B-Instruct")
parser.add_argument("--data", default="./datasets")
parser.add_argument("--epochs", type=int, default=3)
parser.add_argument("--batch", type=int, default=4)
parser.add_argument("--grad-accum", type=int, default=8)
parser.add_argument("--lr", type=float, default=2e-4)
parser.add_argument("--max-len", type=int, default=8192)
parser.add_argument("--no-4bit", action="store_true", help="Disable 4-bit quantization")
parser.add_argument(
"--no-components",
action="store_true",
help="Train LoRA only (freeze novel components)",
)
parser.add_argument("--no-8bit-optim", action="store_true", help="Disable 8-bit optimizer")
parser.add_argument("--prune-every", type=int, default=500, help="Prune MoE experts every N steps (0=off)")
parser.add_argument("--mod-sparsity", type=float, default=0.01, help="MoD sparsity loss weight")
args = parser.parse_args()
train(
output_dir=args.output,
base_model=args.base,
data_dir=args.data,
num_epochs=args.epochs,
per_device_batch_size=args.batch,
gradient_accumulation_steps=args.grad_accum,
learning_rate=args.lr,
max_seq_len=args.max_len,
use_4bit=not args.no_4bit,
train_components=not args.no_components,
use_8bit_optimizer=not args.no_8bit_optim,
prune_experts_every=args.prune_every,
mod_sparsity_weight=args.mod_sparsity,
)