284 lines
8.4 KiB
Python
284 lines
8.4 KiB
Python
#!/usr/bin/env python3
|
||||
|
|
"""
|
|||
|
|
M5.5 — verl Training Harness
|
|||
|
|
|
|||
|
|
Trains the memory controller on exported trajectories using verl.
|
|||
|
|
Supports both trajectory-level and turn-level policy gradient.
|
|||
|
|
|
|||
|
|
Usage:
|
|||
|
|
python verl-training-harness.py \\
|
|||
|
|
--corpus-path corpus/trajectories.jsonl \\
|
|||
|
|
--output-dir ./checkpoints \\
|
|||
|
|
--num-epochs 3 \\
|
|||
|
|
--batch-size 8
|
|||
|
|
|
|||
|
|
Prerequisites:
|
|||
|
|
- verl installed: pip install verl
|
|||
|
|
- transformers, peft, trl installed
|
|||
|
|
- CUDA/GPU available
|
|||
|
|
- Corpus exported from M5.3
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
import argparse
|
|||
|
|
import json
|
|||
|
|
import logging
|
|||
|
|
from pathlib import Path
|
|||
|
|
from typing import List, Dict, Any
|
|||
|
|
|
|||
|
|
import torch
|
|||
|
|
from transformers import AutoTokenizer, AutoModelForCausalLM
|
|||
|
|
from peft import LoraConfig, get_peft_model
|
|||
|
|
from torch.utils.data import Dataset, DataLoader
|
|||
|
|
|
|||
|
|
logger = logging.getLogger(__name__)
|
|||
|
|
|
|||
|
|
|
|||
|
|
class TrajectoryDataset(Dataset):
|
|||
|
|
"""Loads JSONL trajectory format for training."""
|
|||
|
|
|
|||
|
|
def __init__(self, corpus_path: str, tokenizer=None):
|
|||
|
|
self.corpus_path = Path(corpus_path)
|
|||
|
|
self.trajectories = []
|
|||
|
|
self.tokenizer = tokenizer
|
|||
|
|
self._load_trajectories()
|
|||
|
|
|
|||
|
|
def _load_trajectories(self):
|
|||
|
|
"""Load trajectories from JSONL."""
|
|||
|
|
with open(self.corpus_path) as f:
|
|||
|
|
for line in f:
|
|||
|
|
traj = json.loads(line)
|
|||
|
|
self.trajectories.append(traj)
|
|||
|
|
logger.info(f"Loaded {len(self.trajectories)} trajectories")
|
|||
|
|
|
|||
|
|
def __len__(self):
|
|||
|
|
return len(self.trajectories)
|
|||
|
|
|
|||
|
|
def __getitem__(self, idx: int) -> Dict[str, Any]:
|
|||
|
|
"""Return a trajectory with tokenized prompts/responses."""
|
|||
|
|
traj = self.trajectories[idx]
|
|||
|
|
|
|||
|
|
turns = traj.get("turns", [])
|
|||
|
|
rewards = {
|
|||
|
|
"r_update": [],
|
|||
|
|
"r_exit": traj.get("r_exit", 0.0),
|
|||
|
|
"r_format": traj.get("r_format", 1.0),
|
|||
|
|
"r_outcome": traj.get("r_outcome"),
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
# Collect turn rewards
|
|||
|
|
for turn in turns:
|
|||
|
|
rewards["r_update"].append(turn.get("r_update", 0))
|
|||
|
|
|
|||
|
|
return {
|
|||
|
|
"trajectory_id": traj.get("trajectory_id"),
|
|||
|
|
"turns": turns,
|
|||
|
|
"rewards": rewards,
|
|||
|
|
"num_turns": len(turns),
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
|
|||
|
|
def build_model(model_name: str, lora_rank: int = 32):
|
|||
|
|
"""Build base model + LoRA adapter."""
|
|||
|
|
|
|||
|
|
# Load tokenizer
|
|||
|
|
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
|||
|
|
if tokenizer.pad_token is None:
|
|||
|
|
tokenizer.pad_token = tokenizer.eos_token
|
|||
|
|
|
|||
|
|
# Load base model
|
|||
|
|
model = AutoModelForCausalLM.from_pretrained(
|
|||
|
|
model_name,
|
|||
|
|
torch_dtype=torch.float16,
|
|||
|
|
device_map="auto",
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
# Configure LoRA
|
|||
|
|
lora_config = LoraConfig(
|
|||
|
|
r=lora_rank,
|
|||
|
|
lora_alpha=32,
|
|||
|
|
target_modules=["q_proj", "v_proj", "k_proj", "o_proj"],
|
|||
|
|
lora_dropout=0.05,
|
|||
|
|
bias="none",
|
|||
|
|
task_type="CAUSAL_LM",
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
# Apply LoRA
|
|||
|
|
model = get_peft_model(model, lora_config)
|
|||
|
|
|
|||
|
|
return model, tokenizer
|
|||
|
|
|
|||
|
|
|
|||
|
|
class PolicyGradientTrainer:
|
|||
|
|
"""Trains with α-blended trajectory + turn level rewards."""
|
|||
|
|
|
|||
|
|
def __init__(
|
|||
|
|
self,
|
|||
|
|
model,
|
|||
|
|
tokenizer,
|
|||
|
|
learning_rate: float = 5e-5,
|
|||
|
|
trajectory_weight: float = 0.9,
|
|||
|
|
turn_weight: float = 0.1,
|
|||
|
|
):
|
|||
|
|
self.model = model
|
|||
|
|
self.tokenizer = tokenizer
|
|||
|
|
self.optimizer = torch.optim.AdamW(
|
|||
|
|
model.parameters(),
|
|||
|
|
lr=learning_rate,
|
|||
|
|
)
|
|||
|
|
self.trajectory_weight = trajectory_weight
|
|||
|
|
self.turn_weight = turn_weight
|
|||
|
|
|
|||
|
|
def compute_trajectory_loss(
|
|||
|
|
self,
|
|||
|
|
turns: List[Dict],
|
|||
|
|
rewards: Dict,
|
|||
|
|
) -> torch.Tensor:
|
|||
|
|
"""Compute trajectory-level loss (α term)."""
|
|||
|
|
|
|||
|
|
r_exit = torch.tensor(rewards["r_exit"], dtype=torch.float32)
|
|||
|
|
r_format = torch.tensor(rewards["r_format"], dtype=torch.float32)
|
|||
|
|
|
|||
|
|
# Trajectory reward: weighted combination
|
|||
|
|
traj_reward = 0.7 * r_exit + 0.3 * r_format
|
|||
|
|
|
|||
|
|
return traj_reward
|
|||
|
|
|
|||
|
|
def compute_turn_loss(
|
|||
|
|
self,
|
|||
|
|
turns: List[Dict],
|
|||
|
|
rewards: Dict,
|
|||
|
|
) -> torch.Tensor:
|
|||
|
|
"""Compute per-turn loss (1-α term)."""
|
|||
|
|
|
|||
|
|
r_updates = torch.tensor(
|
|||
|
|
rewards["r_update"],
|
|||
|
|
dtype=torch.float32,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
# Average turn reward
|
|||
|
|
turn_loss = -r_updates.mean() # Negative because we minimize loss
|
|||
|
|
|
|||
|
|
return turn_loss
|
|||
|
|
|
|||
|
|
def train_step(self, batch: Dict[str, Any]) -> float:
|
|||
|
|
"""Single training step on a trajectory."""
|
|||
|
|
|
|||
|
|
turns = batch["turns"]
|
|||
|
|
rewards = batch["rewards"]
|
|||
|
|
|
|||
|
|
# Compute loss components
|
|||
|
|
traj_loss = self.compute_trajectory_loss(turns, rewards)
|
|||
|
|
turn_loss = self.compute_turn_loss(turns, rewards)
|
|||
|
|
|
|||
|
|
# Blend losses
|
|||
|
|
loss = (
|
|||
|
|
self.trajectory_weight * traj_loss +
|
|||
|
|
self.turn_weight * turn_loss
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
# Backward pass
|
|||
|
|
self.optimizer.zero_grad()
|
|||
|
|
loss.backward()
|
|||
|
|
torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)
|
|||
|
|
self.optimizer.step()
|
|||
|
|
|
|||
|
|
return loss.item()
|
|||
|
|
|
|||
|
|
|
|||
|
|
def train(
|
|||
|
|
corpus_path: str,
|
|||
|
|
model_name: str = "Qwen/Qwen2.5-3B-Instruct",
|
|||
|
|
output_dir: str = "./checkpoints",
|
|||
|
|
num_epochs: int = 3,
|
|||
|
|
batch_size: int = 8,
|
|||
|
|
lora_rank: int = 32,
|
|||
|
|
learning_rate: float = 5e-5,
|
|||
|
|
):
|
|||
|
|
"""Main training loop."""
|
|||
|
|
|
|||
|
|
# Setup
|
|||
|
|
output_path = Path(output_dir)
|
|||
|
|
output_path.mkdir(parents=True, exist_ok=True)
|
|||
|
|
|
|||
|
|
logger.info(f"Building model: {model_name}")
|
|||
|
|
model, tokenizer = build_model(model_name, lora_rank)
|
|||
|
|
|
|||
|
|
logger.info(f"Loading corpus: {corpus_path}")
|
|||
|
|
dataset = TrajectoryDataset(corpus_path, tokenizer)
|
|||
|
|
dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
|
|||
|
|
|
|||
|
|
logger.info(f"Initializing trainer with lr={learning_rate}")
|
|||
|
|
trainer = PolicyGradientTrainer(
|
|||
|
|
model,
|
|||
|
|
tokenizer,
|
|||
|
|
learning_rate=learning_rate,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
# Training loop
|
|||
|
|
total_loss = 0.0
|
|||
|
|
total_steps = 0
|
|||
|
|
|
|||
|
|
for epoch in range(num_epochs):
|
|||
|
|
logger.info(f"Epoch {epoch+1}/{num_epochs}")
|
|||
|
|
epoch_loss = 0.0
|
|||
|
|
|
|||
|
|
for step, batch in enumerate(dataloader):
|
|||
|
|
loss = trainer.train_step(batch)
|
|||
|
|
epoch_loss += loss
|
|||
|
|
total_loss += loss
|
|||
|
|
total_steps += 1
|
|||
|
|
|
|||
|
|
if step % 10 == 0:
|
|||
|
|
logger.info(f" Step {step}: loss={loss:.4f}")
|
|||
|
|
|
|||
|
|
avg_epoch_loss = epoch_loss / len(dataloader)
|
|||
|
|
logger.info(f"Epoch {epoch+1} avg loss: {avg_epoch_loss:.4f}")
|
|||
|
|
|
|||
|
|
# Save checkpoint
|
|||
|
|
checkpoint_path = output_path / f"memory-v{epoch+1}"
|
|||
|
|
checkpoint_path.mkdir(parents=True, exist_ok=True)
|
|||
|
|
model.save_pretrained(checkpoint_path / "adapter")
|
|||
|
|
tokenizer.save_pretrained(checkpoint_path / "tokenizer")
|
|||
|
|
logger.info(f"Saved checkpoint: {checkpoint_path}")
|
|||
|
|
|
|||
|
|
# Final summary
|
|||
|
|
avg_loss = total_loss / total_steps
|
|||
|
|
logger.info(f"Training complete!")
|
|||
|
|
logger.info(f"Total steps: {total_steps}")
|
|||
|
|
logger.info(f"Average loss: {avg_loss:.4f}")
|
|||
|
|
logger.info(f"Best checkpoint: {output_path / f'memory-v{num_epochs}'}")
|
|||
|
|
|
|||
|
|
return {
|
|||
|
|
"final_loss": avg_loss,
|
|||
|
|
"steps_trained": total_steps,
|
|||
|
|
"epochs": num_epochs,
|
|||
|
|
"checkpoint_path": str(output_path / f"memory-v{num_epochs}"),
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
|
|||
|
|
if __name__ == "__main__":
|
|||
|
|
parser = argparse.ArgumentParser(description="Train memory controller with verl")
|
|||
|
|
parser.add_argument("--corpus-path", required=True, help="Path to JSONL corpus")
|
|||
|
|
parser.add_argument("--output-dir", default="./checkpoints", help="Output directory")
|
|||
|
|
parser.add_argument("--model", default="Qwen/Qwen2.5-3B-Instruct")
|
|||
|
|
parser.add_argument("--num-epochs", type=int, default=3)
|
|||
|
|
parser.add_argument("--batch-size", type=int, default=8)
|
|||
|
|
parser.add_argument("--lora-rank", type=int, default=32)
|
|||
|
|
parser.add_argument("--learning-rate", type=float, default=5e-5)
|
|||
|
|
|
|||
|
|
args = parser.parse_args()
|
|||
|
|
|
|||
|
|
logging.basicConfig(level=logging.INFO)
|
|||
|
|
|
|||
|
|
result = train(
|
|||
|
|
corpus_path=args.corpus_path,
|
|||
|
|
model_name=args.model,
|
|||
|
|
output_dir=args.output_dir,
|
|||
|
|
num_epochs=args.num_epochs,
|
|||
|
|
batch_size=args.batch_size,
|
|||
|
|
lora_rank=args.lora_rank,
|
|||
|
|
learning_rate=args.learning_rate,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
print(json.dumps(result, indent=2))
|