TransVi / app.py
Priyasi's picture
Update app.py
4fdde39 verified
Raw History Blame Contribute Delete
5.25 kB
from flask import Flask, request, render_template
import torch
import re
import numpy as np
from transformers import AutoTokenizer, AutoModelForSequenceClassification
import os
import logging
# Configure logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
app = Flask(__name__)
# Set multiple cache environment variables to ensure they're all pointing to writable locations
cache_dir = "/app/cache"
os.environ["TRANSFORMERS_CACHE"] = cache_dir
os.environ["HF_HOME"] = cache_dir
os.environ["HUGGINGFACE_HUB_CACHE"] = cache_dir
# Name of your finetuned model on Hugging Face
MODEL_NAME = "Priyasi/TransVi_3"
KMER_SIZE = 3 # Must match the k-mer size used in pretraining
# Virus label mapping - 0-indexed to match model output
LABELS = {
1: "SARS-COV-1",
2: "MERS",
3: "SARS-COV-2",
4: "Ebola",
5: "Dengue",
6: "Influenza"
}
# Global variables to store model and tokenizer
tokenizer = None
model = None
def load_model():
"""Load the tokenizer and model"""
global tokenizer, model
try:
logger.info(f"Loading tokenizer and model from {MODEL_NAME}...")
logger.info(f"Using cache directory: {cache_dir}")
logger.info(f"TRANSFORMERS_CACHE: {os.environ.get('TRANSFORMERS_CACHE')}")
logger.info(f"HF_HOME: {os.environ.get('HF_HOME')}")
# Load with explicit cache directory and additional parameters
tokenizer = AutoTokenizer.from_pretrained(
MODEL_NAME,
cache_dir=cache_dir,
local_files_only=False,
use_auth_token=False
)
model = AutoModelForSequenceClassification.from_pretrained(
MODEL_NAME,
cache_dir=cache_dir,
local_files_only=False,
use_auth_token=False
)
model.eval()
logger.info("Model and tokenizer loaded successfully!")
except Exception as e:
logger.error(f"Error loading model: {str(e)}")
raise e
def filter_seq(seq):
"""
Filter out non-ATGC characters and convert U to T for RNA sequences.
"""
seq = str(seq).upper()
seq = seq.replace('U', 'T') # Convert Uracil to Thymine for compatibility
return re.sub(r'[^ATGC]', '', seq)
def create_kmers(sequence, k=3):
"""
Create k-mer representation from the sequence.
Returns a string of k-mers separated by spaces.
"""
if len(sequence) < k:
return ""
kmers = [sequence[i:i+k] for i in range(len(sequence) - k + 1)]
return ' '.join(kmers)
@app.route('/', methods=['GET'])
def home():
return render_template('index.html')
@app.route('/predict', methods=['POST'])
def predict():
global tokenizer, model
# Ensure model is loaded
if tokenizer is None or model is None:
return render_template('index.html',
prediction="Model not loaded. Please try again later.",
error=True)
try:
sequence = request.form.get("sequence")
if not sequence:
return render_template('index.html',
prediction="No sequence provided.",
error=True)
# Preprocess the sequence
filtered_sequence = filter_seq(sequence)
if len(filtered_sequence) < KMER_SIZE:
return render_template('index.html',
prediction=f"Sequence too short after filtering (must be at least {KMER_SIZE} valid characters).",
error=True)
kmers = create_kmers(filtered_sequence, k=KMER_SIZE)
# Tokenize input (using max_length and padding settings similar to training)
inputs = tokenizer(kmers, return_tensors="pt", max_length=512,
padding='max_length', truncation=True)
# Make prediction using the loaded model
with torch.no_grad():
outputs = model(**inputs)
logits = outputs.logits
probabilities = torch.softmax(logits, dim=1).tolist()[0]
predicted_class_idx = int(np.argmax(probabilities))
prediction = LABELS.get(predicted_class_idx, "Unknown")
confidence = probabilities[predicted_class_idx]
return render_template('index.html',
prediction=prediction,
confidence=f"{confidence:.2%}",
error=False)
except Exception as e:
logger.error(f"Error during prediction: {str(e)}")
return render_template('index.html',
prediction="An error occurred during prediction. Please try again.",
error=True)
@app.route('/health', methods=['GET'])
def health_check():
"""Health check endpoint for container orchestration"""
return {"status": "healthy", "model_loaded": model is not None}
if __name__ == '__main__':
# Load model at startup
load_model()
# Get port from environment variable (useful for deployment platforms)
port = int(os.environ.get('PORT', 7860))
# For production deployment
app.run(debug=False, host='0.0.0.0', port=port)