File size: 4,003 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
"""
Inject our Signal-Faithful rationales into OpenTSLM's expected CSV schema, so their
training pipeline trains on OUR rationales with everything else held fixed.

OpenTSLM's HAR loader (har_cot_loader.py) expects, at
    OpenTSLM/src/data/har_cot/har_cot_{train,val,test}_cot.csv
columns:  x_axis, y_axis, z_axis, label, rationale
where the three *_axis columns are stringified arrays and `rationale` is the CoT the
model is trained to produce (OpenTSLMSP.compute_loss teacher-forces on it).

This script keeps the original time series + label and swaps in our faithful rationale,
matched row-for-row (the splits are produced in the same order). Run the resulting
training twice — once with the original CSVs, once with ours — to get the downstream
faithful-vs-baseline comparison (paper Section "Downstream HAR Classification").

Usage:
    python grpo/inject_faithful_data.py \
        --orig OpenTSLM/src/data/har_cot/har_cot_train_cot.csv \
        --ours /path/to/faithful_har_cot_train.csv \
        --out  OpenTSLM/src/data/har_cot/har_cot_train_cot.csv \
        --ours-rationale-col our_rationale     # or 'rationale'
        [--backup]   # save the original alongside as *.orig.csv first

Repeat for val/test. For ECG, our faithful CSV uses (sample_id, our_rationale, ...) and
ECGQACoTQADataset has its own loader; adapt --ours-rationale-col and the merge key with
--key sample_id accordingly.
"""

import argparse
import os
import shutil

import pandas as pd


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--orig", required=True, help="OpenTSLM's original CSV (time series + label)")
    ap.add_argument("--ours", required=True, help="our faithful CSV containing the rationale")
    ap.add_argument("--out", required=True, help="destination CSV (OpenTSLM schema)")
    ap.add_argument("--ours-rationale-col", default="our_rationale",
                    help="rationale column name in --ours (try 'our_rationale' or 'rationale')")
    ap.add_argument("--key", default=None,
                    help="optional join key present in both CSVs (e.g. sample_id). "
                         "If omitted, rows are matched by position.")
    ap.add_argument("--backup", action="store_true", help="copy --orig to *.orig.csv first")
    args = ap.parse_args()

    orig = pd.read_csv(args.orig)
    ours = pd.read_csv(args.ours)
    print(f"orig: {len(orig)} rows, cols={list(orig.columns)}")
    print(f"ours: {len(ours)} rows, cols={list(ours.columns)}")

    if args.ours_rationale_col not in ours.columns:
        raise SystemExit(
            f"--ours-rationale-col '{args.ours_rationale_col}' not in {list(ours.columns)}"
        )

    merged = orig.copy()
    if args.key:
        if args.key not in orig.columns or args.key not in ours.columns:
            raise SystemExit(f"--key '{args.key}' must be in both CSVs")
        rmap = dict(zip(ours[args.key], ours[args.ours_rationale_col]))
        before = len(merged)
        merged = merged[merged[args.key].isin(rmap)].copy()
        merged["rationale"] = merged[args.key].map(rmap)
        print(f"joined on {args.key}: kept {len(merged)}/{before} rows")
    else:
        if len(orig) != len(ours):
            raise SystemExit(
                f"row count mismatch ({len(orig)} vs {len(ours)}); use --key for a safe join"
            )
        merged["rationale"] = ours[args.ours_rationale_col].values

    n_empty = merged["rationale"].isna().sum()
    if n_empty:
        print(f"WARNING: {n_empty} rows have no rationale after merge")

    if args.backup and os.path.abspath(args.orig) == os.path.abspath(args.out):
        bak = args.orig.replace(".csv", ".orig.csv")
        if not os.path.exists(bak):
            shutil.copy(args.orig, bak)
            print(f"backed up original -> {bak}")

    os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
    merged.to_csv(args.out, index=False)
    print(f"wrote {len(merged)} rows -> {args.out}")


if __name__ == "__main__":
    main()