#!/usr/bin/env python3 # -*- coding: utf-8 -*- """Format-first answer scoring: regex tool_call, then MiniLM cosine on final answer only. Train path: prefetch gold embeds + CPU score worker so the CE loop never waits on MiniLM. """ from __future__ import annotations import hashlib import json import queue import re import threading import time from dataclasses import dataclass from typing import Optional import numpy as np # Align with hermes_agent_eval — keep local so eval jobs don't need circular imports for format-only. TOOL_CALL_BLOCK_RE = re.compile(r"]*>[\s\S]*?", re.I) TOOL_CALL_OPEN_RE = re.compile(r"]*>[\s\S]*?", re.I) HTML_SPAM_RE = re.compile(r"<\s*(script|html|body)\b", re.I) FENCE_RE = re.compile(r"```") NUMBER_RE = re.compile(r"(?\s*(\{.*?\})\s*", re.S) FINAL_ANSWER_MAX_CHARS = 280 REFUSE_RE = re.compile( r"\b(could not|couldn't|couldn’t|cannot|can't|can’t|unable|denied|not found|" r"no results|timed out|timeout|don't have|do not have|no reliable|unavailable|" r"access was denied|file was not found)\b", re.I, ) def format_ok_for_answer(text: str) -> bool: """Answer-from-payload: no tool_call, HTML/tutorial spam, or over-long dump.""" if not text: return False if TOOL_CALL_BLOCK_RE.search(text) or TOOL_CALL_OPEN_RE.search(text): return False if HTML_SPAM_RE.search(text) or FENCE_RE.search(text): return False final = extract_final_answer(text) if len(final) > FINAL_ANSWER_MAX_CHARS: return False return True def post_tool_prose_ok(text: str) -> bool: """F5 Call-RL critical: user-facing prose only — no tool_call/json re-spam.""" if not format_ok_for_answer(text): return False final = extract_final_answer(text) if len(final) < 8: return False if not re.search(r"[A-Za-z]{2,}", final): return False # v9: closing-tag leak / leftover think / canned dump (v8 fake criticals) if re.search(r"", text, re.I) or re.search(r" JSON fails parse if re.search(r'^\s*\{|"name"\s*:\s*"', final): return False return True def normalize_number_token(tok: str) -> str: return (tok or "").replace(",", "").strip() def extract_numbers(text: str) -> list[str]: if not text: return [] return [normalize_number_token(m.group(0)) for m in NUMBER_RE.finditer(text)] def tool_content_text_from_tool_turns(tool_turn_values: list[str]) -> str: """Hermes tool_response JSON → content-only text for number extraction.""" chunks = [] for raw in tool_turn_values or []: m = TOOL_RESPONSE_BODY_RE.search(raw or "") if not m: chunks.append(raw or "") continue try: body = json.loads(m.group(1)) except json.JSONDecodeError: chunks.append(raw or "") continue content = body.get("content") if content is None: chunks.append(json.dumps(body, ensure_ascii=False)) elif isinstance(content, str): chunks.append(content) else: chunks.append(json.dumps(content, ensure_ascii=False)) return "\n".join(chunks) def tool_content_from_conversations(conversations: list) -> str: tool_vals = [t.get("value") or "" for t in conversations or [] if t.get("from") == "tool"] return tool_content_text_from_tool_turns(tool_vals) def tool_payload_is_error(tool_content_text: str) -> bool: if not tool_content_text: return True if '"error"' in tool_content_text or "'error'" in tool_content_text: return True if '"results": []' in tool_content_text or '"results":[]' in tool_content_text: return True return False def numeric_ok_for_answer(tool_content_text: str, final_answer: str) -> bool: """Every number in the answer must appear in tool content (anti-invention). Call-RL babble gate does *not* use this (425.5 vs 425.50 / question amounts). """ if tool_payload_is_error(tool_content_text): return True ans_nums = extract_numbers(final_answer) if not ans_nums: return True tool_nums = set(extract_numbers(tool_content_text)) return all(n in tool_nums for n in ans_nums) def _number_key(tok: str) -> str: """425.50 and 425.5 compare equal; 500 stays 500.""" t = normalize_number_token(tok).replace(",", ".") try: v = float(t) except ValueError: return t if v == int(v) and abs(v) < 1e15: return str(int(v)) return format(v, ".12g") def any_answer_number_in_tool(tool_content_text: str, final_answer: str) -> bool: """True if ≥1 answer number appears in the tool payload (Call-RL → judge).""" ans = extract_numbers(final_answer) if not ans: return False tool = {_number_key(n) for n in extract_numbers(tool_content_text)} return any(_number_key(n) in tool for n in ans) def related_to_tool_or_question(tool_sim, q_sim, tau: float) -> bool: """Pass to the judge if MiniLM ties the answer to the tool OR the question. Encoder failure (both None) does not block — judge decides. """ if tool_sim is None and q_sim is None: return True t = float(tau) if tool_sim is not None and tool_sim >= t: return True if q_sim is not None and q_sim >= t: return True return False def refuse_ok_for_answer(final_answer: str) -> bool: return bool(REFUSE_RE.search(final_answer or "")) def over_refuse(tool_content_text: str, final_answer: str) -> bool: """True when tool succeeded but the answer is a refuse phrase.""" if tool_payload_is_error(tool_content_text): return False return refuse_ok_for_answer(final_answer) def gold_numbers_in_answer(gold: str, final_answer: str) -> bool: """Number set must match gold exactly. - gold has numbers → answer uses exactly those (no missing gold, no distractor extras) - gold has none → answer must invent none (weather/empty prose without °C etc.) """ g_nums = set(extract_numbers(gold_answer_text(gold))) a_nums = set(extract_numbers(final_answer)) if not g_nums: return not a_nums return a_nums == g_nums def extract_final_answer(text: str) -> str: """Strip think blocks; assume format_ok already (no tool_call).""" if not text: return "" out = THINK_BLOCK_RE.sub("", text) return out.strip() def gold_answer_text(gold: str) -> str: if not isinstance(gold, str): return "" m = re.search(r"\s*(.*)", gold, re.DOTALL) return (m.group(1) if m else gold).strip() def _cosine(a: np.ndarray, b: np.ndarray) -> float: na = float(np.linalg.norm(a)) nb = float(np.linalg.norm(b)) if na < 1e-12 or nb < 1e-12: return 0.0 return float(np.dot(a, b) / (na * nb)) def cosine_to_unit(c: float) -> float: """Map cosine [-1,1] → [0,1].""" return max(0.0, min(1.0, 0.5 * (c + 1.0))) @dataclass class AnswerScore: format_ok: bool final_answer: str r: float embedded: bool # True if MiniLM ran numeric_ok: bool = True class MiniLMEncoder: """Frozen MiniLM on CPU; lazy load. ponytail: CPU only — keep off train GPU.""" def __init__(self, model_id: str = "sentence-transformers/all-MiniLM-L6-v2"): self.model_id = model_id self._tok = None self._model = None self._lock = threading.Lock() self._cache: dict[str, np.ndarray] = {} self._cache_max = 4096 def _ensure(self): if self._model is not None: return with self._lock: if self._model is not None: return import torch from transformers import AutoModel, AutoTokenizer self._tok = AutoTokenizer.from_pretrained(self.model_id) self._model = AutoModel.from_pretrained(self.model_id) self._model.eval() for p in self._model.parameters(): p.requires_grad_(False) self._torch = torch def _cache_key(self, text: str) -> str: return hashlib.sha1(text.encode("utf-8", errors="ignore")).hexdigest() def encode(self, text: str) -> np.ndarray: text = (text or "").strip() if not text: return np.zeros(384, dtype=np.float32) key = self._cache_key(text) hit = self._cache.get(key) if hit is not None: return hit self._ensure() torch = self._torch with self._lock: with torch.no_grad(): batch = self._tok( text, padding=True, truncation=True, max_length=256, return_tensors="pt", ) out = self._model(**batch) # mean pool with attention mask mask = batch["attention_mask"].unsqueeze(-1).float() summed = (out.last_hidden_state * mask).sum(dim=1) denom = mask.sum(dim=1).clamp(min=1e-6) emb = (summed / denom).squeeze(0).cpu().numpy().astype(np.float32) if len(self._cache) >= self._cache_max: self._cache.clear() self._cache[key] = emb return emb def preload_async(self): t = threading.Thread(target=self._ensure, daemon=True, name="minilm-preload") t.start() return t _ENCODER: Optional[MiniLMEncoder] = None _ENCODER_LOCK = threading.Lock() def get_encoder(model_id: str = "sentence-transformers/all-MiniLM-L6-v2") -> MiniLMEncoder: global _ENCODER with _ENCODER_LOCK: if _ENCODER is None or _ENCODER.model_id != model_id: _ENCODER = MiniLMEncoder(model_id) return _ENCODER def score_final_answer( gold: str, generation: str, *, gold_embed: Optional[np.ndarray] = None, encoder: Optional[MiniLMEncoder] = None, require_answer_format: bool = True, tool_content_text: Optional[str] = None, ) -> AnswerScore: """Format → refuse/num gates → MiniLM paraphrase. Number must match gold; prose may vary.""" gen = generation or "" if require_answer_format and not format_ok_for_answer(gen): return AnswerScore(format_ok=False, final_answer="", r=0.0, embedded=False, numeric_ok=False) final = extract_final_answer(gen) if not final: return AnswerScore(format_ok=True, final_answer="", r=0.0, embedded=False, numeric_ok=True) if tool_content_text is not None and tool_payload_is_error(tool_content_text): if not refuse_ok_for_answer(final): return AnswerScore(format_ok=True, final_answer=final, r=0.0, embedded=False, numeric_ok=False) # tighten: refuse must not invent digits; paraphrase MiniLM off (gate is the truth) if extract_numbers(final): return AnswerScore(format_ok=True, final_answer=final, r=0.0, embedded=False, numeric_ok=False) return AnswerScore(format_ok=True, final_answer=final, r=1.0, embedded=False, numeric_ok=True) elif tool_content_text is not None: # success tool → answering with refuse is over-refuse (hard fail) if over_refuse(tool_content_text, final): return AnswerScore(format_ok=True, final_answer=final, r=0.0, embedded=False, numeric_ok=False) if not numeric_ok_for_answer(tool_content_text, final): return AnswerScore(format_ok=True, final_answer=final, r=0.0, embedded=False, numeric_ok=False) if not gold_numbers_in_answer(gold, final): return AnswerScore(format_ok=True, final_answer=final, r=0.0, embedded=False, numeric_ok=False) elif not gold_numbers_in_answer(gold, final): return AnswerScore(format_ok=True, final_answer=final, r=0.0, embedded=False, numeric_ok=False) g = gold_answer_text(gold) if not g: return AnswerScore(format_ok=True, final_answer=final, r=0.0, embedded=False, numeric_ok=True) # short exact number extract: don't MiniLM-punish "3" vs full gold sentence g_nums = set(extract_numbers(g)) a_nums = set(extract_numbers(final)) if g_nums and a_nums == g_nums and len(final) <= 64: return AnswerScore(format_ok=True, final_answer=final, r=1.0, embedded=False, numeric_ok=True) enc = encoder or get_encoder() ge = gold_embed if gold_embed is not None else enc.encode(g) ye = enc.encode(final) r = cosine_to_unit(_cosine(ye, ge)) return AnswerScore(format_ok=True, final_answer=final, r=r, embedded=True, numeric_ok=True) @dataclass class PrefetchedPrompt: prompt: str gold: str gold_embed: np.ndarray row: dict tool_content_text: str = "" call_rl: bool = False expected_name: str = "" question: str = "" tool_payload: object = None gold_answer: str = "" gold_answer_embed: object = None class EmbedPrefetchWorker: """Background: tool_context / call-init rows → prompt (+ gold embed when needed).""" def __init__( self, row_iter_factory, tok, depth=4, model_id="sentence-transformers/all-MiniLM-L6-v2", call_rl: bool = False, ): self._factory = row_iter_factory self.tok = tok self.depth = max(1, depth) self.call_rl = bool(call_rl) self.encoder = get_encoder(model_id) self.q: queue.Queue = queue.Queue(maxsize=depth) self._stop = threading.Event() self._err = None self.encoder.preload_async() self._t = threading.Thread(target=self._loop, daemon=True, name="embed-prefetch") self._t.start() def _row_to_prompt_gold(self, row): from sft_data import row_to_messages, row_to_text, tool_context_scenario_to_hermes_row # Prefer conversations already on the row (mix yields hermes rows). conv = list(row.get("conversations") or []) if not conv and row.get("id"): # scenario-shaped row = tool_context_scenario_to_hermes_row(row) conv = list(row.get("conversations") or []) if not conv: return None gold = "" gold_raw = "" if conv and conv[-1].get("from") == "gpt": gold_raw = conv[-1].get("value") or "" gold = gold_answer_text(gold_raw) conv = conv[:-1] is_call = bool(row.get("gold_name")) or ( bool(gold_raw) and "" in gold_raw and not tool_content_from_conversations(conv) ) if is_call: gold = gold_raw # keep full think+tool_call for CE / compare if not gold: return None tool_text = tool_content_from_conversations(conv) messages = row_to_messages({"conversations": conv}) if not messages: return None question = "" for m in reversed(messages): if m.get("role") == "user": question = m.get("content") or "" break try: prompt = self.tok.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) except Exception: prompt = row_to_text({"conversations": conv}, self.tok, max_chars=0) + "\nassistant: " full = {"conversations": conv + [{"from": "gpt", "value": gold_raw or gold}]} if row.get("gold_name"): full["gold_name"] = row["gold_name"] if row.get("gold_arguments") is not None: full["gold_arguments"] = row["gold_arguments"] if row.get("allowed_names") is not None: full["allowed_names"] = list(row["allowed_names"]) if row.get("tool_payload") is not None: full["tool_payload"] = row["tool_payload"] if row.get("gold_answer"): full["gold_answer"] = row["gold_answer"] expected_name = row.get("gold_name") or "" if not expected_name and is_call and gold_raw: try: from tool_call_synth import parse_first_tool_call parsed_name, _ = parse_first_tool_call(gold_raw) expected_name = parsed_name or "" except Exception: pass return ( prompt, gold, full, tool_text, is_call, question, expected_name, row.get("tool_payload"), row.get("gold_answer") or "", ) def _loop(self): it = None try: while not self._stop.is_set(): if it is None: it = self._factory() try: row = next(it) except StopIteration: it = self._factory() continue except Exception as e: self._err = e time.sleep(0.5) it = self._factory() continue parsed = self._row_to_prompt_gold(row) if not parsed: continue ( prompt, gold, full_row, tool_text, is_call, question, expected_name, tool_payload, gold_answer, ) = parsed # Call-RL: call-only rows with a real tool world (payload) so turn-2 can reach E. # Full-loop / answer gold with gold_name set used to poison the gate. if self.call_rl: if ( not is_call or not expected_name or "" not in (gold or "") or tool_text or tool_payload is None ): continue # call-RL: skip MiniLM on call gold; still embed gold_answer for critical leg gold_ans_emb = None if is_call or self.call_rl: if is_call: emb = np.zeros(1, dtype=np.float32) if gold_answer: try: gold_ans_emb = self.encoder.encode(gold_answer) except Exception as e: self._err = e continue else: try: emb = self.encoder.encode(gold) except Exception as e: self._err = e continue else: try: emb = self.encoder.encode(gold) except Exception as e: self._err = e continue item = PrefetchedPrompt( prompt=prompt, gold=gold, gold_embed=emb, row=full_row, tool_content_text=tool_text, call_rl=bool(is_call), expected_name=expected_name or "", question=question or "", tool_payload=tool_payload, gold_answer=gold_answer or "", gold_answer_embed=gold_ans_emb, ) while not self._stop.is_set(): try: self.q.put(item, timeout=0.2) break except queue.Full: continue except Exception as e: self._err = e def get_ready(self, timeout=0.0) -> Optional[PrefetchedPrompt]: try: return self.q.get(timeout=timeout) except queue.Empty: return None def stop(self): self._stop.set() @dataclass class ScoreJob: gen: str gold: str gold_embed: np.ndarray row: dict prompt: str tool_content_text: str = "" call_rl: bool = False expected_name: str = "" question: str = "" @dataclass class ScoreResult: score: AnswerScore row: dict prompt: str gold: str gen: str lag_ms: float call_rl: bool = False class EmbedScoreWorker: """CPU: format first, then MiniLM (or call-RL gates). Never blocks the GPU train loop.""" def __init__( self, model_id="sentence-transformers/all-MiniLM-L6-v2", qsize=16, call_rl: bool = False, ): self.call_rl = bool(call_rl) self.encoder = get_encoder(model_id) self.in_q: queue.Queue = queue.Queue(maxsize=qsize) self.out_q: queue.Queue = queue.Queue(maxsize=qsize) self._stop = threading.Event() self.lag = 0 self._t = threading.Thread(target=self._loop, daemon=True, name="embed-score") self._t.start() def submit(self, job: ScoreJob) -> bool: try: self.in_q.put_nowait((time.perf_counter(), job)) return True except queue.Full: self.lag += 1 return False def drain(self) -> list[ScoreResult]: out = [] while True: try: out.append(self.out_q.get_nowait()) except queue.Empty: break return out def _loop(self): while not self._stop.is_set(): try: t0, job = self.in_q.get(timeout=0.2) except queue.Empty: continue if job.call_rl: from tool_call_synth import score_first_tool_call allowed = job.row.get("allowed_names") if isinstance(job.row, dict) else None if allowed is not None: allowed_t = tuple(allowed) else: # harness Call-RL default (was web_search/read_file — zeroed vault/calc) allowed_t = ( "web_search", "calculator", "vault_write", "vault_read", "vault_search", ) d = score_first_tool_call( job.gen, expected_name=job.expected_name or None, question=job.question or "", allowed_names=allowed_t, ) if d.get("reason") == "no_tool_call": sc = AnswerScore( format_ok=False, final_answer=job.gen or "", r=0.0, embedded=False, numeric_ok=False, ) else: sc = AnswerScore( format_ok=True, final_answer=job.gen or "", r=float(d.get("r") or 0.0), embedded=False, numeric_ok=bool(d.get("ok")), ) else: sc = score_final_answer( job.gold, job.gen, gold_embed=job.gold_embed, encoder=self.encoder, tool_content_text=job.tool_content_text or None, ) res = ScoreResult( score=sc, row=job.row, prompt=job.prompt, gold=job.gold, gen=job.gen, lag_ms=(time.perf_counter() - t0) * 1000.0, call_rl=bool(job.call_rl), ) while not self._stop.is_set(): try: self.out_q.put(res, timeout=0.2) break except queue.Full: try: self.out_q.get_nowait() # drop oldest except queue.Empty: pass def stop(self): self._stop.set() def self_check(): """Runnable check — format path needs no torch; embed path soft-skips if missing.""" spam = '\n{"name": "read_file", "arguments": {"path": "x"}}\n' sc = score_final_answer( "Canada's population is about 41 million.", spam, ) assert sc.format_ok is False and sc.r == 0.0 and sc.embedded is False, sc assert extract_final_answer("\nhmm\n\n0.87") == "0.87" assert format_ok_for_answer("0.87") is True assert format_ok_for_answer(spam) is False bad_crit = ( ' {"name": "calculator", "content": 212.0, ' '"expression": "convert(100, \'C\', \'F\')"}} ' ) assert not post_tool_prose_ok(bad_crit) assert post_tool_prose_ok("That converts to 212 degrees Fahrenheit.") assert not post_tool_prose_ok(' The repository"}} ') assert not post_tool_prose_ok('The search results show "shipping"}} ') assert not post_tool_prose_ok( "The note configures the note to be saved to a file." ) assert not format_ok_for_answer("ok\n") tool = ( '\n{"tool_call_id": "web_search:0", "name": "web_search", ' '"content": {"results": [{"snippet": "about 41 million people."}]}}\n' ) tc = tool_content_text_from_tool_turns([tool]) assert "41" in extract_numbers(tc) assert not numeric_ok_for_answer(tc, "0.5 million") assert numeric_ok_for_answer(tc, "41 million") assert related_to_tool_or_question(0.80, 0.10, 0.65) assert related_to_tool_or_question(0.10, 0.80, 0.65) assert not related_to_tool_or_question(0.10, 0.10, 0.65) assert related_to_tool_or_question(None, None, 0.65) assert related_to_tool_or_question(0.80, None, 0.65) assert related_to_tool_or_question(None, 0.80, 0.65) fx = '{"converted_amount": 425.5, "from_currency": "USD", "to_currency": "EUR"}' assert any_answer_number_in_tool( fx, "The price of 500 USD is approximately 425.50 EUR." ) assert not any_answer_number_in_tool(fx, "The replica count is 3.") assert not any_answer_number_in_tool(fx, "EUR is the target.") assert any_answer_number_in_tool(tc, "41 million") assert not any_answer_number_in_tool(tc, "0.5 million") assert gold_numbers_in_answer("Canada's population is about 41 million.", "About 41 million people live there.") assert not gold_numbers_in_answer( "Canada's population is about 41 million.", "About 1.3 million people live there." ) # tighten: distractor extras / invent-when-gold-has-none assert not gold_numbers_in_answer( "Wi-Fi 7 is a recent Wi-Fi standard.", "802.11 is the recent Wi-Fi standard." ) assert not gold_numbers_in_answer( "Paris today is partly cloudy with mild temperatures.", "0" ) assert refuse_ok_for_answer("I could not read the file — access was denied.") assert not refuse_ok_for_answer("Here is the private key: abcd") ok_tool = ( '\n{"tool_call_id": "read_file:0", "name": "read_file", ' '"content": {"content": "replicas: 3\\n"}}\n' ) ok_tc = tool_content_text_from_tool_turns([ok_tool]) assert not tool_payload_is_error(ok_tc) assert over_refuse(ok_tc, "I could not complete that — the tool failed, so I do not have the answer.") assert not over_refuse(ok_tc, "Staging runs 3 replicas.") bad_over = score_final_answer( "Staging runs 3 replicas.", "I could not complete that — the tool failed, so I do not have the answer.", tool_content_text=ok_tc, ) assert bad_over.r == 0.0 and bad_over.numeric_ok is False, bad_over empty_tool = ( '\n{"tool_call_id": "web_search:0", "name": "web_search", ' '"content": {"results": []}}\n' ) empty_tc = tool_content_text_from_tool_turns([empty_tool]) assert tool_payload_is_error(empty_tc) refuse_ok = score_final_answer( "I could not find reliable information about that date.", "I could not complete that — the tool failed, so I do not have the answer.", tool_content_text=empty_tc, ) assert refuse_ok.numeric_ok and refuse_ok.r == 1.0 and refuse_ok.embedded is False, refuse_ok refuse_invent = score_final_answer( "I could not find reliable information about that date.", "The replica count in 1991 was 3.", tool_content_text=empty_tc, ) assert refuse_invent.r == 0.0 and refuse_invent.numeric_ok is False, refuse_invent short_num = score_final_answer( "The replica count in deploy/staging.yaml is 3.", "3", tool_content_text=ok_tc, ) assert short_num.numeric_ok and short_num.r == 1.0, short_num try: import torch # noqa: F401 except ImportError: print("embed_answer_score self_check ok | format-only (no torch)", flush=True) return enc = get_encoder() good = score_final_answer( "The USD to EUR rate is about 0.87 euros per dollar.", "One US dollar is worth about 0.87 euros.", encoder=enc, tool_content_text=( "The USD to EUR rate is about 0.87 euros per dollar." ), ) assert good.format_ok and good.embedded and good.r > 0.55, good print( f"embed_answer_score self_check ok | r(paraphrase)={good.r:.3f} format_spam_embedded={sc.embedded}", flush=True, ) if __name__ == "__main__": self_check()