File size: 6,853 Bytes
60b21d3 | 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 158 159 | """
grpo_run.py — standalone RLVR/GRPO runner for OpenTSLM.
This is the script you actually launch. It drives grpo_trainer.GRPOTrainer + reward.py
on the OpenTSLM model — NO edits to curriculum_learning.py. SFT (arms A0/A1) is done
with their curriculum_learning.py unmodified; this script does the GRPO stage (A2/A3),
starting from the SFT checkpoint.
Flow:
1. load OpenTSLMSP(llm_id) + enable LoRA + load the SFT checkpoint
2. build GRPO items = OpenTSLM-formatted sample (prompt + raw time series) augmented
with {gold_label, facts} so the reward can score answer-correctness + faithfulness
3. optimizer over encoder + projector + LoRA (LLM backbone stays frozen)
4. for each batch: loss, stats = grpo.grpo_loss(batch); backward; clip; step; log; save
GPU-blocked to run; needs the HAR faithful CSV + our signal facts present.
source /mnt/nvme0/adinath/timeagent/venv/bin/activate
export CUDA_VISIBLE_DEVICES=<free_gpu> HF_HOME=/mnt/nvme0/adinath/timeagent/hf_cache
python grpo/grpo_run.py --dataset har \
--llm_id meta-llama/Llama-3.2-1B \
--sft_ckpt results/Llama3_2_1B/OpenTSLMSP/stage3_cot/best.pt \
--facts /mnt/nvme0/adinath/timeagent/signal_facts_har_full.json
"""
import os, sys, json, argparse
import torch
from torch.optim import AdamW
from torch.nn.utils import clip_grad_norm_
from torch.utils.data import DataLoader
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src"))
sys.path.insert(0, os.path.dirname(__file__))
from opentslm.model.llm.OpenTSLMSP import OpenTSLMSP
from grpo_trainer import GRPOTrainer, GRPOConfig
import reward as R
# dataset class + faithfulness scorer per dataset
from opentslm.time_series_datasets.har_cot.HARCoTQADataset import HARCoTQADataset
DATASETS = {
"har": dict(cls=HARCoTQADataset, scorer=R.HAR_SCORER, label_key="label"),
# "ecg": dict(cls=ECGQACoTQADataset, scorer=ECG_SCORER, label_key="answer"), # add later
}
def build_grpo_items(ds_obj, facts_by_idx, label_key):
"""Augment each OpenTSLM-formatted sample with gold_label + facts for the reward.
ds_obj.dataset is the list of formatted items (PromptWithAnswer.to_dict(): has the
prompt + time series the model needs). We attach the activity/answer label and the
computed signal facts, joined by position (the loader formats rows in CSV order).
NOTE: confirm the join when HAR lands — if signal_facts entries carry a 'sample_id',
prefer joining on that over positional index.
"""
raw = ds_obj.dataset # list of dicts
items = []
for i, sample in enumerate(raw):
it = dict(sample) # keep prompt + time series fields intact
it["gold_label"] = sample.get(label_key) or facts_by_idx.get(i, {}).get("label", "")
it["facts"] = facts_by_idx.get(i, {}).get("facts", {})
items.append(it)
return items
def make_optimizer(model, lr=1e-6, wd=0.01):
groups = [
{"params": list(model.encoder.parameters()), "lr": lr, "weight_decay": wd},
{"params": list(model.projector.parameters()), "lr": lr, "weight_decay": wd},
]
if getattr(model, "lora_enabled", False):
groups.append({"params": model.get_lora_parameters(), "lr": lr, "weight_decay": wd})
return AdamW(groups)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--dataset", choices=list(DATASETS), default="har")
ap.add_argument("--llm_id", default="meta-llama/Llama-3.2-1B")
ap.add_argument("--sft_ckpt", required=True, help="checkpoint from the SFT (A1) stage")
ap.add_argument("--facts", required=True, help="our signal_facts_*.json for the reward")
ap.add_argument("--out", default="/mnt/nvme0/adinath/timeagent/grpo_out")
ap.add_argument("--batch_size", type=int, default=2)
ap.add_argument("--num_rollouts", type=int, default=8)
ap.add_argument("--max_new_tokens", type=int, default=400)
ap.add_argument("--lr", type=float, default=1e-6)
ap.add_argument("--max_steps", type=int, default=2000)
ap.add_argument("--save_every", type=int, default=200)
# A3 ablation: set --w_faith 0 to train with answer-only reward
ap.add_argument("--w_answer", type=float, default=0.7)
ap.add_argument("--w_faith", type=float, default=0.3)
args = ap.parse_args()
os.makedirs(args.out, exist_ok=True)
spec = DATASETS[args.dataset]
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Loading OpenTSLMSP({args.llm_id}) + LoRA + SFT ckpt {args.sft_ckpt}")
model = OpenTSLMSP(llm_id=args.llm_id, device=device)
model.enable_lora()
model.load_from_file(args.sft_ckpt) # restore SFT encoder/projector/LoRA
model.train()
# facts keyed by sample index (or sample_id) for the reward
with open(args.facts) as f:
facts_list = json.load(f)
facts_by_idx = {i: e for i, e in enumerate(facts_list)}
print("Building GRPO dataset...")
ds_obj = spec["cls"]("train", EOS_TOKEN=model.get_eos_token())
items = build_grpo_items(ds_obj, facts_by_idx, spec["label_key"])
print(f"GRPO items: {len(items)}")
weights = {"answer": args.w_answer, "faith": args.w_faith}
scorer = spec["scorer"]
reward_fn = lambda comp, item: R.dual_reward(
comp, item.get("gold_label", ""), item.get("facts", {}), scorer, weights=weights)
cfg = GRPOConfig(num_rollouts=args.num_rollouts, max_new_tokens=args.max_new_tokens,
w_answer=args.w_answer, w_faith=args.w_faith)
grpo = GRPOTrainer(model, reward_fn=reward_fn, config=cfg)
optimizer = make_optimizer(model, lr=args.lr)
loader = DataLoader(items, batch_size=args.batch_size, shuffle=True,
collate_fn=lambda x: x) # batch = list of item dicts
print("=" * 60)
print(f"GRPO | dataset={args.dataset} G={cfg.num_rollouts} lr={args.lr} "
f"reward w_answer={args.w_answer} w_faith={args.w_faith}")
print("=" * 60)
step = 0
for batch in loader:
optimizer.zero_grad()
loss, stats = grpo.grpo_loss(batch)
loss.backward()
clip_grad_norm_([p for g in optimizer.param_groups for p in g["params"]],
cfg.max_grad_norm)
optimizer.step()
step += 1
if step % 10 == 0:
print(f"step {step:5d} | loss {loss.item():+.4f} | reward {stats['reward_mean']:+.3f} "
f"| ans {stats['answer_reward']:+.3f} | faith {stats['faith_reward']:.3f}")
if step % args.save_every == 0:
ckpt = os.path.join(args.out, f"grpo_step{step}.pt")
model.store_to_file(ckpt)
print("saved", ckpt)
if step >= args.max_steps:
break
final = os.path.join(args.out, "grpo_final.pt")
model.store_to_file(final)
print("Done. Saved", final)
if __name__ == "__main__":
main()
|