DecisionLab / tests /test_arthur_io.py
RealFalconsAI's picture
Upload 39 files
66ee87e verified
Raw History Blame Contribute Delete
3.98 kB
"""Arthur's input layout and output calibration (app/arthur_io.py): pure Python, copied from arthur_v0_8_0.ipynb."""
import json
import unittest
import math
from app.arthur_io import (CLS, LAYOUT, OPT, PAD, SEP, QTYPES, assemble, bucket, normalize_question, probabilities,
softmax, state_text)
class NormalizeQuestionTest(unittest.TestCase):
def test_noul_is_yes_no_with_true_first(self):
self.assertEqual(normalize_question({"type": "noul", "instructions": "Angry?"}),
("noul", "Angry?", [True, False], ["Yes", "No"]))
def test_noul_labels_can_be_renamed(self):
q = {"type": "noul", "question": "Q", "labels": {"true": "Refund", "false": "Keep"}}
self.assertEqual(normalize_question(q)[3], ["Refund", "Keep"])
def test_choice_dict_becomes_key_colon_description(self):
q = {"type": "choice", "instructions": "Team?", "criteria": {"bug": "Something is broken", "sales": ""}}
self.assertEqual(normalize_question(q), ("choice", "Team?", ["bug", "sales"], ["bug: Something is broken", "sales"]))
def test_choice_list(self):
self.assertEqual(normalize_question({"instructions": "x", "options": ["a", "b"]})[2:], (["a", "b"], ["a", "b"]))
def test_score_levels_are_indices(self):
q = {"type": "score", "instructions": "x", "criteria": ["Low", "High"]}
self.assertEqual(normalize_question(q), ("score", "x", [0, 1], ["Low", "High"]))
class StateTextTest(unittest.TestCase):
def test_dict_state_is_json_keeping_unicode(self):
self.assertEqual(state_text({"msg": "café"}), json.dumps({"msg": "café"}, ensure_ascii=False))
def test_text_and_empty(self):
self.assertEqual((state_text("hi"), state_text(None)), ("hi", ""))
class AssembleTest(unittest.TestCase):
def test_layout_matches_the_notebook(self):
ids, windows = assemble({"question": "Q?", "options": ["a", "b"], "state": "s"})
b = lambda ch: ord(ch) + 1
self.assertEqual(ids, [CLS, b("Q"), b("?"), SEP, OPT, b("a"), PAD, PAD, OPT, b("b"), SEP, b("s"), SEP])
self.assertEqual(windows, [1, 2])
def test_every_option_window_starts_on_an_opt_token(self):
ids, windows = assemble({"question": "Which?", "options": ["alpha", "b", "gamma delta"], "state": "x" * 50})
self.assertEqual([ids[w * 4] for w in windows], [OPT, OPT, OPT])
def test_long_state_is_cut_to_the_token_budget(self):
ids, _ = assemble({"question": "q", "options": ["a", "b"], "state": "x" * 10000})
self.assertEqual(len(ids), 4 * LAYOUT["max_len"])
def test_many_options_get_the_long_budget(self):
ids, _ = assemble({"question": "q", "options": [f"o{i}" for i in range(30)], "state": "x" * 20000})
self.assertEqual(len(ids), 4 * LAYOUT["long_max_len"])
def test_utf8_bytes_are_shifted_by_one(self):
ids, _ = assemble({"question": "é", "options": ["a", "b"], "state": ""})
self.assertEqual(ids[1:3], [0xC3 + 1, 0xA9 + 1])
class CalibrationTest(unittest.TestCase):
TEMPS = [[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0], [9.0, 10.0, 11.0, 12.0]]
def test_buckets(self):
self.assertEqual([bucket(k) for k in (2, 3, 5, 6, 12, 13, 40)], [0, 1, 1, 2, 2, 3, 3])
def test_softmax_with_temperature(self):
p = softmax([2.0, 0.0], 2.0)
self.assertAlmostEqual(p[0], math.exp(1) / (math.exp(1) + 1))
self.assertAlmostEqual(sum(p), 1.0)
def test_temperature_is_chosen_by_type_and_option_count(self):
z = [1.0, 0.0, 0.0]
self.assertEqual(list(probabilities(z, "score", self.TEMPS)), list(softmax(z, 10.0))) # score row, bucket 1
self.assertEqual(list(probabilities([1.0, 0.0], "noul", self.TEMPS)), list(softmax([1.0, 0.0], 5.0)))
def test_qtype_indices_match_the_training_notebook(self):
self.assertEqual(QTYPES, {"choice": 0, "noul": 1, "score": 2})
if __name__ == "__main__":
unittest.main()