PepPA / tests /test_science.py
pranamanam's picture
Upload 97 files
98bde72 verified
Raw
History Blame Contribute Delete
6.54 kB
import json,tempfile,unittest
from pathlib import Path
import numpy as np
from peppa.schema import *
from peppa.ternary import equilibrium,fit_cooperativity
from peppa.ptm import fit,loss_gradient
from peppa.metrics import *
from peppa.engine import Engine,Tool
from peppa.trace import Trace,replay
from peppa.structure import af3_input,ternary_interface_features
from peppa.data import normalize_snooppi,connected_splits
from peppa.controllers import ScriptedController
class ScienceTests(unittest.TestCase):
def spec(self):
return DesignSpec(episode_id='test',task='binding',objective='test',targets=[Target(id='T',accession='test',molecule=Molecule(sequence='ACY'),compartment='test')],requirements=[Requirement(endpoint='pKd',unit='logM',direction='ge',threshold=7,assay='test',scale=1)],administration='test')
def decision(self,tool,arguments={}):
return Decision(tool=tool,arguments=arguments,hypothesis='test',evidence_ids=[],decision_summary='test',expected_observation='test')
def test_mass_balance_symmetry(self):
for totals in [(10,50,100),(1,1,1e5),(100,1,1e-3)]:
x=equilibrium(*totals,10,30,20)
np.testing.assert_allclose([x['A']+x['AL']+x['ABL'],x['B']+x['BL']+x['ABL'],x['L']+x['AL']+x['BL']+x['ABL']],totals,rtol=1e-7)
y=equilibrium(totals[1],totals[0],totals[2],30,10,20)
self.assertAlmostEqual(x['ABL'],y['ABL'],places=6)
def test_zero_component(self):
x=equilibrium(10,0,20,5,5,10);self.assertEqual(x['ABL'],0);self.assertAlmostEqual(x['A']+x['AL'],10)
def test_cooperativity_recovery(self):
grid=np.array([[a,30,l] for a in [10,100] for l in np.logspace(-1,3,8)])
y=[equilibrium(*v,15,40,8)['ABL'] for v in grid]
r=fit_cooperativity(grid,y,15,40);self.assertAlmostEqual(r['alpha'],8,places=4)
def test_hook_effect(self):
mid=equilibrium(100,100,100,10,10,20)['ABL'];high=equilibrium(100,100,1e8,10,10,20)['ABL'];self.assertGreater(mid,high)
def test_ptm_gradient(self):
rng=np.random.default_rng(7);b=rng.normal(size=(7,3));t=rng.normal(size=(7,2));theta=rng.normal(size=7);labels=np.array([1,0,1,np.nan,0,1,0])
f,g=loss_gradient(theta,b,t,labels,[(0,1),(2,3)])
numeric=[]
for j in range(len(theta)):
d=np.zeros_like(theta);d[j]=1e-6
numeric.append((loss_gradient(theta+d,b,t,labels,[(0,1),(2,3)])[0]-loss_gradient(theta-d,b,t,labels,[(0,1),(2,3)])[0])/2e-6)
np.testing.assert_allclose(g,numeric,atol=1e-6)
def test_ptm_fit(self):
rng=np.random.default_rng(9);b=rng.normal(size=(80,3));t=rng.normal(size=(80,2));y=(b[:,0]*t[:,0]>0).astype(float)
head,report=fit(b,t,y,l2=.001);self.assertGreater(np.mean((head.logits(b,t)>0)==y),.9)
def test_chemistry_identity(self):
x=Molecule(sequence='ASY');y=Molecule(sequence='ASY',modifications=[Modification(position=3,residue='Y',ccd='PTR')]);self.assertNotEqual(x.identity,y.identity)
a=Molecule(sequence='ASY',bonds=[(1,'N',3,'C')]);b=Molecule(sequence='ASY',bonds=[(3,'C',1,'N')]);self.assertEqual(a.identity,b.identity)
def test_invalid_modification(self):
with self.assertRaises(ValueError):Molecule(sequence='ASY',modifications=[Modification(position=2,residue='Y',ccd='PTR')])
def test_af3_chemistry(self):
s=self.spec();s.targets[0].molecule.modifications=[Modification(position=3,residue='Y',ccd='PTR')]
x=af3_input('test',s.targets,Molecule(sequence='ACD',bonds=[(1,'N',3,'C')]))
self.assertEqual(x['sequences'][0]['protein']['modifications'][0]['ptmType'],'PTR');self.assertEqual(x['bondedAtomPairs'][0],[['B',1,'N'],['B',3,'C']])
def test_weakest_interface(self):
x=ternary_interface_features([[1,.9,.8],[.9,1,.2],[.8,.2,1]]);self.assertEqual(x['weakest_interface'],.2)
def test_score_requires_calibration(self):
m=Measurement(candidate_id='x',endpoint='pKd',value=8,unit='logM',source_id='x',kind='prediction',assay='test')
with self.assertRaises(ValueError):candidate_score([m],self.spec().requirements)
m.lower=7.5;m.upper=8.5;self.assertEqual(candidate_score([m],self.spec().requirements),.5)
def test_selectivity_sign(self):self.assertAlmostEqual(selectivity(10,100),1)
def test_unknown_remains_unknown(self):
self.assertIsNone(normalize_snooppi({'SNOOPPI_final_label':'unknown','partner_A_sequence':'A','partner_B_sequence':'B'})['label'])
def test_connected_split_no_leak(self):
rows=[{'id':'a','target_clusters':['x'],'peptide_clusters':['p'],'source_ids':['1']},{'id':'b','target_clusters':['y'],'peptide_clusters':['q'],'source_ids':['1']},{'id':'c','target_clusters':['y'],'peptide_clusters':['r'],'source_ids':['2']}]
self.assertEqual(len(set(connected_splits(rows))),1)
def test_failed_worker_charged_and_replayed(self):
with tempfile.TemporaryDirectory() as d:
def fail(a,s):raise RuntimeError('deliberate failure')
e=Engine(self.spec(),{'fail':Tool('fail',fail,{'tool_calls':1},'test')},Path(d)/'trace.jsonl')
with self.assertRaises(RuntimeError):e.apply(self.decision('fail'))
r=replay(e.trace.path);self.assertEqual(r['spent']['tool_calls'],1);self.assertEqual(len(r['errors']),1)
def test_endpoint_immutable(self):
with tempfile.TemporaryDirectory() as d:
e=Engine(self.spec(),{},Path(d)/'trace.jsonl')
with self.assertRaises(ValueError):e.apply(self.decision('revise_plan',{'requirements':[]}))
self.assertEqual(e.state['spec']['requirements'][0]['threshold'],7)
def test_trace_tampering(self):
with tempfile.TemporaryDirectory() as d:
p=Path(d)/'trace.jsonl';e=Engine(self.spec(),{},p);s=p.read_text().replace('"threshold":7.0','"threshold":8.0');p.write_text(s)
with self.assertRaises(ValueError):Trace.read(p)
def test_budget_stop(self):
with tempfile.TemporaryDirectory() as d:
s=self.spec();s.budgets['controller_calls']=0;e=Engine(s,{},Path(d)/'trace.jsonl');r=e.run(ScriptedController([]));self.assertTrue(r['stopped'])
def test_redundancy_selection(self):
r=diverse_select({'a':'AAAA','b':'AAAA','c':'CCCC'},{'a':2,'b':1.9,'c':1},2);self.assertEqual(r,['a','c'])
def test_conformal_finite_sample(self):
self.assertTrue(np.isinf(conformal_radius([1],[1],[1],.1)))
self.assertEqual(conformal_radius(np.arange(20),np.arange(20),np.ones(20)),0)
if __name__=='__main__':unittest.main()