carbon-a-database-explorer / tests /test_catalog.py
lvwerra's picture
lvwerra HF Staff
Show the distribution of probabilities per column, not their mean (#11)
0d06695
Raw History Blame Contribute Delete
7.22 kB
import hashlib
import json
from pathlib import Path
import tempfile
import unittest
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
from catalog import Catalog, ROOT
class SampleTests(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.catalog = Catalog()
def test_sample_integrity_and_budget(self):
manifest = self.catalog.manifest
self.assertEqual(hashlib.sha256(self.catalog.path.read_bytes()).hexdigest(), manifest['sha256'])
self.assertEqual(len(self.catalog.records), manifest['rows'])
self.assertLessEqual(manifest['rows'], manifest['max_rows'])
self.assertLessEqual(manifest['bases'], manifest['max_bp'])
for i, record in enumerate(self.catalog.records):
row = self.catalog.segment_table(i).to_pylist()[0]
length = record['segment_end_bp'] - record['segment_start_bp']
for key in ('pred_prob_positive_strand_cds', 'pred_prob_negative_strand_cds'):
values = np.asarray(row[key], dtype=float)
self.assertEqual(len(values), length)
self.assertTrue(np.all(np.isfinite(values)))
self.assertTrue(np.all((values >= 0) & (values <= 1)))
def test_real_accession_lookup(self):
first = self.catalog.records[0]
self.assertIn(0, self.catalog.lookup(first['source_key']))
self.assertIn(0, self.catalog.lookup(' ' + first['record_name'].lower() + ' '))
self.assertIn(0, self.catalog.lookup(first['record_name'].rsplit('.', 1)[0]))
expected = [i for i, r in enumerate(self.catalog.records) if r['assembly_accession'] == first['assembly_accession']]
self.assertCountEqual(self.catalog.lookup(first['assembly_accession']), expected)
self.assertEqual(self.catalog.lookup(first['record_name'].rsplit('.', 1)[0] + '.999999'), [])
self.assertEqual(self.catalog.lookup(''), [])
self.assertEqual(self.catalog.lookup('not-an-accession'), [])
class CoordinateTests(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
directory = Path(self.temp.name)
rows = []
for version in (1, 2):
rows.append({'assembly_accession': f'GCA_000001.{version}',
'record_name': f'TEST000001.{version}',
'source_key': f'GCA_000001.{version}|TEST000001.{version}',
'segment_index': 1, 'segment_start_bp': 100,
'segment_end_bp': 105,
'pred_prob_positive_strand_cds': [0., 0.2, 0.4, 0.6, 0.8],
'pred_prob_negative_strand_cds': [1., 0.8, 0.6, 0.4, 0.2]})
pq.write_table(pa.Table.from_pylist(rows), directory / 'sample.parquet')
(directory / 'manifest.json').write_text('{}')
self.catalog = Catalog(directory)
def tearDown(self):
self.temp.cleanup()
def test_version_resolution(self):
self.assertEqual(self.catalog.lookup('GCA_000001'), [0, 1])
self.assertEqual(self.catalog.lookup('GCA_000001.2'), [1])
self.assertEqual(self.catalog.lookup('TEST000001'), [0, 1])
self.assertEqual(self.catalog.lookup('GCA_000001.3'), [])
def test_exact_window_keeps_both_strands_per_base(self):
frame, step = self.catalog.window(0, 103, 105)
self.assertEqual(step, 1)
self.assertEqual(frame['Position (bp)'].tolist(), [103, 104, 103, 104])
np.testing.assert_allclose(frame[frame.Strand == '+ strand']['P(CDS)'], [0.6, 0.8])
np.testing.assert_allclose(frame[frame.Strand == '− strand']['P(CDS)'], [0.4, 0.2])
def test_binned_window_keeps_the_distribution_not_the_mean(self):
"""A mean hides a part-coding bin; the histogram has to keep both bands."""
rows = 4
frame, step = self.catalog.window(0, 100, 105, max_points=2, hist_rows=rows)
self.assertEqual(step, 3)
self.assertNotIn('Strand', frame.columns)
self.assertEqual(frame['Position (bp)'].tolist()[::rows], [100, 103])
bases = frame['Bases'].to_numpy().reshape(-1, rows)
# Every base is counted exactly once, and only once.
self.assertEqual(bases.sum(), 5)
self.assertEqual(bases.sum(axis=1).tolist(), [3, 2])
# max(P_pos, P_neg) over 100..104 is [1.0, 0.8, 0.6, 0.6, 0.8]; with four
# bands those land in the top band, top, third, third, top.
self.assertEqual(bases.tolist(), [[0, 0, 1, 2], [0, 0, 1, 1]])
np.testing.assert_allclose(frame['Mean P'].to_numpy()[::rows], [0.8, 0.7], atol=1e-6)
def test_binned_window_conserves_every_base(self):
frame, step = self.catalog.window(0, max_points=2, hist_rows=8)
width = self.catalog.records[0]['segment_end_bp'] - self.catalog.records[0]['segment_start_bp']
self.assertGreater(step, 1)
self.assertEqual(frame['Bases'].sum(), width)
self.assertTrue((frame['P(CDS)'] > 0).all() and (frame['P(CDS)'] < 1).all())
def test_invalid_ranges(self):
for start, end in [(99, 103), (100, 106), (103, 103), (104, 102)]:
with self.assertRaises(ValueError):
self.catalog.window(0, start, end)
def test_binary_labels_preserve_exact_transitions_and_coordinates(self):
frame, step = self.catalog.window(0, 100, 105, mode='Binary labels', threshold=0.75)
self.assertEqual(step, 1)
self.assertEqual(frame.Strand.unique().tolist(), ['CDS (either strand)'])
self.assertEqual(frame['Position (bp)'].tolist(), [100, 102, 104, 105])
self.assertEqual(frame['Predicted CDS'].tolist(), [1, 0, 1, 1])
self.assertTrue(set(frame['Predicted CDS']).issubset({0, 1}))
def test_binary_overview_thresholds_before_binning(self):
frame, step = self.catalog.window(0, 100, 105, mode='Binary labels', threshold=0.85,
max_points=2, max_transitions=0)
self.assertEqual(step, 3)
# First bin's combined mean is 0.8, but one base exceeds 0.85.
self.assertEqual(frame['Position (bp)'].tolist(), [100, 103, 105])
self.assertEqual(frame['Predicted CDS'].tolist(), [1, 0, 0])
def test_threshold_limits_and_equality(self):
frame, _ = self.catalog.window(0, mode='Binary labels', threshold=0)
self.assertEqual(set(frame['Predicted CDS']), {1})
frame, _ = self.catalog.window(0, mode='Binary labels', threshold=1)
self.assertEqual(set(frame['Predicted CDS']), {0})
table = pa.Table.from_pylist([{
'pred_prob_positive_strand_cds': [0.5, 0.1, 0.51, 0.49, 0.5],
'pred_prob_negative_strand_cds': [0.1, 0.5, 0.1, 0.1, 0.5]}])
frame, _ = self.catalog.window(0, mode='Binary labels', threshold=0.5, table=table)
self.assertEqual(frame['Position (bp)'].tolist(), [100, 102, 103, 105])
self.assertEqual(frame['Predicted CDS'].tolist(), [0, 1, 0, 0])
for threshold in [-0.1, 1.1, float('nan')]:
with self.assertRaises(ValueError):
self.catalog.window(0, mode='Binary labels', threshold=threshold)
if __name__ == '__main__':
unittest.main()