138 lines
4.8 KiB
Python
138 lines
4.8 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_distillation_loss_is_zero_for_matching_logits_and_backpropagates(self):
|
|
teacher = torch.tensor([[2.0, 0.0, -1.0]])
|
|
student = teacher.clone().requires_grad_(True)
|
|
loss = train.knowledge_distillation_loss(
|
|
torch,
|
|
student,
|
|
teacher,
|
|
temperature=2.0,
|
|
).mean()
|
|
self.assertAlmostEqual(0.0, loss.item(), places=6)
|
|
loss.backward()
|
|
self.assertIsNotNone(student.grad)
|
|
|
|
def test_qat_replacements_keep_checkpoint_keys_and_gradients(self):
|
|
model = torch.nn.Sequential(
|
|
torch.nn.Embedding(16, 8),
|
|
torch.nn.Flatten(),
|
|
torch.nn.Linear(16, 3),
|
|
)
|
|
keys = set(model.state_dict())
|
|
counts = train.enable_quantization_aware_training(torch, model)
|
|
self.assertEqual({"linear": 1, "embedding": 1}, counts)
|
|
self.assertEqual(keys, set(model.state_dict()))
|
|
output = model(torch.tensor([[1, 2]], dtype=torch.long))
|
|
output.sum().backward()
|
|
self.assertIsNotNone(model[0].weight.grad)
|
|
self.assertIsNotNone(model[2].weight.grad)
|
|
|
|
def test_qat_affine_ranges_include_zero(self):
|
|
model = torch.nn.Sequential(torch.nn.Linear(2, 2, bias=False))
|
|
with torch.no_grad():
|
|
model[0].weight.copy_(torch.eye(2))
|
|
train.enable_quantization_aware_training(torch, model)
|
|
output = model(torch.tensor([[1.0, 2.0]]))
|
|
self.assertTrue(
|
|
torch.allclose(output, torch.tensor([[1.0, 2.0]]), atol=0.02),
|
|
output,
|
|
)
|
|
|
|
def test_boundary_training_weight_is_opt_in(self):
|
|
self.assertEqual(
|
|
2.0,
|
|
train.training_weight({"slice": "boundary"}, boundary_weight=2.0),
|
|
)
|
|
self.assertEqual(
|
|
1.0,
|
|
train.training_weight({"slice": "core"}, boundary_weight=2.0),
|
|
)
|
|
|
|
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()
|