""" 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= 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()