timeagent / code /OpenTSLM /grpo /reward.py
roh8exe's picture
Upload folder using huggingface_hub
60b21d3 verified
Raw
History Blame Contribute Delete
6.43 kB
"""
Reward for RLVR/GRPO on OpenTSLM: answer correctness + signal faithfulness.
r_total = w_answer * r_answer + w_faith * r_faith (default 0.7 / 0.3)
r_answer {-1, 0, +1} gold label vs the rationale's "Answer:" (exact / end-of-phrase
token match; -1 wrong, 0 if no answer stated)
r_faith [0, 1] fraction of numeric claims that MATCH ground-truth facts —
identical to our Stage-3 verifier metric (agent3_*.py), so the
RL reward and the reported faithfulness are the same yardstick.
Faithfulness is modality-pluggable via a FaithfulnessScorer built from claim_patterns.
A HAR scorer is provided; ECG/Sleep/WESAD scorers plug in with the same shape (lift the
pattern sets from codebase/agent3_{ecg,wesad}.py).
"""
import re
# --------------------------------------------------------------------------------------
# Completion text extraction (TRL may hand back str | [{"content": ...}] | {"content"})
# --------------------------------------------------------------------------------------
def completion_text(comp) -> str:
if comp is None:
return ""
if isinstance(comp, str):
return comp
if isinstance(comp, list):
return comp[0].get("content", "") if comp and isinstance(comp[0], dict) else str(comp)
if isinstance(comp, dict):
return comp.get("content", "")
return str(comp)
# --------------------------------------------------------------------------------------
# R_answer — answer correctness
# --------------------------------------------------------------------------------------
def _norm_tokens(s):
return re.findall(r"[a-z0-9]+", str(s).lower())
def answer_reward(text: str, gold_label: str) -> float:
"""+1 if the stated answer matches gold, -1 if a different answer is stated, 0 if
none found. Exact / end-of-phrase token match (NOT naive substring, so short
answers like 'no' aren't matched inside unrelated words)."""
m = list(re.finditer(r"Answer:\s*(.+?)\s*$", text, re.IGNORECASE | re.MULTILINE))
if not m:
return 0.0
pred = _norm_tokens(m[-1].group(1))
gold = _norm_tokens(gold_label)
if not pred or not gold:
return 0.0
if pred == gold or (len(pred) >= len(gold) and pred[-len(gold):] == gold):
return 1.0
return -1.0
# --------------------------------------------------------------------------------------
# Faithfulness scorer (modality-pluggable)
# --------------------------------------------------------------------------------------
def _parse_number(s):
try:
return float(str(s).replace(",", "").replace(" ", ""))
except Exception:
return None
class FaithfulnessScorer:
"""claim_patterns: list of (regex_with_one_capture_group, [[fact_key, ...]]).
Keys ending in '*' are prefix-expanded over the facts dict. Mirrors
compute_faithfulness in codebase/agent3_*.py."""
def __init__(self, claim_patterns, tol: float = 0.15):
self.claim_patterns = [(re.compile(p, re.IGNORECASE), keys) for p, keys in claim_patterns]
self.tol = tol
def extract_claims(self, text):
claims = []
for pattern, key_groups in self.claim_patterns:
for m in pattern.finditer(text):
for i, val_str in enumerate(m.groups()):
if val_str is None:
continue
val = _parse_number(val_str)
if val is None:
continue
keys = key_groups[i] if i < len(key_groups) else key_groups[-1]
claims.append((val, keys if isinstance(keys, list) else [keys]))
return claims
@staticmethod
def _pool(facts, fact_keys):
pool = []
for key in fact_keys:
if key.endswith("*"):
prefix = key[:-1]
for fk, fv in facts.items():
if fk.startswith(prefix) and isinstance(fv, (int, float)) and fv is not None:
pool.append(float(fv))
elif key in facts and isinstance(facts[key], (int, float)) and facts[key] is not None:
pool.append(float(facts[key]))
return pool
def _matches(self, val, pool):
for f in pool:
if f == 0:
if abs(val) < 1:
return True
elif abs(val - f) / abs(f) <= self.tol:
return True
return False
def score(self, text, facts) -> float:
claims = self.extract_claims(text)
if not claims:
return 0.0
verified = sum(1 for val, keys in claims if self._matches(val, self._pool(facts, keys)))
return verified / len(claims)
# HAR claim patterns (Hz, m/s^2, sec, peaks) — matches the HAR Stage-3 verifier.
HAR_CLAIM_PATTERNS = [
(r"([\d\.]+)\s*Hz", [["dominant_freq_x", "dominant_freq_y", "dominant_freq_z", "stride_freq_hz"]]),
(r"([\d\.]+)\s*m/s", [["mean_x", "mean_y", "mean_z", "std_x", "std_y", "std_z",
"smv_mean", "smv_std", "smv_max", "dynamic_acc_mean",
"dynamic_acc_max"]]),
(r"([\d\.]+)\s*(?:sec|seconds|s\b)", [["stride_interval_sec"]]),
(r"([\d]+)\s*(?:strides|peaks|steps)", [["n_strides", "n_peaks"]]),
]
HAR_SCORER = FaithfulnessScorer(HAR_CLAIM_PATTERNS)
# --------------------------------------------------------------------------------------
# Composite reward (answer + faithfulness)
# --------------------------------------------------------------------------------------
WEIGHTS = {"answer": 0.7, "faith": 0.3}
def dual_reward(comp, gold_label, facts, scorer, weights=None, clamp=5.0) -> dict:
"""r = w_answer*r_answer + w_faith*r_faith. `comp` may be str | message-list | dict.
Returns components + NaN-safe, clamped 'r_total'."""
w = weights or WEIGHTS
text = completion_text(comp)
try:
r_ans = answer_reward(text, gold_label)
r_fai = scorer.score(text, facts) if facts else 0.0
total = w["answer"] * r_ans + w["faith"] * r_fai
except Exception as e:
print("reward error:", e)
r_ans = r_fai = 0.0
total = 0.0
if total != total: # NaN guard
total = 0.0
total = max(-clamp, min(clamp, total))
return {"r_answer": r_ans, "r_faith": r_fai, "r_total": total}