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