timeagent / code /OpenTSLM /grpo /grpo_run.py
roh8exe's picture
Upload folder using huggingface_hub
60b21d3 verified
Raw
History Blame Contribute Delete
6.85 kB
"""
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()