Download tests/test_catalog.py from HuggingFaceBio/carbon-a-database-explorer: direct link, hf CLI and curl.
- Browser
- Download file 7.22 kB
-
https://huggingface.co/spaces/HuggingFaceBio/carbon-a-database-explorer/resolve/main/tests/test_catalog.py
- Command line
-
hf download hf://spaces/HuggingFaceBio/carbon-a-database-explorer/tests/test_catalog.py
-
curl -L -o test_catalog.py https://huggingface.co/spaces/HuggingFaceBio/carbon-a-database-explorer/resolve/main/tests/test_catalog.py
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): | |
| 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() | |