| """Pinned ESM-2 protein embedding generation for external BGCs.""" |
|
|
| from __future__ import annotations |
|
|
| import json |
| import time |
| from pathlib import Path |
| from typing import Any |
|
|
| import h5py |
| import numpy as np |
| import torch |
| from Bio import SeqIO |
|
|
| from .artifacts import sha256_file, write_json_immutable |
|
|
|
|
| MODEL_NAME = "esm2_t33_650M_UR50D" |
| MODEL_LAYER = 33 |
| EMBEDDING_DIMENSION = 1280 |
| MAX_SEQUENCE_LENGTH = 1022 |
|
|
|
|
| def read_fasta(path: str | Path) -> list[tuple[str, str]]: |
| with Path(path).open("r", encoding="utf-8") as handle: |
| records = [(record.id, str(record.seq)) for record in SeqIO.parse(handle, "fasta")] |
| identifiers = [identifier for identifier, _ in records] |
| if len(identifiers) != len(set(identifiers)): |
| raise ValueError("External FASTA identifiers must be unique") |
| if not records: |
| raise ValueError("External FASTA is empty") |
| return records |
|
|
|
|
| def generate_esm2_embeddings( |
| fasta_path: str | Path, |
| output_h5: str | Path, |
| metadata_path: str | Path, |
| batch_size: int = 16, |
| resume: bool = False, |
| allow_cpu: bool = False, |
| ) -> dict[str, Any]: |
| try: |
| import esm |
| except ImportError as error: |
| raise RuntimeError("Install the pinned fair-esm==2.0.0 dependency") from error |
| if not torch.cuda.is_available() and not allow_cpu: |
| raise RuntimeError("CUDA is required unless --allow-cpu is explicitly supplied") |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| output = Path(output_h5) |
| partial = output.with_suffix(output.suffix + ".partial") |
| metadata_output = Path(metadata_path) |
| if output.exists() or metadata_output.exists(): |
| raise FileExistsError("Refusing to overwrite completed external embeddings") |
| if partial.exists() and not resume: |
| raise FileExistsError(f"Partial output exists; pass --resume: {partial}") |
| output.parent.mkdir(parents=True, exist_ok=True) |
|
|
| sequences = read_fasta(fasta_path) |
| already_done: set[str] = set() |
| if resume and partial.exists(): |
| with h5py.File(partial, "r") as handle: |
| already_done = set(handle.keys()) |
| remaining = [(key, value) for key, value in sequences if key not in already_done] |
| remaining.sort(key=lambda item: (len(item[1]), item[0])) |
|
|
| model, alphabet = esm.pretrained.esm2_t33_650M_UR50D() |
| model = model.eval().to(device) |
| if device.type == "cuda": |
| model = model.half() |
| converter = alphabet.get_batch_converter() |
| started = time.time() |
| truncated = sum(len(sequence) > MAX_SEQUENCE_LENGTH for _, sequence in sequences) |
| with h5py.File(partial, "a") as handle: |
| for start in range(0, len(remaining), batch_size): |
| batch = remaining[start : start + batch_size] |
| prepared = [] |
| for identifier, sequence in batch: |
| prepared.append((identifier, sequence[:MAX_SEQUENCE_LENGTH])) |
| labels, _, tokens = converter(prepared) |
| tokens = tokens.to(device) |
| with torch.no_grad(), torch.autocast( |
| device_type=device.type, |
| dtype=torch.float16 if device.type == "cuda" else torch.bfloat16, |
| enabled=device.type == "cuda", |
| ): |
| representations = model( |
| tokens, repr_layers=[MODEL_LAYER], return_contacts=False |
| )["representations"][MODEL_LAYER] |
| for row, (identifier, sequence) in enumerate(prepared): |
| embedding = representations[row, 1 : len(sequence) + 1].mean(dim=0) |
| handle.create_dataset(identifier, data=embedding.float().cpu().numpy().astype(np.float16)) |
| handle.flush() |
| partial.replace(output) |
| metadata = { |
| "schema_version": 1, |
| "model": MODEL_NAME, |
| "fair_esm_version": "2.0.0", |
| "representation_layer": MODEL_LAYER, |
| "embedding_dimension": EMBEDDING_DIMENSION, |
| "maximum_sequence_length": MAX_SEQUENCE_LENGTH, |
| "proteins": len(sequences), |
| "truncated_proteins": truncated, |
| "device": str(device), |
| "elapsed_seconds": time.time() - started, |
| "fasta_sha256": sha256_file(fasta_path), |
| "h5_sha256": sha256_file(output), |
| } |
| write_json_immutable(metadata_output, metadata) |
| return metadata |
|
|