Skip to content

Commit 2ab81af

Browse files
authored
Fix test_smoke_molecule_pgd and test_standard_pgd regressions (#54)
- Restore 20-molecule SMILES lists and subsample_size=8 in test_molecule_metrics.py. The camera-ready commit (#48) shrunk these to 10 molecules with subsample_size=4, which is too small for StratifiedKFold(n_splits=4). - Use np.isclose(rtol=1e-3) instead of exact equality in test_standard_pgd. The shared TabPFN classifier instance in StandardPGD causes minor float divergence (~1e-5) compared to fresh per-descriptor instances.
1 parent 5707aa2 commit 2ab81af

2 files changed

Lines changed: 25 additions & 3 deletions

File tree

tests/test_molecule_metrics.py

Lines changed: 21 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -20,22 +20,42 @@
2020
"CC(C)Cc1ccc(cc1)C(C)C(=O)O",
2121
"CC1(C)SC2C(NC(=O)C2=O)C1(C)C(=O)N",
2222
"C1C(=O)N(C2=CC=CC=C12)C3=CC=C(C=C3)C(F)(F)F",
23+
"CCCCCCOc1ccc(C(=O)C=Cc2c(C=Cc3ccc(OC)cc3)cc(OC)cc2OC)cc1",
2324
"O=C(Nc1nc(-c2ccc(Cl)s2)cs1)c1ccncc1",
25+
"COc1nc(N(C)C)ncc1-n1nc2c(c1C(C)C)C(c1ccc(C#N)c(F)c1)N(c1c[nH]c(=O)c(Cl)c1)C2=O",
2426
"Cc1ncc([N+](=O)[O-])n1CC(=O)Nc1ccccc1",
27+
"CCOC(=O)N=C(NC(C)C)c1ccc(-c2ccc(-c3ccc(C(=NC(=O)OCC)NC(C)C)cc3)o2)cc1",
28+
"COCC(=O)N1CCC(C2CC(C(F)(F)F)n3nc(C)cc3N2)CC1",
29+
"Fc1ccc(C(OCCN2C3CCC2CC(Cc2ccccc2)C3)c2ccc(F)cc2)cc1",
30+
"CC(C)C1CCN(C(=O)C2CCC(=O)N(C3CCCCCC3)C2)CC1",
31+
"CCc1c2c(n(C)c1C)CCCC2=NOC(=O)Nc1ccc(C(C)=O)cc1",
32+
"Cc1cc2c(-c3ccc(S(=O)(=O)NCCO)cc3)ccnc2[nH]1",
2533
"CC(C)N1CCN(C(=O)c2ccc(Oc3ccc(F)cc3)nc2)CC1",
2634
"O=C(Nc1ccc(Cl)c(O)c1)Nc1ccc(Cl)c(Cl)c1",
35+
"O=C1NC(=NN=CC(O)C(O)C(O)C(O)CO)NC1=Cc1ccfo1",
36+
"O=C(NCCO)c1c(O)c2ncc(Cc3ccc(F)cc3)cc2[nH]c1=O",
2737
"NC(=O)C(=O)C(Cc1ccccc1)NC(=O)C1CCN(C(=O)C=Cc2ccncc2)CC1",
2838
]
2939

3040
smiles_b = [
3141
"CC1=C(C=CC=C1)NC2=NC=CC(=N2)NC3=CC=CC=C3C(=O)NC4=CC=CC=N4",
3242
"CN1CCN(C2=CC3=C(C=C2)N=CN3C)C4=CC=CC=C14",
3343
"CN(C)CCCN1C2=CC=CC=C2SC3=CC=CC=C31",
44+
"CC(C)C(C(=O)NCC(C)C)NC(=O)C1=CC=CC=C1C(C)C(C)NC(=O)C2=CN=CC=C2",
3445
"CN1C(=O)CN=C(C2=CC=CC=C12)C3=CC=CC=C3Cl",
3546
"O=C(c1cc(-c2ccc(Cl)cc2Cl)n[nH]1)N1CCCC1",
3647
"COc1cccc(OC)c1C=CC(=O)NC1CCCCC1",
3748
"O=C1NC(O)CCN1C1OC(CO)C(O)C1O",
49+
"Cc1c2ccnc(C(=O)NCCN(C)C)c2cc2c3cc(OC(=O)CCCCC(=O)O)ccc3n(C)c12",
3850
"CC1(C)SC2C(NC(=O)C2=O)C1(C)C(=O)N",
51+
"C1C(=O)N(C2=CC=CC=C12)C3=CC=C(C=C3)C(F)(F)F",
52+
"CCCCCCOc1ccc(C(=O)C=Cc2c(C=Cc3ccc(OC)cc3)cc(OC)cc2OC)cc1",
53+
"O=C(Nc1nc(-c2ccc(Cl)s2)cs1)c1ccncc1",
54+
"COc1nc(N(C)C)ncc1-n1nc2c(c1C(C)C)C(c1ccc(C#N)c(F)c1)N(c1c[nH]c(=O)c(Cl)c1)C2=O",
55+
"Cc1ncc([N+](=O)[O-])n1CC(=O)Nc1ccccc1",
56+
"CCOC(=O)N=C(NC(C)C)c1ccc(-c2ccc(-c3ccc(C(=NC(=O)OCC)NC(C)C)cc3)o2)cc1",
57+
"COCC(=O)N1CCC(C2CC(C(F)(F)F)n3nc(C)cc3N2)CC1",
58+
"Fc1ccc(C(OCCN2C3CCC2CC(Cc2ccccc2)C3)c2ccc(F)cc2)cc1",
3959
"O=C(NCC1CCCO1)c1ccc2c(=O)n(-c3ccccc3)c(=S)[nH]c2c1",
4060
"CCNc1nc(C#N)nc(N2CCCCC2)n1",
4161
]
@@ -103,5 +123,5 @@ def test_smoke_molecule_pgd():
103123
metric = MoleculePGD(mols_a)
104124
metric.compute(mols_b)
105125

106-
metric = MoleculePGDInterval(mols_a, subsample_size=4, num_samples=4)
126+
metric = MoleculePGDInterval(mols_a, subsample_size=8, num_samples=4)
107127
metric.compute(mols_b)

tests/test_polygraphdiscrepancy.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import numpy as np
12
import pytest
23

34
from sklearn.linear_model import LogisticRegression
@@ -234,9 +235,10 @@ def test_standard_pgd(dense_graphs, sparse_graphs):
234235
}
235236

236237
for name, (_, individual_result) in individual_results.items():
238+
joint = result["subscores"][name]
237239
assert isinstance(individual_result, float)
238-
assert individual_result == result["subscores"][name], (
239-
f"Individual result {individual_result} for descriptor {name} does not match the overall result {result['subscores'][name]}"
240+
assert np.isclose(individual_result, joint, rtol=1e-3), (
241+
f"Individual result {individual_result} for descriptor {name} does not match the overall result {joint}"
240242
)
241243

242244
metric = StandardPGDInterval(dense_graphs, subsample_size=10, num_samples=4)

0 commit comments

Comments
 (0)