| """ |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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) |
|
|
|
|
| |
| |
| |
|
|
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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 = [ |
| (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) |
|
|
|
|
| |
| |
| |
|
|
| 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: |
| total = 0.0 |
| total = max(-clamp, min(clamp, total)) |
| return {"r_answer": r_ans, "r_faith": r_fai, "r_total": total} |
|
|