Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 39 additions & 1 deletion colabfold/batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -405,6 +405,13 @@ def get(self, x: str, ext:str) -> Path:
self.files[self.tag].append([x,ext,file])
return file

def exists(self, x: str, ext:str) -> Path:
file = self.result_dir.joinpath(f"{self.prefix}_{x}_{self.tag}.{ext}")
if file.is_file():
return True
else:
return False

def set_tag(self, tag):
self.tag = tag

Expand Down Expand Up @@ -435,6 +442,7 @@ def predict_structure(
save_single_representations: bool = False,
save_pair_representations: bool = False,
save_recycles: bool = False,
keep_existing_results: bool = True,
calc_extra_ptm: bool = False,
use_probs_extra: bool = True,
):
Expand Down Expand Up @@ -526,6 +534,32 @@ def callback(result, recycles):

return_representations = save_all or save_single_representations or save_pair_representations

if keep_existing_results and files.exists("unrelaxed","pdb"):
logger.info(f"{tag} was found. Loading scores and skipping!")
if files.exists("scores","json"):
with files.get("scores","json").open("r") as handle:
_scores = json.load(handle)
if is_complex:
mean_scores.append(0.8 * _scores['iptm'] + 0.2 * _scores['ptm'])
else:
mean_scores.append(np.mean(_scores['plddt']))
print_line = ""
conf.append({})
for x,y in [["plddt","pLDDT"],["ptm","pTM"],["iptm","ipTM"]]:
if x in _scores:
if x == "plddt":
print_line += f" {y}={np.mean(_scores[x]):.3g}"
_scores["mean_plddt"] = np.mean(_scores[x])
x = "mean_plddt"
else:
print_line += f" {y}={_scores[x]:.3g}"
conf[-1][x] = float(_scores[x])
conf[-1]["print_line"] = print_line
files.get("unrelaxed","pdb")
continue
else:
logger.error(f"Scores file for {tag} not found")

# predict
result, recycles = \
model_runner.predict(input_features,
Expand Down Expand Up @@ -648,6 +682,9 @@ def callback(result, recycles):

rank, metric = [],[]
result_files = []
if rank_by == "skip":
logger.info(f"Skipping ranking and exiting")
exit()
logger.info(f"reranking models by '{rank_by}' metric")
model_rank = np.array(mean_scores).argsort()[::-1]
for n, key in enumerate(model_rank):
Expand Down Expand Up @@ -2079,7 +2116,7 @@ def main():
help='Choose metric to rank the "--num-models" predicted models.',
type=str,
default="auto",
choices=["auto", "plddt", "ptm", "iptm", "multimer"],
choices=["auto", "plddt", "ptm", "iptm", "multimer", "skip"],
)
output_group.add_argument(
"--stop-at-score",
Expand Down Expand Up @@ -2239,6 +2276,7 @@ def comma_separated_list(arg_string):
version += f" ({commit})"

logger.info(f"Running colabfold {version}")
logger.info(f"Version modded by SSchott to use CBCSrv local MMSeqs2 and keep existing results!")

data_dir = Path(args.data or default_data_dir)

Expand Down
2 changes: 1 addition & 1 deletion colabfold/download.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@

# The data dir location logic switches between a version with and one without "params" because alphafold
# always internally joins "params". (We should probably patch alphafold)
default_data_dir = Path(appdirs.user_cache_dir(__package__ or "colabfold"))
default_data_dir = Path("/apps/local/cbclab/COLABFOLD/colabfold")

def download(url, params_dir, size_queue, progress_queue):
try:
Expand Down
12 changes: 6 additions & 6 deletions colabfold/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,15 +22,15 @@

If you're sure you want to run without a GPU, pass `--cpu`"""

DEFAULT_API_SERVER = "https://api.colabfold.com"
DEFAULT_API_SERVER = "http://192.168.62.51:9091"

ACCEPT_DEFAULT_TERMS = \
"""
WARNING: You are welcome to use the default MSA server, however keep in mind that it's a
limited shared resource only capable of processing a few thousand MSAs per day. Please
submit jobs only from a single IP address. We reserve the right to limit access to the
server case-by-case when usage exceeds fair use. If you require more MSAs: You can
precompute all MSAs with `colabfold_search` or host your own API and pass it to `--host-url`
WARNING: You are using the internal CBCSRV MMseqs2 server, with a local ColabFold 1.5.5 installation.
This server was last updated 22.10.2024. The databases or the ColabFold installation might be too old
at the time of use. Please be aware, and consider getting a newer version by youself if that is the case!
You can use the internal MMseqs2 server API by using http://192.168.62.51:9091 as the host url in the
colabfold_batch script. If the server does not respond, ask around!
"""

class TqdmHandler(logging.StreamHandler):
Expand Down