Batch-effects-leaderboard / tests /test_proba_alignment.py
spell0's picture
Fix BERNN probability and batch mapping
173a895
Raw
History Blame Contribute Delete
1.42 kB
import unittest
import numpy as np
from src.code_challenge import _aligned_proba_frame
class TestProbabilityAlignment(unittest.TestCase):
def test_rejects_encoded_classes_that_do_not_match_decoded_labels(self):
frame = _aligned_proba_frame(
np.array([[0.7, 0.2, 0.1], [0.1, 0.6, 0.3]]),
classes=[0, 1, 2],
labels=["AA", "Bio", "FA"],
n_rows=2,
)
self.assertIsNone(frame)
def test_aligns_decoded_classes_to_label_order(self):
frame = _aligned_proba_frame(
np.array([[0.2, 0.7, 0.1], [0.6, 0.1, 0.3]]),
classes=["Bio", "AA", "FA"],
labels=["AA", "Bio", "FA"],
n_rows=2,
)
self.assertIsNotNone(frame)
self.assertEqual(list(frame.columns), ["AA", "Bio", "FA"])
self.assertEqual(frame["AA"].tolist(), [0.7, 0.1])
self.assertEqual(frame["Bio"].tolist(), [0.2, 0.6])
def test_accepts_positional_probabilities_when_classes_are_missing(self):
frame = _aligned_proba_frame(
np.array([[0.2, 0.7, 0.1]]),
classes=None,
labels=["AA", "Bio", "FA"],
n_rows=1,
)
self.assertIsNotNone(frame)
self.assertEqual(list(frame.columns), ["AA", "Bio", "FA"])
self.assertEqual(frame.iloc[0].tolist(), [0.2, 0.7, 0.1])
if __name__ == "__main__":
unittest.main()