ProRiboGen / api /encode_protein.py
TimelessAEther's picture
Upload ProRiboGen inference package and checkpoints
6dd9839 verified
Raw
History Blame Contribute Delete
2.84 kB
#!/usr/bin/env python3
"""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()