File size: 12,769 Bytes
d766458 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 | # Copyright 2025 ByteDance and/or its affiliates.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import json
import logging
import os
from abc import ABC, abstractmethod
from functools import cached_property
from typing import Optional
import numpy as np
import pandas as pd
from pxdbench.metrics import diversity, secondary
from pxdbench.tools.af2.af2_predictor import AF2ComplexPredictor, AF2MonomerPredictor
from pxdbench.tools.ptx.interface import ProtenixAPI
from pxdbench.tools.registry import get_backend
logger = logging.getLogger(__name__)
class BaseTask(ABC):
task_type: str
task_name: str
backend_spec: Optional[str] = None
def __init__(self, input_data, cfg, device_id: int, seed: int):
self.cfg = cfg
self.device_id = device_id
self.seed = seed
assert "pdb_dir" in input_data
assert "pdb_names" in input_data
self.pdb_dir = input_data["pdb_dir"]
self.pdb_names = input_data["pdb_names"]
self.out_dir = input_data.get("out_dir", os.path.dirname(self.pdb_dir))
self.sample_fn = input_data.get("sample_fn", "sample_level_output.csv")
self.summary_fn = input_data.get("summary_fn", "summary_output.json")
# Default values
self.num_seqs = cfg.get("num_seqs", 4)
self.use_gt_seq = cfg.get("use_gt_seq", False)
self._ptx_mini_inst: Optional[ProtenixAPI] = None
self._ptx_inst: Optional[ProtenixAPI] = None
# Check pdb_paths
self.process_pdb_paths()
def get_device(self):
if self.device_id >= 0:
return f"cuda:{self.device_id}"
else:
return "cpu"
@cached_property
def ptx_factory(self):
return get_backend(self.backend_spec)
def get_ptx(self, is_large: bool = False) -> ProtenixAPI:
if is_large:
if self._ptx_inst is None:
self._ptx_inst = self.ptx_factory(
cfg=self.cfg.tools.ptx,
device=self.get_device(),
)
return self._ptx_inst
if self._ptx_mini_inst is None:
self._ptx_mini_inst = self.ptx_factory(
cfg=self.cfg.tools.ptx_mini,
device=self.get_device(),
)
return self._ptx_mini_inst
@abstractmethod
def design_sequence(self):
pass
@abstractmethod
def run(self):
"""Execute the task"""
pass
@staticmethod
def summary_from_df(
sample_df: pd.DataFrame,
exclude_keys=["name", "seq_idx", "sequence"],
other_metrics={},
):
"""
Compute summary metrics from a DataFrame.
Args:
sample_df (pd.DataFrame): Input DataFrame containing sample data.
exclude_keys (list, optional): Columns to exclude from summary. Defaults to ["name", "seq_idx", "sequence"].
other_metrics (dict, optional): Additional metrics to include. Defaults to {}.
Returns:
dict: A dictionary containing computed summary metrics.
"""
metrics = {}
for col in sample_df.columns:
if col in exclude_keys:
continue
col_data = sample_df[col].dropna()
# numeric column (float/int)
if pd.api.types.is_numeric_dtype(sample_df[col]):
values = col_data.values
metrics[f"{col}.avg"] = float(np.mean(values))
# metrics[f"{col}.std"] = float(np.std(values))
# list of numbers (object column but with list values)
elif pd.api.types.is_object_dtype(col_data):
if all(isinstance(x, list) for x in col_data):
try:
# Flatten and compute avg/std per sample
per_sample_mean = col_data.apply(lambda x: np.mean(x))
metrics[f"{col}.avg"] = float(np.mean(per_sample_mean))
# metrics[f"{col}.std"] = float(np.std(per_sample_mean))
except:
pass # fallback if something is not list of numbers
if "_success" in col:
metrics[f"{col}.count"] = int(np.sum(sample_df[col].astype(bool)))
metrics.update(other_metrics)
return metrics
@staticmethod
def compute_success_rate(filters_cfg, metrics: pd.DataFrame) -> pd.DataFrame:
"""
Compute success rate for each filter based on metrics.
Args:
filters_cfg (dict): Configuration for filters.
metrics (pd.DataFrame): DataFrame containing metrics.
Returns:
pd.DataFrame: Updated DataFrame with success rate metrics.
"""
for filter_name, filter_details in filters_cfg.items():
missing = [k for k in filter_details.keys() if k not in metrics.columns]
def row_success(row):
for metric_name, (sym, thres) in filter_details.items():
if metric_name not in row:
continue
value = row[metric_name]
if value is None:
return None
if isinstance(value, list):
# check whether there is any sample pass the filter
value = min(value) if sym == "<" else max(value)
if sym == "<" and value >= thres:
return 0
if sym == ">" and value <= thres:
return 0
return 1
if missing:
print(
f"Missing columns {missing} for filter '{filter_name}'. Available columns: {list(metrics.columns)}"
)
metrics[f"{filter_name}_success"] = None
metrics[f"{filter_name}_success_ignore_missing"] = metrics.apply(
row_success, axis=1
)
else:
metrics[f"{filter_name}_success"] = metrics.apply(row_success, axis=1)
metrics[f"{filter_name}_success_ignore_missing"] = metrics[
f"{filter_name}_success"
]
return metrics
def process_pdb_paths(self):
"""
Validate and filter PDB file paths based on their existence.
This method checks if each PDB file specified in self.pdb_names exists in the
directory specified by self.pdb_dir. It filters out any PDB names that don't
correspond to existing files and updates self.pdb_names to only contain valid names.
Logs a warning message for each PDB file that is not found and skipped.
"""
# Check if pdb_paths are valid
valid_pdb_names = []
for name in self.pdb_names:
pdb_path = os.path.join(self.pdb_dir, name + ".pdb")
# File exists
if not os.path.exists(pdb_path):
logger.warning(
f"pdb_path {pdb_path} does not exist. Will skip this file."
)
continue
valid_pdb_names.append(name)
self.pdb_names = list(valid_pdb_names)
return
def cal_diversity(self, pdb_names=None, binder_chain=None):
"""
Calculate diversity of PDB structures. May be slow.
Args:
pdb_names (list, optional): List of PDB names to consider. Defaults to None.
binder_chain (str, optional): Chain ID of the binder. Defaults to None.
Returns:
float: Diversity value.
"""
if self.eval_diversity:
all_names = self.pdb_names if pdb_names is None else pdb_names
pdb_paths = [
os.path.join(self.pdb_dir, name + ".pdb") for name in all_names
]
div = diversity.compute_diversity(pdb_paths, binder_chain)
else:
div = -1
return div
def cal_secondary(self, results, chain_id=None):
"""
Calculate secondary structure metrics for PDB structures.
Args:
results (list): List of dictionaries containing PDB structure information.
chain_id (str, optional): Chain ID of the binder. Defaults to None.
"""
for item in results:
pdb_path = os.path.join(self.pdb_dir, item["name"] + ".pdb")
alpha, beta, loop = secondary.cacl_secondary_structure(pdb_path, chain_id)
Rg, ref_ratio = secondary.get_chain_rg(pdb_path, chain_id)
item.update(
{
"alpha": alpha,
"beta": beta,
"loop": loop,
"Rg": Rg,
"ref_ratio": ref_ratio,
}
)
def af2_complex_predict(self, data_list, save_dir, verbose=True):
"""
Run AF2 complex prediction.
Args:
data_list (list): List of data samples.
save_dir (str): Directory to save predictions.
verbose (bool, optional): Whether to print verbose output. Defaults to True.
"""
assert self.task_type in ["binder"]
predictor = AF2ComplexPredictor(
self.cfg.tools.af2,
device_id=self.device_id,
verbose=verbose,
seed=self.seed,
)
predictor.predict(
input_dir=self.pdb_dir,
save_dir=save_dir,
design_pdb_dir=self.pdb_dir,
data_list=data_list,
cond_chain=",".join(self.cond_chains),
binder_chain=",".join(self.binder_chains),
)
def af2_monomer_predict(self, data_list, save_dir, verbose=True):
"""
Run AF2 monomer prediction.
Args:
data_list (list): List of data samples.
save_dir (str): Directory to save predictions.
verbose (bool, optional): Whether to print verbose output. Defaults to True.
"""
assert self.task_type in ["binder", "ligand_binder"]
predictor = AF2MonomerPredictor(
self.cfg.tools.af2,
device_id=self.device_id,
verbose=verbose,
seed=self.seed,
)
predictor.predict(
save_dir=save_dir,
design_pdb_dir=self.pdb_dir,
data_list=data_list,
binder_chain=self.binder_chains[0],
)
def protenix_predict(self, data_list, orig_seqs=None, is_large=False):
"""
Run Protenix prediction.
Args:
data_list (list): List of data samples.
is_large (bool, optional): Whether to use the large model. Defaults to False.
"""
ptx_cfg = self.cfg.tools.ptx if is_large else self.cfg.tools.ptx_mini
ptx_filter = self.get_ptx(is_large)
dump_dir = os.path.join(
self.out_dir, "ptx_pred" if is_large else "ptx_mini_pred"
)
# HARDCODE binder chain idx
binder_chain_idx = 0 if self.binder_chains[0] == "A" else None
json_path = ptx_filter.prepare_json(
self.pdb_dir,
data_list,
dump_dir=dump_dir,
binder_chain_idx=binder_chain_idx,
orig_seqs=orig_seqs,
use_template=ptx_cfg.get("use_template", False),
)
pred_pdb_paths = ptx_filter.predict(
input_json_path=json_path,
design_pdb_dir=self.pdb_dir,
data_list=data_list,
dump_dir=dump_dir,
seed=self.seed,
N_sample=ptx_cfg.N_sample,
N_step=ptx_cfg.N_step,
step_scale_eta=ptx_cfg.step_scale_eta,
gamma0=ptx_cfg.gamma0,
N_cycle=ptx_cfg.N_cycle,
binder_chain_idx=binder_chain_idx,
use_msa=ptx_cfg.get("use_msa", True),
suffix="_mini" if not is_large else "",
)
return pred_pdb_paths
|