| """ |
| 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 |
|
|
| |
| from opentslm.time_series_datasets.har_cot.HARCoTQADataset import HARCoTQADataset |
| DATASETS = { |
| "har": dict(cls=HARCoTQADataset, scorer=R.HAR_SCORER, label_key="label"), |
| |
| } |
|
|
|
|
| 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 |
| items = [] |
| for i, sample in enumerate(raw): |
| it = dict(sample) |
| 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) |
| |
| 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) |
| model.train() |
|
|
| |
| 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) |
|
|
| 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() |
|
|