Files
nucleic-purpose-classifier/tests/test_train.py
T

89 lines
3.0 KiB
Python

import sys
import unittest
from pathlib import Path
import torch
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
import train
class MetricsTests(unittest.TestCase):
def test_classification_metrics_include_every_label(self):
actual = list(range(8))
predicted = [0, 1, 2, 3, 4, 5, 6, 0]
metrics = train.classification_metrics(actual, predicted)
self.assertEqual(7 / 8, metrics["accuracy"])
self.assertEqual(0.0, metrics["perPurposeRecall"]["writing"])
self.assertEqual(1.0, metrics["perPurposeRecall"]["planning"])
def test_thresholds_preserve_confidence_nesting(self):
probabilities = [0.99, 0.95, 0.85, 0.75, 0.65]
margins = [0.95, 0.85, 0.60, 0.40, 0.20]
correct = [True, True, True, False, False]
thresholds = train.choose_confidence_thresholds(
probabilities,
margins,
correct,
high_precision=1.0,
accepted_precision=0.75,
)
self.assertGreaterEqual(
thresholds["high"]["minimumScore"],
thresholds["medium"]["minimumScore"],
)
self.assertGreater(thresholds["medium"]["validationAcceptedCoverage"], 0)
def test_expected_calibration_error_is_zero_for_perfect_extremes(self):
self.assertEqual(
0.0,
train.expected_calibration_error([1.0, 0.0], [True, False]),
)
class FixedShapeTokenizerTests(unittest.TestCase):
class Tokenizer:
pad_token_id = 0
cls_token_id = 1
sep_token_id = 2
padding_side = "right"
model_input_names = ["input_ids", "attention_mask", "token_type_ids"]
def __call__(self, texts, **_):
return {
"input_ids": [
list(range(10, 10 + int(text.split()[-1]))) for text in texts
]
}
@staticmethod
def num_special_tokens_to_add(pair=False):
return 3 if pair else 2
def test_long_input_keeps_head_and_tail_in_fixed_pair_shape(self):
encoded = train.encode_fixed_shape(
self.Tokenizer(), ["tokens 200"], torch
)
self.assertEqual((1, train.MAX_LENGTH), tuple(encoded["input_ids"].shape))
row = encoded["input_ids"][0].tolist()
self.assertEqual(list(range(10, 10 + train.HEAD_TOKENS)), row[1:64])
self.assertEqual(2, row[64])
self.assertEqual(
list(range(10 + 200 - train.TAIL_TOKENS, 10 + 200)),
row[65:127],
)
self.assertEqual(2, row[127])
self.assertEqual(1, encoded["token_type_ids"][0, 65].item())
def test_short_input_is_padded_as_one_sequence(self):
encoded = train.encode_fixed_shape(self.Tokenizer(), ["tokens 3"], torch)
self.assertEqual([1, 10, 11, 12, 2], encoded["input_ids"][0, :5].tolist())
self.assertEqual(0, encoded["attention_mask"][0, 5].item())
if __name__ == "__main__":
unittest.main()