Download embed_answer_score.py from ProjectScugnizz/scugnizz-llama-training: direct link, hf CLI and curl.
- Browser
- Download file 30.3 kB
-
https://huggingface.co/ProjectScugnizz/scugnizz-llama-training/resolve/main/embed_answer_score.py
- Command line
-
hf download hf://ProjectScugnizz/scugnizz-llama-training/embed_answer_score.py
-
curl -L -o embed_answer_score.py https://huggingface.co/ProjectScugnizz/scugnizz-llama-training/resolve/main/embed_answer_score.py
30.3 kB
| #!/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"<tool_call\b[^>]*>[\s\S]*?</tool_call>", re.I) | |
| TOOL_CALL_OPEN_RE = re.compile(r"<tool_call\b", re.I) | |
| THINK_BLOCK_RE = re.compile(r"<think\b[^>]*>[\s\S]*?</think>", 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"(?<![A-Za-z])\d+(?:[.,]\d+)?") | |
| TOOL_RESPONSE_BODY_RE = re.compile(r"<tool_response>\s*(\{.*?\})\s*</tool_response>", 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"</tool_call>", text, re.I) or re.search(r"<think\b", final, re.I): | |
| return False | |
| if re.search( | |
| r"The search results show|The note configures|The calculator is a calculator|" | |
| r"The repository|The top hit mentions|partly cloudy|12\s*°\s*C|I saved the note", | |
| final, | |
| re.I, | |
| ): | |
| return False | |
| # malformed {"name": ...} leaks when </tool_call> 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"</think>\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))) | |
| 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) | |
| 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 "<tool_call>" 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 "<tool_call>" 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() | |
| 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 = "" | |
| 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 = '<tool_call>\n{"name": "read_file", "arguments": {"path": "x"}}\n</tool_call>' | |
| 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("<think>\nhmm\n</think>\n0.87") == "0.87" | |
| assert format_ok_for_answer("0.87") is True | |
| assert format_ok_for_answer(spam) is False | |
| bad_crit = ( | |
| '<tool_call> {"name": "calculator", "content": 212.0, ' | |
| '"expression": "convert(100, \'C\', \'F\')"}} </tool_call>' | |
| ) | |
| 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('<think> The repository"}} </tool_call>') | |
| assert not post_tool_prose_ok('The search results show "shipping"}} </tool_call>') | |
| 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<script>alert(1)</script>") | |
| tool = ( | |
| '<tool_response>\n{"tool_call_id": "web_search:0", "name": "web_search", ' | |
| '"content": {"results": [{"snippet": "about 41 million people."}]}}\n</tool_response>' | |
| ) | |
| 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 = ( | |
| '<tool_response>\n{"tool_call_id": "read_file:0", "name": "read_file", ' | |
| '"content": {"content": "replicas: 3\\n"}}\n</tool_response>' | |
| ) | |
| 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 = ( | |
| '<tool_response>\n{"tool_call_id": "web_search:0", "name": "web_search", ' | |
| '"content": {"results": []}}\n</tool_response>' | |
| ) | |
| 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() | |