Skip to content

Commit d45c84f

Browse files
committed
check in zero shot cl eval.
1 parent 46ed24d commit d45c84f

14 files changed

Lines changed: 2409660 additions & 33 deletions

backend/src/folde/campaign.py

Lines changed: 3 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -461,11 +461,9 @@ def simulate_campaign(
461461
cache_key
462462
]
463463
naturalness_ensemble_df = entire_naturalness_df[
464-
(
465-
"wt_marginal"
466-
if model_config.few_shot_naturalness_column is None
467-
else model_config.few_shot_naturalness_column
468-
)
464+
["wt_marginal"]
465+
if model_config.naturalness_columns is None
466+
else model_config.naturalness_columns
469467
]
470468
embedding_series = embedding_df[
471469
"embedding" if model_config.embedding_column is None else model_config.embedding_column

backend/src/folde/types.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,14 +7,17 @@ class FolDEModelConfig(BaseModel):
77
name: str
88
data_split_mode: Optional[str] = None # Can be "1-VS-REST" or "2-VS-REST" etc.
99
naturalness_model_id: str
10+
naturalness_columns: Optional[List[str]] = None
1011
embedding_model_id: str
1112
embedding_column: Optional[str] = None
1213
zero_shot_model_name: str
1314
zero_shot_model_params: Dict[str, Any]
1415
few_shot_model_name: str
15-
few_shot_naturalness_column: Optional[str] = None
1616
few_shot_model_params: Dict[str, Any]
1717

18+
# DEPRECATED.
19+
few_shot_naturalness_column: Optional[str] = None
20+
1821

1922
class MutantMetrics(BaseModel):
2023
"""Stores dense information about each mutants tested in the simulation."""

backend/src/folde/zero_shot_models.py

Lines changed: 11 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -113,20 +113,7 @@ def get_top_n(
113113
self.selection_debug_info = {"sorts": {}}
114114

115115
if self.decision_mode == "constantliar" or self.decision_mode == "krigingbeliever":
116-
if self.lie_noise_stddev_multiplier is not None:
117-
lie_noise_stddev_multiplier = self.lie_noise_stddev_multiplier
118-
elif self.lie_noise_stddev_multiplier_schedule is not None:
119-
lie_noise_stddev_multiplier = self.lie_noise_stddev_multiplier_schedule[0]
120-
self.lie_noise_stddev_multiplier_schedule = (
121-
self.lie_noise_stddev_multiplier_schedule[1:]
122-
)
123-
logging.info(
124-
f"Using lie_noise_stddev_multiplier {lie_noise_stddev_multiplier} for this round and {self.lie_noise_stddev_multiplier_schedule} remaining."
125-
)
126-
else:
127-
raise ValueError(
128-
"Either lie_noise_stddev_multiplier or lie_noise_stddev_multiplier_schedule must be set."
129-
)
116+
assert self.lie_noise_stddev_multiplier is not None
130117

131118
ensemble_scores = get_consensus_scores(ensemble_of_predictions, "mean")
132119
pred_df = {
@@ -142,7 +129,7 @@ def get_top_n(
142129
cl_considerations.to_numpy(),
143130
cl_considerations.index.to_numpy(),
144131
n,
145-
lie_noise_stddev_multiplier=lie_noise_stddev_multiplier,
132+
lie_noise_stddev_multiplier=self.lie_noise_stddev_multiplier,
146133
choice_of_baseline="min" if self.decision_mode == "constantliar" else "mean",
147134
ucb_beta=self.ucb_beta if self.ucb_beta is not None else 0.0,
148135
)
@@ -229,41 +216,37 @@ def __init__(self, **kwargs):
229216
"""
230217
super().__init__(**kwargs)
231218
self.is_pretrained = False
232-
self.single_mutant_naturalness_ensemble: Optional[pd.DataFrame] = None
219+
self.single_mutant_naturalness_df: Optional[pd.DataFrame] = None
233220
self.knns: Dict[str, KNeighborsRegressor] = {}
234221

235222
def pretrain(
236223
self,
237224
naturalness_df: pd.DataFrame,
238225
embedding_series: pd.Series,
239226
) -> "NaturalnessZeroShotModel":
227+
if self.is_pretrained:
228+
raise ValueError("Model is already pretrained.")
240229

241230
self.knns = {}
242231
for naturalness_column in naturalness_df.columns:
243232
naturalness_series = naturalness_df[naturalness_column]
244233
assert naturalness_series.index.equals(embedding_series.index)
245234
assert embedding_series.index.is_unique, "embedding_series contains duplicate indices"
246235

247-
if self.is_pretrained:
248-
raise ValueError("Model is already pretrained.")
249236

250237
assert not naturalness_series.isna().any(), "naturalness_series contains NANs"
251238

252-
logging.info(
253-
f"Pretraining model with naturalness data with {len(naturalness_series)} naturalness measurements."
254-
)
255-
256239
X = np.array([np.array(emb) for emb in embedding_series.values])
257240
y = naturalness_series.values
258241

259242
# Fit KNN regressor
260243
knn = KNeighborsRegressor(n_neighbors=5)
261244
knn.fit(X, y)
262-
self.is_pretrained = True
263245
self.knns[naturalness_column] = knn
264246

265247
# Also store the naturalness series for prediction later.
266-
self.single_mutant_naturalness_ensemble = naturalness_df
248+
self.single_mutant_naturalness_df = naturalness_df
249+
self.is_pretrained = True
267250

268251
return self
269252

@@ -279,9 +262,9 @@ def predict(
279262
Returns:
280263
Array of prediction scores based on naturalness
281264
"""
282-
assert self.single_mutant_naturalness_ensemble is not None
265+
assert self.single_mutant_naturalness_df is not None
283266
assert len(naturalness_df.columns) == len(self.knns)
284-
assert len(naturalness_df.columns) == len(self.single_mutant_naturalness_ensemble)
267+
assert set(naturalness_df.columns) == set(self.single_mutant_naturalness_df.columns)
285268

286269
# Get the ensemble of naturalness scores with missing values imputed.
287270

@@ -290,7 +273,7 @@ def predict(
290273
for naturalness_column in naturalness_df.columns:
291274

292275
naturalness_series = naturalness_df[naturalness_column]
293-
single_mutant_naturalness_series = self.single_mutant_naturalness_ensemble[
276+
single_mutant_naturalness_series = self.single_mutant_naturalness_df[
294277
naturalness_column
295278
]
296279
knn = self.knns[naturalness_column]
@@ -314,6 +297,7 @@ def get_naturalness(seq_id, direct_naturalness) -> float:
314297
)
315298
return computed_naturalness
316299

300+
naturalness_series.index.name = 'seq_id'
317301
computed_naturalness_series = naturalness_series.reset_index(name="wt_marginal").apply(
318302
lambda r: get_naturalness(r.seq_id, r.wt_marginal), axis=1
319303
)

backend/src/notebooks/jacob/250713_naturalness_param_scan.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,10 @@
6464
ModelDiff(name="650m", diffs={"naturalness_model_id": "650m"}),
6565
ModelDiff(name="3b", diffs={"naturalness_model_id": "3b"}),
6666
ModelDiff(name="15b", diffs={"naturalness_model_id": "15b"}),
67+
ModelDiff(name="v1", diffs={
68+
"naturalness_model_id": "v1",
69+
"few_shot_naturaless_column": ['wt_marginal_1', 'wt_marginal_2', 'wt_marginal_3', 'wt_marginal_4', 'wt_marginal_5'],
70+
}),
6771
],
6872
)
6973

backend/src/notebooks/jacob/250718_train_ablation.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -196,6 +196,17 @@
196196
"few_shot_model_params.shrink_and_perturb_params": [0.8, 0.0001],
197197
},
198198
),
199+
ModelDiff(name="with-esmv1", diffs={
200+
"naturalness_model_id": "v1",
201+
"few_shot_naturaless_column": ['wt_marginal_1', 'wt_marginal_2', 'wt_marginal_3', 'wt_marginal_4', 'wt_marginal_5'],
202+
}),
203+
ModelDiff(name="with-esmv1-zsclS025", diffs={
204+
"naturalness_model_id": "v1",
205+
"few_shot_naturaless_column": ['wt_marginal_1', 'wt_marginal_2', 'wt_marginal_3', 'wt_marginal_4', 'wt_marginal_5'],
206+
207+
"zero_shot_model_params.decision_mode": "constantliar",
208+
"zero_shot_model_params.lie_noise_stddev_multiplier": 0.25,
209+
}),
199210
],
200211
)
201212

Lines changed: 126 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,126 @@
1+
# Configure logging
2+
import logging
3+
4+
logging.basicConfig(
5+
level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
6+
)
7+
logger = logging.getLogger()
8+
9+
from folde.campaign import simulate_campaigns_with_config_checkpoints
10+
from folde.types import FolDEModelConfig, ModelDiff
11+
from folde.util import apply_diff_list_to_config
12+
13+
folde_train_dms_ids = [
14+
"ANCSZ_Hobbs_2022",
15+
"BLAT_ECOLX_Firnberg_2014",
16+
"CBS_HUMAN_Sun_2020",
17+
"HEM3_HUMAN_Loggerenberg_2023",
18+
"HSP82_YEAST_Flynn_2019",
19+
"HXK4_HUMAN_Gersing_2022_activity",
20+
"OXDA_RHOTO_Vanella_2023_activity",
21+
"PPM1D_HUMAN_Miller_2022",
22+
"SHOC2_HUMAN_Kwon_2022",
23+
]
24+
25+
# Example configuration
26+
NAME = "250917_zero_shot_cl"
27+
28+
random_config = FolDEModelConfig(
29+
name="Random",
30+
naturalness_model_id="600m",
31+
embedding_model_id="300m",
32+
zero_shot_model_name="RandomZeroShotModel",
33+
zero_shot_model_params={},
34+
few_shot_model_name="RandomFewShotModel",
35+
few_shot_model_params={},
36+
)
37+
38+
random_forest_config = FolDEModelConfig(
39+
name="RandomToRandomForest",
40+
naturalness_model_id="600m",
41+
embedding_model_id="300m",
42+
zero_shot_model_name="RandomZeroShotModel",
43+
zero_shot_model_params={},
44+
few_shot_model_name="RandomForestFewShotModel",
45+
few_shot_model_params={
46+
"n_estimators": 100,
47+
"criterion": "friedman_mse",
48+
"max_depth": None,
49+
"min_samples_split": 2,
50+
"min_samples_leaf": 1,
51+
"min_weight_fraction_leaf": 0.0,
52+
"max_features": 1.0,
53+
"max_leaf_nodes": None,
54+
"min_impurity_decrease": 0.0,
55+
"bootstrap": True,
56+
"oob_score": False,
57+
"n_jobs": None,
58+
"verbose": 0,
59+
"warm_start": False,
60+
"ccp_alpha": 0.0,
61+
"max_samples": None,
62+
},
63+
)
64+
65+
folde_config = FolDEModelConfig(
66+
name="FolDE",
67+
# Required parameters
68+
naturalness_model_id="600m", # ESM-2 650M model
69+
embedding_model_id="300m", # Same model for embeddings
70+
zero_shot_model_name="NaturalnessZeroShotModel",
71+
zero_shot_model_params={},
72+
# Few-shot model configuration (used after first round)
73+
few_shot_model_name="TorchMLPFewShotModel",
74+
few_shot_model_params={
75+
"pretrain": True,
76+
"pretrain_epochs": 50,
77+
"ensemble_size": 5,
78+
"embedding_dim": 960,
79+
"hidden_dims": [100, 50],
80+
"dropout": 0.2,
81+
"learning_rate": 3e-4,
82+
"weight_decay": 1e-5,
83+
"train_epochs": 200,
84+
"train_patience": 40,
85+
"val_frequency": 10,
86+
"do_validation_with_pair_fraction": 0.2,
87+
"decision_mode": "constantliar",
88+
"lie_noise_stddev_multiplier": 0.25,
89+
},
90+
)
91+
92+
mlp_config_list = apply_diff_list_to_config(
93+
folde_config,
94+
[
95+
ModelDiff(name="no-constantliar", diffs={"few_shot_model_params.decision_mode": "mean"}),
96+
ModelDiff(name="with-esmv1", diffs={
97+
"naturalness_model_id": "1v",
98+
"naturalness_columns": ['wt_marginal_1', 'wt_marginal_2', 'wt_marginal_3', 'wt_marginal_4', 'wt_marginal_5'],
99+
}),
100+
ModelDiff(name="with-esmv1-zsclS025", diffs={
101+
"naturalness_model_id": "1v",
102+
"naturalness_columns": ['wt_marginal_1', 'wt_marginal_2', 'wt_marginal_3', 'wt_marginal_4', 'wt_marginal_5'],
103+
104+
"zero_shot_model_params.decision_mode": "constantliar",
105+
"zero_shot_model_params.lie_noise_stddev_multiplier": 0.25,
106+
}),
107+
],
108+
)
109+
110+
config_list = [random_config, random_forest_config] + mlp_config_list
111+
112+
print(f"Config 1/{len(config_list)}:")
113+
print(config_list[0].model_dump_json(indent=2))
114+
115+
results = simulate_campaigns_with_config_checkpoints(
116+
eval_prefix=NAME,
117+
dms_ids=folde_train_dms_ids,
118+
config_list=config_list,
119+
checkpoint_dir="notebooks/jacob/model_evals",
120+
round_size=16,
121+
number_of_simulations=10,
122+
activity_column="DMS_score",
123+
max_rounds=6,
124+
random_seed=42,
125+
num_workers=10,
126+
)

0 commit comments

Comments
 (0)