File size: 3,981 Bytes
66ee87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
"""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()