Download app.py from Priyasi/TransVi: direct link, hf CLI and curl.
- Browser
- Download file 5.25 kB
-
https://huggingface.co/spaces/Priyasi/TransVi/resolve/main/app.py
- Command line
-
hf download hf://spaces/Priyasi/TransVi/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Priyasi/TransVi/resolve/main/app.py
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) | |
| def home(): | |
| return render_template('index.html') | |
| 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) | |
| 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) |