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)