rnaseek-full / tests /test_sft_data_mapping.py
schen647's picture
Clarify SFT and RLHF workflows, fix UTR preset, and consolidate shared regression base
7155827 verified
Raw History Blame Contribute Delete
1.51 kB
"""Keep the corrected UTR SFT mapping separate from scalar regression inputs."""
import hashlib
import importlib.util
from pathlib import Path
import unittest
ROOT = Path(__file__).resolve().parents[1]
spec = importlib.util.spec_from_file_location('sft_mapping', ROOT / 'train_sft.py')
sft = importlib.util.module_from_spec(spec)
spec.loader.exec_module(sft)
class SFTDataMappingTests(unittest.TestCase):
def test_utr_wet_uses_fixed_utr_generation_data(self):
preset = sft.PROJECTS['utr-wet']
self.assertEqual(preset['directory'], 'utrgen')
self.assertEqual(preset['length'], 512)
self.assertEqual(preset['batch'], 2)
self.assertEqual(preset['steps'], -1)
expected = {
'train.txt': '48f7e5986a4d5816bdc5f5dc8cab2f2acdc8ba6872879c9d3fff0799a49def4a',
'valid.txt': '139b7462887a63eabacb7273948bf896699eb5a1668b9a9381d595646efe0bb1',
}
for name, digest in expected.items():
path = sft.local_path(preset['directory'] + '/' + name)
self.assertEqual(hashlib.sha256(path.read_bytes()).hexdigest(), digest)
with path.open() as stream:
self.assertTrue(all(line.startswith('~$predict_utr') for line in stream if line.strip()))
self.assertEqual(path.read_bytes(), (ROOT / 'utrgen' / name).read_bytes())
def test_ribozyme_wet_mapping_is_unchanged(self):
self.assertEqual(sft.PROJECTS['ribozyme-wet']['directory'], 'ribozymegen-figure7/filteredwet')