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