| #!/usr/bin/env bash |
| set -euo pipefail |
|
|
| |
| |
|
|
| 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}" |
|
|
| INPUT_CSV="${INPUT_CSV:-${PROJECT_DIR}/scripts/demo/sample_dataset.csv}" |
| OUTPUT_CSV="${OUTPUT_CSV:-${PROJECT_DIR}/outputs/sample_predictions.csv}" |
| OFFLINE="${OFFLINE:-1}" |
|
|
| cd "${PROJECT_DIR}" |
| mkdir -p "$(dirname "${OUTPUT_CSV}")" |
|
|
| export HF_HOME |
| export PYTHONPATH="${PROJECT_DIR}/model:${PYTHONPATH:-}" |
| export INPUT_CSV |
| export OUTPUT_CSV |
| 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 |
| |
| source "$(conda info --base)/etc/profile.d/conda.sh" |
| conda activate "${CONDA_ENV}" |
| fi |
|
|
| python - <<'PY' |
| import os |
|
|
| import pandas as pd |
| import torch |
| from tqdm import tqdm |
| from transformers import AutoTokenizer, BigBirdForMaskedLM |
|
|
| from CodonTransformer.CodonPrediction import predict_dna_sequence |
|
|
| input_csv = os.environ["INPUT_CSV"] |
| output_csv = os.environ["OUTPUT_CSV"] |
| 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}") |
| print(f"Input CSV: {input_csv}") |
|
|
| 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) |
|
|
| dataset = pd.read_csv(input_csv) |
| if "Unnamed: 0" in dataset.columns: |
| dataset = dataset.drop(columns=["Unnamed: 0"]) |
|
|
| required_columns = {"protein_sequence", "organism"} |
| missing = required_columns - set(dataset.columns) |
| if missing: |
| raise ValueError(f"Input CSV is missing required columns: {sorted(missing)}") |
|
|
| dataset["predicted_dna"] = "" |
| for index, row in tqdm(dataset.iterrows(), total=len(dataset), desc="Predicting"): |
| output = predict_dna_sequence( |
| protein=row["protein_sequence"], |
| organism=row["organism"], |
| device=device, |
| tokenizer=tokenizer, |
| model=model, |
| attention_type="original_full", |
| deterministic=True, |
| ) |
| dataset.loc[index, "predicted_dna"] = output.predicted_dna |
|
|
| dataset.to_csv(output_csv, index=False) |
| print(f"Saved predictions to {output_csv}") |
| PY |
|
|