File size: 2,702 Bytes
53e66de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
#!/usr/bin/env bash
set -euo pipefail

# Batch deterministic inference from a CSV file.
# 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}"

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
    # shellcheck disable=SC1091
    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