Skip to content

Commit 945ce03

Browse files
committed
Add 10 epoch train.
1 parent 7924273 commit 945ce03

7 files changed

Lines changed: 780453 additions & 121 deletions

backend/folde/few_shot_models.py

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
)
2929
from app.helpers.sequence_util import sort_seq_id_list
3030
from folde.util import (
31+
NaturalnessImputer,
3132
cluster_sort_seq_ids,
3233
constant_liar_sample,
3334
get_consensus_scores,
@@ -233,6 +234,53 @@ def get_debug_info(self) -> Dict[str, Any]:
233234
return {}
234235

235236

237+
@register_few_shot_model
238+
class NaturalnessFewShotModel(FewShotModel):
239+
"""Use naturalness scores as activity scores."""
240+
241+
def __init__(self, random_state: int, **kwargs):
242+
super().__init__(**kwargs)
243+
self.random_state = random_state
244+
self.naturalness_imputer = NaturalnessImputer()
245+
246+
def fit(
247+
self,
248+
naturalness_df: pd.DataFrame,
249+
embedding_series: pd.Series,
250+
measured_activity_series: pd.Series,
251+
test_naturalness_df: pd.DataFrame | None = None,
252+
test_embedding_series: pd.Series | None = None,
253+
test_activity_series: Optional[pd.Series] = None,
254+
) -> "NaturalnessFewShotModel":
255+
return self
256+
257+
def pretrain(
258+
self,
259+
naturalness_df: pd.DataFrame,
260+
embedding_series: pd.Series,
261+
) -> "NaturalnessFewShotModel":
262+
self.naturalness_imputer.pretrain(naturalness_df, embedding_series)
263+
return self
264+
265+
def predict(
266+
self, naturalness_df: pd.DataFrame, embedding_series: Optional[pd.Series] = None
267+
) -> List[pd.Series]:
268+
"""Predict using naturalness scores.
269+
270+
Args:
271+
naturalness_series: Series containing naturalness scores, some of which may be NAN.
272+
embedding_series: Optional Series containing protein embeddings
273+
274+
Returns:
275+
Array of prediction scores based on naturalness
276+
"""
277+
assert embedding_series is not None
278+
return self.naturalness_imputer.impute(naturalness_df, embedding_series)
279+
280+
def get_debug_info(self) -> Dict[str, Any]:
281+
return {}
282+
283+
236284
# TODO(jacob): Implement GP Model
237285
# https://scikit-learn.org/stable/modules/generated/sklearn.gaussian_process.GaussianProcessRegressor.html
238286

backend/folde/util.py

Lines changed: 110 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,11 @@
11
import json
22
import logging
33
from pathlib import Path
4-
from typing import Any, List, Tuple, Union, cast
4+
from typing import Any, Dict, List, Optional, Tuple, Union, cast
55

66
import numpy as np
77
import pandas as pd
8+
from sklearn.neighbors import KNeighborsRegressor
89
import torch
910
from numpy.typing import NDArray
1011
from pandas import DataFrame
@@ -14,18 +15,10 @@
1415
from scipy.special import softmax
1516
from sklearn.metrics import recall_score
1617

18+
from app.helpers.sequence_util import is_homolog_seq_id
1719
from folde.types import FolDEModelConfig, ModelDiff, ModelEvaluation
1820

1921
DMS_SHORTNAMES = {
20-
# 'A0A140D2T1_ZIKV_Sourisseau_2019': "A0A140D2T1_ZIKV",
21-
# 'A0A2Z5U3Z0_9INFA_Doud_2016': "A0A2Z5U3Z0_9INFA",
22-
# 'ADRB2_HUMAN_Jones_2020': "ADRB2",
23-
# 'BLAT_ECOLX_Stiffler_2015': "BLAT_ECOLX",
24-
# 'C6KNH7_9INFA_Lee_2018': "C6KNH7",
25-
# 'IF1_ECOLI_Kelsic_2016': "IF1_ECOLI",
26-
# 'MK01_HUMAN_Brenan_2016': "MK01_HUMAN",
27-
# 'P53_HUMAN_Giacomelli_2018_Null_Etoposide': "P53",
28-
# 'PHOT_CHLRE_Chen_2023': "PHOT_CHLRE",
2922
"BLAT_ECOLX_Firnberg_2014": "BLAT_ECOLX",
3023
"ANCSZ_Hobbs_2022": "ANCSZ",
3124
"HXK4_HUMAN_Gersing_2022_activity": "HXK4_HUMAN",
@@ -36,26 +29,20 @@
3629
"P53_HUMAN_Giacomelli_2018_Null_Nutlin": "P53_Null",
3730
"HSP82_YEAST_Flynn_2019": "HSP82_YEAST",
3831
"P53_HUMAN_Giacomelli_2018_WT_Nutlin": "P53_WT",
39-
# "MK01_HUMAN_Brenan_2016": "MK01_HUMAN",
4032
"HEM3_HUMAN_Loggerenberg_2023": "HEM3_HUMAN",
4133
"PPM1D_HUMAN_Miller_2022": "PPM1D_HUMAN",
4234
"SPG1_STRSG_Olson_2014": "SPG1",
43-
# VALIDATION TARGETS
4435
"ADRB2_HUMAN_Jones_2020": "ADRB2_HUMAN",
4536
"P53_HUMAN_Giacomelli_2018_Null_Nutlin": "P53_HUMAN_Null",
4637
"P53_HUMAN_Giacomelli_2018_WT_Nutlin": "P53_HUMAN_WT",
4738
"MK01_HUMAN_Brenan_2016": "MK01_HUMAN",
4839
"KCNJ2_MOUSE_Coyote-Maestas_2022_function": "KCNJ2_MOUSE",
4940
"CAS9_STRP1_Spencer_2017_positive": "CAS9_STRP1",
5041
"SC6A4_HUMAN_Young_2021": "SC6A4_HUMAN",
51-
# 'OXDA_RHOTO_Vanella_2023_expression',
52-
# 'HSP82_YEAST_Mishra_2016',
5342
"PTEN_HUMAN_Mighell_2018": "PTEN_HUMAN",
5443
"S22A1_HUMAN_Yee_2023_activity": "S22A1_HUMAN",
5544
"KKA2_KLEPN_Melnikov_2014": "KKA2_KLEPN",
56-
# Include some "easy to engineer" targets.
5745
"PPARG_HUMAN_Majithia_2016": "PPARG_HUMAN",
58-
# 'P53_HUMAN_Giacomelli_2018_Null_Etoposide',
5946
"MET_HUMAN_Estevam_2023": "MET_HUMAN",
6047
"MTHR_HUMAN_Weile_2021": "MTHR_HUMAN",
6148
"LGK_LIPST_Klesmith_2015": "LGK_LIPST",
@@ -595,3 +582,110 @@ def load_checkpointed_model_eval(
595582
raise ValueError(f"Failed to find any matching files.")
596583

597584
return composite_result
585+
586+
587+
class NaturalnessImputer(object):
588+
def __init__(self):
589+
self.is_pretrained = False
590+
self.single_mutant_naturalness_df: Optional[pd.DataFrame] = None
591+
self.knns: Dict[str, KNeighborsRegressor] = {}
592+
593+
def pretrain(self, naturalness_df: pd.DataFrame, embedding_series: pd.Series):
594+
if self.is_pretrained:
595+
raise ValueError("Model is already pretrained.")
596+
597+
self.knns = {}
598+
for naturalness_column in naturalness_df.columns:
599+
naturalness_series = naturalness_df[naturalness_column]
600+
assert naturalness_series.index.equals(embedding_series.index)
601+
assert embedding_series.index.is_unique, "embedding_series contains duplicate indices"
602+
603+
assert not naturalness_series.isna().any(), "naturalness_series contains NANs"
604+
605+
X = np.array([np.array(emb) for emb in embedding_series.values])
606+
y = naturalness_series.values
607+
608+
# Fit KNN regressor
609+
knn = KNeighborsRegressor(n_neighbors=5)
610+
knn.fit(X, y)
611+
self.knns[naturalness_column] = knn
612+
613+
# Also store the naturalness series for prediction later.
614+
self.single_mutant_naturalness_df = naturalness_df
615+
self.is_pretrained = True
616+
617+
def impute(self, naturalness_df: pd.DataFrame, embedding_series: pd.Series) -> List[pd.Series]:
618+
assert self.single_mutant_naturalness_df is not None
619+
assert len(naturalness_df.columns) == len(self.knns)
620+
assert set(naturalness_df.columns) == set(self.single_mutant_naturalness_df.columns)
621+
622+
# Get the ensemble of naturalness scores with missing values imputed.
623+
624+
ensemble_of_computed_naturalness: List[pd.Series] = []
625+
626+
for naturalness_column in naturalness_df.columns:
627+
628+
naturalness_series = naturalness_df[naturalness_column]
629+
single_mutant_naturalness_series = self.single_mutant_naturalness_df[naturalness_column]
630+
knn = self.knns[naturalness_column]
631+
632+
def get_naturalness(seq_id, direct_naturalness) -> float:
633+
"""Try computing naturalness for mutants even if none was provided by extrapolating for multimutants."""
634+
if direct_naturalness is not None and not pd.isna(direct_naturalness):
635+
return direct_naturalness
636+
637+
if is_homolog_seq_id(seq_id):
638+
return np.nan
639+
640+
# Break it down into single mutants.
641+
seq_id_parts = seq_id.split("_")
642+
643+
# For multimutants, we compute naturalness as the product of the naturalness of the single mutants.
644+
computed_naturalness = single_mutant_naturalness_series.loc[seq_id_parts].sum()
645+
if pd.isna(computed_naturalness):
646+
raise ValueError(
647+
f"Computed naturalness is NAN for {seq_id} with parts {seq_id_parts}"
648+
)
649+
return computed_naturalness
650+
651+
naturalness_series.index.name = "seq_id"
652+
computed_naturalness_series = naturalness_series.reset_index(name="wt_marginal").apply(
653+
lambda r: get_naturalness(r.seq_id, r.wt_marginal), axis=1
654+
)
655+
computed_naturalness_series.index = naturalness_series.index
656+
657+
# Do KNN imputation to fill in NANs from homologs.
658+
if computed_naturalness_series.isna().any():
659+
if not self.is_pretrained:
660+
raise ValueError(
661+
"Model is not pretrained, so cannot fill in NANs from homologs."
662+
)
663+
664+
logging.info(
665+
f"Filling in NANs from homologs for {computed_naturalness_series.isna().sum()}/{len(computed_naturalness_series)} naturalness values."
666+
)
667+
assert embedding_series is not None
668+
embedding_array = np.array([np.array(emb) for emb in embedding_series.values])
669+
naturalness_array = computed_naturalness_series.values
670+
671+
# Find indices for known and missing
672+
missing_mask = computed_naturalness_series.isna().to_numpy()
673+
X_missing = embedding_array[missing_mask]
674+
675+
imputed_values = knn.predict(X_missing)
676+
677+
# Fill in the missing values
678+
imputed_naturalness = naturalness_array.copy()
679+
imputed_naturalness[missing_mask] = imputed_values
680+
681+
# Convert back to Series
682+
computed_naturalness_series = pd.Series(
683+
imputed_naturalness, index=naturalness_series.index
684+
)
685+
686+
if computed_naturalness_series.isna().any():
687+
raise ValueError(
688+
f"Computed naturalness series still has NANs: {computed_naturalness_series.isna().sum()}/{len(computed_naturalness_series)}"
689+
)
690+
ensemble_of_computed_naturalness.append(computed_naturalness_series)
691+
return ensemble_of_computed_naturalness

backend/folde/zero_shot_models.py

Lines changed: 5 additions & 103 deletions
Original file line numberDiff line numberDiff line change
@@ -15,8 +15,7 @@
1515
import pandas as pd
1616
from sklearn.neighbors import KNeighborsRegressor
1717

18-
from app.helpers.sequence_util import is_homolog_seq_id
19-
from folde.util import constant_liar_sample, get_consensus_scores, internal_sample_n_indices
18+
from folde.util import constant_liar_sample, get_consensus_scores, internal_sample_n_indices, NaturalnessImputer
2019

2120
# Registry of available zero-shot models
2221
_ZERO_SHOT_MODELS = {}
@@ -218,38 +217,14 @@ def __init__(self, **kwargs):
218217
**kwargs: Additional parameters
219218
"""
220219
super().__init__(**kwargs)
221-
self.is_pretrained = False
222-
self.single_mutant_naturalness_df: Optional[pd.DataFrame] = None
223-
self.knns: Dict[str, KNeighborsRegressor] = {}
220+
self.naturalness_imputer = NaturalnessImputer()
224221

225222
def pretrain(
226223
self,
227224
naturalness_df: pd.DataFrame,
228225
embedding_series: pd.Series,
229226
) -> "NaturalnessZeroShotModel":
230-
if self.is_pretrained:
231-
raise ValueError("Model is already pretrained.")
232-
233-
self.knns = {}
234-
for naturalness_column in naturalness_df.columns:
235-
naturalness_series = naturalness_df[naturalness_column]
236-
assert naturalness_series.index.equals(embedding_series.index)
237-
assert embedding_series.index.is_unique, "embedding_series contains duplicate indices"
238-
239-
assert not naturalness_series.isna().any(), "naturalness_series contains NANs"
240-
241-
X = np.array([np.array(emb) for emb in embedding_series.values])
242-
y = naturalness_series.values
243-
244-
# Fit KNN regressor
245-
knn = KNeighborsRegressor(n_neighbors=5)
246-
knn.fit(X, y)
247-
self.knns[naturalness_column] = knn
248-
249-
# Also store the naturalness series for prediction later.
250-
self.single_mutant_naturalness_df = naturalness_df
251-
self.is_pretrained = True
252-
227+
self.naturalness_imputer.pretrain(naturalness_df, embedding_series)
253228
return self
254229

255230
def predict(
@@ -264,81 +239,8 @@ def predict(
264239
Returns:
265240
Array of prediction scores based on naturalness
266241
"""
267-
assert self.single_mutant_naturalness_df is not None
268-
assert len(naturalness_df.columns) == len(self.knns)
269-
assert set(naturalness_df.columns) == set(self.single_mutant_naturalness_df.columns)
270-
271-
# Get the ensemble of naturalness scores with missing values imputed.
272-
273-
ensemble_of_computed_naturalness: List[pd.Series] = []
274-
275-
for naturalness_column in naturalness_df.columns:
276-
277-
naturalness_series = naturalness_df[naturalness_column]
278-
single_mutant_naturalness_series = self.single_mutant_naturalness_df[naturalness_column]
279-
knn = self.knns[naturalness_column]
280-
281-
def get_naturalness(seq_id, direct_naturalness) -> float:
282-
"""Try computing naturalness for mutants even if none was provided by extrapolating for multimutants."""
283-
if direct_naturalness is not None and not pd.isna(direct_naturalness):
284-
return direct_naturalness
285-
286-
if is_homolog_seq_id(seq_id):
287-
return np.nan
288-
289-
# Break it down into single mutants.
290-
seq_id_parts = seq_id.split("_")
291-
292-
# For multimutants, we compute naturalness as the product of the naturalness of the single mutants.
293-
computed_naturalness = single_mutant_naturalness_series.loc[seq_id_parts].sum()
294-
if pd.isna(computed_naturalness):
295-
raise ValueError(
296-
f"Computed naturalness is NAN for {seq_id} with parts {seq_id_parts}"
297-
)
298-
return computed_naturalness
299-
300-
naturalness_series.index.name = "seq_id"
301-
computed_naturalness_series = naturalness_series.reset_index(name="wt_marginal").apply(
302-
lambda r: get_naturalness(r.seq_id, r.wt_marginal), axis=1
303-
)
304-
computed_naturalness_series.index = naturalness_series.index
305-
306-
# Do KNN imputation to fill in NANs from homologs.
307-
if computed_naturalness_series.isna().any():
308-
if not self.is_pretrained:
309-
raise ValueError(
310-
"Model is not pretrained, so cannot fill in NANs from homologs."
311-
)
312-
313-
logging.info(
314-
f"Filling in NANs from homologs for {computed_naturalness_series.isna().sum()}/{len(computed_naturalness_series)} naturalness values."
315-
)
316-
assert embedding_series is not None
317-
embedding_array = np.array([np.array(emb) for emb in embedding_series.values])
318-
naturalness_array = computed_naturalness_series.values
319-
320-
# Find indices for known and missing
321-
missing_mask = computed_naturalness_series.isna().to_numpy()
322-
X_missing = embedding_array[missing_mask]
323-
324-
imputed_values = knn.predict(X_missing)
325-
326-
# Fill in the missing values
327-
imputed_naturalness = naturalness_array.copy()
328-
imputed_naturalness[missing_mask] = imputed_values
329-
330-
# Convert back to Series
331-
computed_naturalness_series = pd.Series(
332-
imputed_naturalness, index=naturalness_series.index
333-
)
334-
335-
if computed_naturalness_series.isna().any():
336-
raise ValueError(
337-
f"Computed naturalness series still has NANs: {computed_naturalness_series.isna().sum()}/{len(computed_naturalness_series)}"
338-
)
339-
ensemble_of_computed_naturalness.append(computed_naturalness_series)
340-
341-
return ensemble_of_computed_naturalness
242+
assert embedding_series is not None
243+
return self.naturalness_imputer.impute(naturalness_df, embedding_series)
342244

343245
def get_debug_info(self) -> Dict[str, Any]:
344246
"""Get debug information about the model.

0 commit comments

Comments
 (0)