Files
poimen-memory/verl-training-harness.py
T

284 lines
8.4 KiB
Python
Raw Normal View History

#!/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))