CodonTransformer / scripts /slurm /run_inference_single.sh
anzhi2710gmailcom's picture
Upload folder using huggingface_hub
53e66de verified
Raw
History Blame Contribute Delete
1.98 kB
#!/usr/bin/env bash
set -euo pipefail
# Single deterministic inference for one protein sequence.
# Run this inside an allocated/interactive GPU session. No SLURM resources are requested here.
PROJECT_DIR="${PROJECT_DIR:-/public/home/scnb9biwet/jiangqq/CodonTransformer-main}"
HF_HOME="${HF_HOME:-/public/home/scnb9biwet/.cache/huggingface}"
CONDA_ENV="${CONDA_ENV:-struct-evo}"
PROTEIN="${PROTEIN:-MFWY}"
ORGANISM="${ORGANISM:-Escherichia coli general}"
OFFLINE="${OFFLINE:-1}"
cd "${PROJECT_DIR}"
export HF_HOME
export PYTHONPATH="${PROJECT_DIR}/model:${PYTHONPATH:-}"
export PROTEIN
export ORGANISM
export OFFLINE
export PYTHONFAULTHANDLER=1
if [[ "${OFFLINE}" == "1" ]]; then
export HF_HUB_OFFLINE=1
export TRANSFORMERS_OFFLINE=1
fi
if [[ -n "${CONDA_ENV}" ]] && command -v conda >/dev/null 2>&1; then
# shellcheck disable=SC1091
source "$(conda info --base)/etc/profile.d/conda.sh"
conda activate "${CONDA_ENV}"
fi
python - <<'PY'
import os
import torch
from transformers import AutoTokenizer, BigBirdForMaskedLM
from CodonTransformer.CodonJupyter import format_model_output
from CodonTransformer.CodonPrediction import predict_dna_sequence
protein = os.environ["PROTEIN"]
organism = os.environ["ORGANISM"]
local_files_only = os.environ.get("OFFLINE", "1") == "1"
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"HF_HOME: {os.environ.get('HF_HOME')}")
print(f"Device: {device}")
print(f"Local files only: {local_files_only}")
tokenizer = AutoTokenizer.from_pretrained(
"adibvafa/CodonTransformer",
local_files_only=local_files_only,
)
model = BigBirdForMaskedLM.from_pretrained(
"adibvafa/CodonTransformer",
local_files_only=local_files_only,
).to(device)
output = predict_dna_sequence(
protein=protein,
organism=organism,
device=device,
tokenizer=tokenizer,
model=model,
attention_type="original_full",
deterministic=True,
)
print(format_model_output(output))
PY