# patched_app.py import os import io import csv import subprocess import pandas as pd import joblib import numpy as np import torch import esm import logging from flask import Flask, request, render_template, send_file, session, redirect, url_for # ---------- Logging Setup ---------- logger = logging.getLogger(__name__) logging.basicConfig( level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s", handlers=[logging.FileHandler("app.log"), logging.StreamHandler()] ) # ---------- Paths ---------- BASE_DIR = os.environ.get("BASE_DIR", os.getcwd()) MODEL_PATH = os.path.join(BASE_DIR, "best_model_LR.sav") ENCODER_PATH = os.path.join(BASE_DIR, "label_encoder.pkl") BLAST_PATH = os.environ.get("BLAST_PATH", "/usr/bin/blastp") # Default for Linux/HF BLAST_DB = os.environ.get("BLAST_DB", os.path.join(BASE_DIR, "pathway_db")) MAPPING_FILE = os.environ.get("MAPPING_FILE", os.path.join(BASE_DIR, "pathway_map.csv")) BITSCORE_THRESHOLD = 80.0 CONFIDENCE_THRESHOLD = 0.7 # threshold for accepting gene family predictions RESULT_CSV = os.path.join(os.getcwd(), "combined_results.csv") # ---------- Load Models ---------- model = None encoder = None esm_model = None batch_converter = None try: if os.path.exists(MODEL_PATH): model = joblib.load(MODEL_PATH) logger.info("LR model loaded successfully.") else: logger.warning(f"Model file not found at {MODEL_PATH}") except Exception as e: logger.exception(f"Failed to load LR model: {e}") try: if os.path.exists(ENCODER_PATH): encoder = joblib.load(ENCODER_PATH) logger.info("Label encoder loaded.") else: logger.warning(f"Encoder file not found at {ENCODER_PATH}") except Exception as e: logger.exception(f"Failed to load label encoder: {e}") try: logger.info("Loading ESM2 model...") esm_model, alphabet = esm.pretrained.load_model_and_alphabet("esm2_t6_8M_UR50D") esm_model.eval() esm_model = esm_model.to("cpu") batch_converter = alphabet.get_batch_converter() logger.info("ESM2 model loaded and set to CPU.") except Exception as e: logger.exception(f"Failed to load ESM model: {e}") # ---------- Functions ---------- def esm2_320_embed(sequence): """Generate ESM2 embeddings for a protein sequence""" if esm_model is None or batch_converter is None: raise RuntimeError("ESM model not available") try: batch_labels, batch_strs, batch_tokens = batch_converter([("seq1", sequence)]) with torch.no_grad(): results = esm_model(batch_tokens, repr_layers=[6], return_contacts=False) token_representations = results["representations"][6] seq_repr = token_representations[0, 1:-1].detach().cpu().numpy() return seq_repr.mean(axis=0) except Exception as e: logger.error(f"Error generating embedding: {e}") raise def run_blast_and_get_dataframe(temp_fasta, blast_output): """Run BLAST search and return results as DataFrame""" try: cmd = [ BLAST_PATH, "-query", temp_fasta, "-db", BLAST_DB, "-out", blast_output, "-outfmt", "6 qseqid sseqid pident length evalue bitscore" ] logger.info(f"Running BLAST: {' '.join(cmd)}") proc = subprocess.run(cmd, capture_output=True, text=True, timeout=300) if proc.returncode != 0: logger.error(f"BLAST error: {proc.stderr}") if not os.path.exists(blast_output) or os.path.getsize(blast_output) == 0: logger.warning("BLAST output is empty") return pd.DataFrame(columns=["qseqid", "sseqid", "pident", "length", "evalue", "bitscore"]) cols = ["qseqid", "sseqid", "pident", "length", "evalue", "bitscore"] df = pd.read_csv(blast_output, sep="\t", names=cols) df["bitscore"] = df["bitscore"].astype(float) logger.info(f"BLAST completed with {len(df)} hits") return df except subprocess.TimeoutExpired: logger.error("BLAST search timed out") return pd.DataFrame(columns=["qseqid", "sseqid", "pident", "length", "evalue", "bitscore"]) except Exception as e: logger.exception(f"Failed to read BLAST output: {e}") return pd.DataFrame(columns=["qseqid", "sseqid", "pident", "length", "evalue", "bitscore"]) # ---------- Flask app ---------- app = Flask(__name__) app.secret_key = os.environ.get('SECRET_KEY', 'your-secret-key-here') latest_predictions = [] @app.after_request def add_cache_control(response): """Add cache control headers to prevent caching""" response.headers["Cache-Control"] = "no-cache, no-store, must-revalidate" response.headers["Pragma"] = "no-cache" response.headers["Expires"] = "0" return response @app.route("/", methods=["GET", "POST"]) def index(): global latest_predictions latest_predictions = [] predictions = [] if request.method == "POST": sequences = [] headers = [] try: # --- Parse uploaded file --- uploaded_file = request.files.get("fasta_file") if uploaded_file and uploaded_file.filename != "": seq = "" header = "" for line in uploaded_file: line = line.decode().strip() if not line: continue if line.startswith(">"): if seq: sequences.append(seq) headers.append(header) seq = "" header = line[1:] # keep full header else: seq += line if seq: sequences.append(seq) headers.append(header) # --- Parse textarea --- sequence_text = request.form.get("sequence_text", "").strip() if sequence_text: seq = "" header = "" for line in sequence_text.splitlines(): line = line.strip() if not line: continue if line.startswith(">"): if seq: sequences.append(seq) headers.append(header) seq = "" header = line[1:] # keep full header else: seq += line if seq: sequences.append(seq) headers.append(header if header else "sequence_from_text") if not sequences: return "⚠️ No valid sequences provided.", 400 logger.info(f"Processing {len(sequences)} sequences") # --- Gene family prediction with confidence check --- try: features = [esm2_320_embed(seq) for seq in sequences] features = np.array(features) gene_preds = [] if model is None: logger.warning("Model not loaded, using 'Unknown' for all predictions") gene_preds = ["Unknown"] * len(sequences) else: probs = model.predict_proba(features) max_probs = probs.max(axis=1) pred_indices = probs.argmax(axis=1) for idx, prob in zip(pred_indices, max_probs): if prob >= CONFIDENCE_THRESHOLD: gene = encoder.inverse_transform([idx])[0] if encoder else str(idx) else: gene = "Unknown" gene_preds.append(gene) logger.info(f"Gene family predictions: {len([p for p in gene_preds if p != 'Unknown'])} confident") except Exception as e: logger.exception("Gene family prediction failed") gene_preds = ["Unknown"] * len(sequences) # --- Save temp fasta with BLAST-safe headers --- temp_fasta = os.path.join(os.getcwd(), "temp_input.fasta") header_map = {} with open(temp_fasta, "w") as f: for h, s in zip(headers, sequences): safe_h = h.replace(" ", "_") header_map[safe_h] = h f.write(f">{safe_h}\n{s}\n") # --- Run BLAST --- blast_output = os.path.join(os.getcwd(), "blast_results.txt") blast_df = run_blast_and_get_dataframe(temp_fasta, blast_output) # --- Load mapping file --- map_df = pd.DataFrame() if os.path.exists(MAPPING_FILE): try: map_df = pd.read_csv(MAPPING_FILE, dtype=str) logger.info(f"Mapping file loaded with {len(map_df)} entries") except Exception as e: logger.exception(f"Failed to read mapping file: {e}") # --- Pick best hits --- combined_results = [] if not blast_df.empty: top_hits_idx = blast_df.groupby("qseqid")["bitscore"].idxmax() top_hits = blast_df.loc[top_hits_idx] else: top_hits = pd.DataFrame(columns=blast_df.columns) # --- Build results --- for i, orig_header in enumerate(headers): safe_h = orig_header.replace(" ", "_") top_hit_row = top_hits[top_hits["qseqid"] == safe_h] if not top_hit_row.empty: row = top_hit_row.iloc[0] top_hit = row["sseqid"] bitscore = float(row["bitscore"]) pathways = "No Pathways Found" if not map_df.empty and "Entry" in map_df.columns and "Pathways" in map_df.columns: paths = map_df.loc[map_df["Entry"] == top_hit, "Pathways"] if (not paths.empty) and bitscore >= BITSCORE_THRESHOLD: pathways = paths.values[0] else: top_hit = "No Hit" bitscore = 0.0 pathways = "No Pathways Found" gene_family = gene_preds[i] if i < len(gene_preds) else "Unknown" combined_results.append([orig_header, top_hit, bitscore, pathways, gene_family]) # --- Save results --- result_df = pd.DataFrame( combined_results, columns=["Query/Header", "Top Hit", "Bitscore", "Pathways", "Predicted Gene Family"] ) result_df.to_csv(RESULT_CSV, index=False) latest_predictions = combined_results # Store results in session and redirect to results page session['predictions'] = combined_results session['result_summary'] = { 'total_sequences': len(sequences), 'successful_predictions': len([p for p in gene_preds if p != "Unknown"]), 'total_pathways_found': len([r for r in combined_results if r[3] != "No Pathways Found"]) } logger.info(f"Processing complete. Results saved to {RESULT_CSV}") return redirect(url_for('results')) except Exception as e: logger.exception("Error during sequence processing") return f"⚠️ An error occurred: {str(e)}", 500 return render_template("index.html", predictions=[]) @app.route("/download_csv") def download_csv(): global latest_predictions if not latest_predictions: return "⚠️ No predictions to download.", 400 try: output = io.StringIO() writer = csv.writer(output) writer.writerow(["Query/Header", "Top Hit", "Bitscore", "Pathways", "Predicted Gene Family"]) writer.writerows(latest_predictions) output.seek(0) logger.info("CSV file generated for download") return send_file( io.BytesIO(output.getvalue().encode()), mimetype="text/csv", as_attachment=True, download_name="combined_predictions.csv" ) except Exception as e: logger.exception("Error generating CSV") return f"⚠️ Error generating file: {str(e)}", 500 @app.route("/results") def results(): predictions = session.get('predictions', []) summary = session.get('result_summary', {}) if not predictions: return redirect(url_for('index')) return render_template("results.html", predictions=predictions, summary=summary) @app.route("/health") def health(): """Health check endpoint for Hugging Face Spaces""" return {"status": "healthy"}, 200 if __name__ == "__main__": host = os.environ.get("HOST", "0.0.0.0") port = int(os.environ.get("PORT", 7860)) debug = os.environ.get("DEBUG", "False").lower() == "true" logger.info(f"Starting Flask app on {host}:{port} (debug={debug})") logger.info(f"Model path: {MODEL_PATH}") logger.info(f"BLAST path: {BLAST_PATH}") app.run(host=host, port=port, debug=debug)