Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 11 additions & 5 deletions scripts/ANN.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,16 @@
# Import libraries
import os
import numpy as np
from Bio.Alphabet import IUPAC
from keras.optimizers import Adam
# from Bio.Alphabet import IUPAC <-- REMOVED (Deprecated)
from tensorflow.keras.optimizers import Adam # <-- UPDATED for modern Tensorflow
from contextlib import redirect_stdout

# Import custom functions
from utils import one_hot_encoder, create_ann, \
plot_ROC_curve, plot_PR_curve, calc_stat

# Define the protein alphabet manually since Bio.Alphabet is removed
PROTEIN_ALPHABET = "ACDEFGHIKLMNPQRSTVWY"

def ANN_classification(dataset, filename, save_model=False):
"""
Expand Down Expand Up @@ -37,13 +39,16 @@ def ANN_classification(dataset, filename, save_model=False):
X_val = dataset.val.loc[:, 'AASeq'].values

# One hot encode the sequences
X_train = [one_hot_encoder(s=x, alphabet=IUPAC.protein) for x in X_train]
# UPDATED: Using PROTEIN_ALPHABET string instead of IUPAC.protein
X_train = [one_hot_encoder(s=x, alphabet=PROTEIN_ALPHABET) for x in X_train]
X_train = [x.flatten('F') for x in X_train]
X_train = np.asarray(X_train)
X_test = [one_hot_encoder(s=x, alphabet=IUPAC.protein) for x in X_test]

X_test = [one_hot_encoder(s=x, alphabet=PROTEIN_ALPHABET) for x in X_test]
X_test = [x.flatten('F') for x in X_test]
X_test = np.asarray(X_test)
X_val = [one_hot_encoder(s=x, alphabet=IUPAC.protein) for x in X_val]

X_val = [one_hot_encoder(s=x, alphabet=PROTEIN_ALPHABET) for x in X_val]
X_val = [x.flatten('F') for x in X_val]
X_val = np.asarray(X_val)

Expand All @@ -56,6 +61,7 @@ def ANN_classification(dataset, filename, save_model=False):
ANN_classifier = create_ann()

# Compiling the ANN
# Note: Learning rate is small, suitable for fine-tuning or noisy data
ada_optimizer = Adam(learning_rate=0.0001)
ANN_classifier.compile(
optimizer=ada_optimizer, loss='binary_crossentropy',
Expand Down