| |
| """Encode an amino-acid sequence to a ProRiboGen-compatible VESM3B H5.""" |
| from __future__ import annotations |
|
|
| import argparse |
| import os |
| import re |
| import subprocess |
| import sys |
| import tempfile |
| from pathlib import Path |
|
|
| from common import PKG_ROOT, device, load_env, resolve, vesm_paths |
|
|
|
|
| def write_fasta(path: Path, p_id: str, sequence: str) -> None: |
| seq = re.sub(r"\s+", "", sequence).upper() |
| if not seq or not re.fullmatch(r"[ACDEFGHIKLMNPQRSTVWY*BXZ]+", seq): |
| raise ValueError("Protein sequence is empty or contains invalid characters") |
| path.write_text(f">{p_id}\n{seq}\n", encoding="utf-8") |
|
|
|
|
| def encode_protein_to_h5( |
| protein_seq: str, |
| *, |
| p_id: str = "QUERY", |
| output_h5: Path | None = None, |
| fp16: bool = True, |
| ) -> Path: |
| load_env() |
| base, weights = vesm_paths() |
| if not base.is_dir(): |
| raise FileNotFoundError(f"VESM 基座不存在: {base}") |
| if not weights.is_file(): |
| raise FileNotFoundError(f"VESM 权重不存在: {weights}") |
|
|
| work = resolve(os.environ.get("WORKSPACE", "workspace")) |
| work.mkdir(parents=True, exist_ok=True) |
| out = output_h5 or (work / f"{p_id}_vesm3b.h5") |
|
|
| with tempfile.TemporaryDirectory(dir=work) as tmp: |
| fa = Path(tmp) / f"{p_id}.fasta" |
| write_fasta(fa, p_id, protein_seq) |
| script = PKG_ROOT / "protein_encoder" / "build_vesm3b_protein_embeddings_h5.py" |
| cmd = [ |
| sys.executable, |
| str(script), |
| "--base-model-dir", |
| str(base), |
| "--vesm-weights", |
| str(weights), |
| "--fasta", |
| str(fa), |
| "--output", |
| str(out), |
| "--device", |
| device(), |
| ] |
| if fp16: |
| cmd.append("--fp16") |
| subprocess.run(cmd, check=True) |
| return out |
|
|
|
|
| def main() -> None: |
| ap = argparse.ArgumentParser(description="Protein AA sequence → VESM3B H5") |
| ap.add_argument("--protein", required=True, help="AA sequence, or @path/to.fasta") |
| ap.add_argument("--p-id", default="QUERY") |
| ap.add_argument("--output", type=Path, default=None) |
| ap.add_argument("--no-fp16", action="store_true") |
| args = ap.parse_args() |
|
|
| load_env() |
| protein = args.protein |
| if protein.startswith("@"): |
| text = Path(protein[1:]).read_text(encoding="utf-8") |
| lines = [ln.strip() for ln in text.splitlines() if ln.strip()] |
| if lines and lines[0].startswith(">"): |
| protein = "".join(lines[1:]) |
| if args.p_id == "QUERY": |
| args.p_id = lines[0][1:].split()[0] |
| else: |
| protein = "".join(lines) |
|
|
| out = encode_protein_to_h5( |
| protein, p_id=args.p_id, output_h5=args.output, fp16=not args.no_fp16 |
| ) |
| print(f"OK -> {out}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|