""" BioAgent: AI-Powered Bioinformatics Demo Deployable on HuggingFace Spaces with Claude Scientific Skills integration """ import os import streamlit as st import anthropic import requests import json import time import pandas as pd import plotly.express as px import plotly.graph_objects as go from functools import lru_cache from typing import Optional import numpy as np from io import StringIO import re from dataclasses import dataclass from typing import Tuple # ============================================================================ # PAGE CONFIG & STYLING # ============================================================================ st.set_page_config( page_title="BioAgent | AI Bioinformatics", page_icon="🧬", layout="wide", initial_sidebar_state="expanded" ) # Custom CSS for a distinctive scientific aesthetic st.markdown(""" """, unsafe_allow_html=True) # ============================================================================ # SKILL MANAGEMENT # ============================================================================ SKILL_BASE_URL = "https://raw.githubusercontent.com/K-Dense-AI/claude-scientific-skills/main" SKILL_REGISTRY = { "biopython": { "path": "scientific-packages/biopython", "description": "Sequence analysis, alignment, parsing biological data", "icon": "🧬" }, "rdkit": { "path": "scientific-packages/rdkit", "description": "Molecular properties, SMILES, fingerprints, drug discovery", "icon": "πŸ’Š" }, "scanpy": { "path": "scientific-packages/scanpy", "description": "Single-cell RNA-seq analysis, clustering, visualization", "icon": "πŸ”¬" }, "pubmed": { "path": "scientific-databases/pubmed", "description": "Literature search, citation retrieval", "icon": "πŸ“š" }, "uniprot": { "path": "scientific-databases/uniprot", "description": "Protein sequences, annotations, functional data", "icon": "πŸ”·" }, "chembl": { "path": "scientific-databases/chembl", "description": "Bioactive molecules, drug targets, bioassays", "icon": "βš—οΈ" }, "pubchem": { "path": "scientific-databases/pubchem", "description": "Chemical structures, properties, bioactivities", "icon": "πŸ§ͺ" } } @lru_cache(maxsize=20) def fetch_skill(skill_key: str) -> str: """Fetch a SKILL.md from K-Dense GitHub repo with caching""" if skill_key not in SKILL_REGISTRY: return "" skill_path = SKILL_REGISTRY[skill_key]["path"] url = f"{SKILL_BASE_URL}/{skill_path}/SKILL.md" try: response = requests.get(url, timeout=10) if response.status_code == 200: return response.text except requests.RequestException: pass return "" def get_combined_skills(skill_keys: list[str]) -> str: """Combine multiple skills into a single context block""" skills_content = [] for key in skill_keys: content = fetch_skill(key) if content: skills_content.append(f"### {SKILL_REGISTRY[key]['icon']} {key.upper()} SKILL\n\n{content}") return "\n\n---\n\n".join(skills_content) # ============================================================================ # API INTEGRATIONS (Real Data) # ============================================================================ class PubChemAPI: """Real PubChem API integration""" BASE_URL = "https://pubchem.ncbi.nlm.nih.gov/rest/pug" @staticmethod def search_compound(query: str, max_results: int = 5) -> list[dict]: """Search for compounds by name""" try: # Get CIDs url = f"{PubChemAPI.BASE_URL}/compound/name/{query}/cids/JSON" response = requests.get(url, timeout=10) if response.status_code != 200: return [] cids = response.json().get("IdentifierList", {}).get("CID", [])[:max_results] compounds = [] for cid in cids: # Get properties prop_url = f"{PubChemAPI.BASE_URL}/compound/cid/{cid}/property/MolecularFormula,MolecularWeight,XLogP,TPSA,HBondDonorCount,HBondAcceptorCount/JSON" prop_response = requests.get(prop_url, timeout=10) if prop_response.status_code == 200: props = prop_response.json().get("PropertyTable", {}).get("Properties", [{}])[0] def safe_float(val, default=0.0): """Safely convert to float, handling None and strings""" if val is None: return default try: return float(val) except (ValueError, TypeError): return default def safe_int(val, default=0): """Safely convert to int, handling None and strings""" if val is None: return default try: return int(val) except (ValueError, TypeError): return default compounds.append({ "cid": cid, "formula": props.get("MolecularFormula") or "N/A", "mw": safe_float(props.get("MolecularWeight")), "logp": safe_float(props.get("XLogP")), "tpsa": safe_float(props.get("TPSA")), "hbd": safe_int(props.get("HBondDonorCount")), "hba": safe_int(props.get("HBondAcceptorCount")), "image_url": f"https://pubchem.ncbi.nlm.nih.gov/rest/pug/compound/cid/{cid}/PNG" }) return compounds except Exception as e: st.error(f"PubChem API error: {e}") return [] @staticmethod def get_structure_image(cid: int) -> str: """Get 2D structure image URL""" return f"https://pubchem.ncbi.nlm.nih.gov/rest/pug/compound/cid/{cid}/PNG?image_size=300x300" class PubMedAPI: """Real PubMed/NCBI E-utilities integration""" BASE_URL = "https://eutils.ncbi.nlm.nih.gov/entrez/eutils" @staticmethod def search_literature(query: str, max_results: int = 10) -> list[dict]: """Search PubMed for articles""" try: # Search search_url = f"{PubMedAPI.BASE_URL}/esearch.fcgi" search_params = { "db": "pubmed", "term": query, "retmax": max_results, "retmode": "json", "sort": "relevance" } search_response = requests.get(search_url, params=search_params, timeout=10) if search_response.status_code != 200: return [] pmids = search_response.json().get("esearchresult", {}).get("idlist", []) if not pmids: return [] # Fetch details fetch_url = f"{PubMedAPI.BASE_URL}/esummary.fcgi" fetch_params = { "db": "pubmed", "id": ",".join(pmids), "retmode": "json" } fetch_response = requests.get(fetch_url, params=fetch_params, timeout=10) if fetch_response.status_code != 200: return [] results = fetch_response.json().get("result", {}) articles = [] for pmid in pmids: if pmid in results: article = results[pmid] articles.append({ "pmid": pmid, "title": article.get("title", "N/A"), "authors": ", ".join([a.get("name", "") for a in article.get("authors", [])[:3]]), "journal": article.get("source", "N/A"), "pubdate": article.get("pubdate", "N/A"), "url": f"https://pubmed.ncbi.nlm.nih.gov/{pmid}/" }) return articles except Exception as e: st.error(f"PubMed API error: {e}") return [] class UniProtAPI: """Real UniProt API integration""" BASE_URL = "https://rest.uniprot.org/uniprotkb" @staticmethod def search_protein(query: str, max_results: int = 5) -> list[dict]: """Search UniProt for proteins""" try: url = f"{UniProtAPI.BASE_URL}/search" params = { "query": query, "format": "json", "size": max_results, "fields": "accession,id,protein_name,organism_name,length,sequence" } response = requests.get(url, params=params, timeout=15) if response.status_code != 200: return [] results = response.json().get("results", []) proteins = [] for r in results: protein_name = r.get("proteinDescription", {}).get("recommendedName", {}).get("fullName", {}).get("value", "N/A") proteins.append({ "accession": r.get("primaryAccession", "N/A"), "entry_name": r.get("uniProtkbId", "N/A"), "protein_name": protein_name, "organism": r.get("organism", {}).get("scientificName", "N/A"), "length": r.get("sequence", {}).get("length", 0), "sequence": r.get("sequence", {}).get("value", "")[:100] + "..." }) return proteins except Exception as e: st.error(f"UniProt API error: {e}") return [] # ============================================================================ # CLAUDE INTEGRATION # ============================================================================ def call_claude_with_skills( client: anthropic.Anthropic, user_prompt: str, skill_keys: list[str], task_context: str = "", model: str = "claude-sonnet-4-20250514" ) -> str: """Call Claude API with scientific skills context injected""" skills_content = get_combined_skills(skill_keys) system_prompt = f"""You are BioAgent, an expert AI bioinformatics assistant created by Excelra. You help researchers with drug discovery, multi-omics analysis, and precision medicine workflows. {f"TASK CONTEXT: {task_context}" if task_context else ""} SCIENTIFIC SKILLS REFERENCE: Use the following skill documentation to inform your responses. Follow the patterns, best practices, and code examples described in these skills. {skills_content} RESPONSE GUIDELINES: 1. Be precise and scientifically accurate 2. When providing code, use the libraries and patterns from the skills above 3. Explain your reasoning step-by-step 4. Cite relevant databases or literature when applicable 5. Format output clearly with appropriate sections 6. If generating analysis code, make it executable and well-commented """ if client is None: return "Error: No API client configured. Please enter your Anthropic API key." try: response = client.messages.create( model=model, max_tokens=4096, system=system_prompt, messages=[{"role": "user", "content": user_prompt}] ) return response.content[0].text except anthropic.APIError as e: return f"API Error: {str(e)}" except Exception as e: return f"Error: {str(e)}" # ============================================================================ # MODULE: DRUG DISCOVERY # ============================================================================ def render_drug_discovery_module(client: anthropic.Anthropic): """Drug Discovery Pipeline Module""" st.markdown("### πŸ’Š Drug Discovery Pipeline") st.markdown("Search compounds, analyze molecular properties, and get AI-powered insights.") # Show active skills st.markdown( 'rdkit' 'pubchem' 'chembl', unsafe_allow_html=True ) col1, col2 = st.columns([2, 1]) with col1: compound_query = st.text_input( "Search Compound", placeholder="e.g., aspirin, imatinib, remdesivir", key="drug_search" ) analysis_type = st.selectbox( "Analysis Type", ["Drug-likeness Assessment", "ADMET Prediction", "Target Identification", "Lead Optimization"] ) with col2: st.markdown("**Quick Examples:**") st.caption("Click to copy, then paste:") st.code("aspirin", language=None) st.code("imatinib", language=None) st.code("metformin", language=None) if st.button("πŸ” Analyze Compound", key="analyze_drug") and compound_query: with st.status("Running Drug Discovery Pipeline...", expanded=True) as status: # Step 1: Fetch from PubChem st.write("πŸ“‘ Querying PubChem database...") compounds = PubChemAPI.search_compound(compound_query, max_results=3) time.sleep(0.5) if not compounds: st.error(f"No compounds found for '{compound_query}'") return st.write(f"βœ“ Found {len(compounds)} compound(s)") # Step 2: Display results st.write("πŸ§ͺ Analyzing molecular properties...") time.sleep(0.3) status.update(label="Pipeline complete!", state="complete", expanded=True) # Results display tabs = st.tabs(["πŸ“Š Properties", "πŸ–ΌοΈ Structures", "πŸ€– AI Analysis"]) with tabs[0]: # Properties table df = pd.DataFrame(compounds) df.columns = ["CID", "Formula", "MW (g/mol)", "LogP", "TPSA (Ε²)", "HBD", "HBA", "Image URL"] # Lipinski's Rule of Five assessment st.markdown("#### Lipinski's Rule of Five") for comp in compounds: violations = 0 checks = [] if comp["mw"] <= 500: checks.append("βœ… MW ≀ 500") else: checks.append("❌ MW > 500") violations += 1 if comp["logp"] <= 5: checks.append("βœ… LogP ≀ 5") else: checks.append("❌ LogP > 5") violations += 1 if comp["hbd"] <= 5: checks.append("βœ… HBD ≀ 5") else: checks.append("❌ HBD > 5") violations += 1 if comp["hba"] <= 10: checks.append("βœ… HBA ≀ 10") else: checks.append("❌ HBA > 10") violations += 1 with st.expander(f"CID {comp['cid']} - {comp['formula']} ({violations} violations)"): col1, col2 = st.columns(2) with col1: for check in checks: st.markdown(check) with col2: st.metric("Molecular Weight", f"{comp['mw']:.2f}") st.metric("LogP", f"{comp['logp']:.2f}") st.dataframe(df.drop(columns=["Image URL"]), width='stretch') with tabs[1]: # Structure images cols = st.columns(min(len(compounds), 3)) for i, comp in enumerate(compounds): with cols[i % 3]: st.image(comp["image_url"], caption=f"CID: {comp['cid']}") with tabs[2]: # AI Analysis with st.spinner("πŸ€– Generating AI analysis..."): prompt = f"""Analyze the compound "{compound_query}" for {analysis_type.lower()}. Available data from PubChem: {json.dumps(compounds[0], indent=2)} Provide: 1. Brief compound overview 2. Key molecular features relevant to {analysis_type.lower()} 3. Potential concerns or advantages 4. Recommended next steps for drug development Be concise but thorough. Use scientific terminology appropriately.""" analysis = call_claude_with_skills( client, prompt, ["rdkit", "pubchem", "chembl"], task_context=f"Drug discovery analysis for {compound_query}" ) st.markdown(analysis) # ============================================================================ # MODULE: SEQUENCE ANALYSIS # ============================================================================ def render_sequence_analysis_module(client: anthropic.Anthropic): """Sequence Analysis Module""" st.markdown("### 🧬 Sequence Analysis") st.markdown("Analyze protein/DNA sequences with AI-powered insights.") st.markdown( 'biopython' 'uniprot', unsafe_allow_html=True ) input_type = st.radio( "Input Type", ["Search UniProt", "Paste Sequence"], horizontal=True ) if input_type == "Search UniProt": # Initialize session state for search results and analysis if "uniprot_results" not in st.session_state: st.session_state.uniprot_results = [] if "protein_analysis" not in st.session_state: st.session_state.protein_analysis = {} col1, col2 = st.columns([3, 1]) with col1: protein_query = st.text_input( "Search Protein", placeholder="e.g., BRCA1, insulin, p53 human", key="protein_search" ) with col2: st.markdown("**Examples:**") st.caption("Click to copy, then paste:") st.code("p53 human", language=None) st.code("EGFR human", language=None) if st.button("πŸ” Search UniProt") and protein_query: with st.spinner("Searching UniProt..."): results = UniProtAPI.search_protein(protein_query) st.session_state.uniprot_results = results st.session_state.protein_analysis = {} # Clear previous analyses st.session_state.uniprot_searched = True # Track that search was performed # Display results from session state (persists across reruns) proteins = st.session_state.uniprot_results if proteins: st.markdown("#### Search Results") # Let user select which protein to analyze protein_options = {f"{p['entry_name']} - {p['protein_name'][:40]}": p for p in proteins} selected_name = st.selectbox( "Select protein to analyze:", options=list(protein_options.keys()), key="protein_selector" ) selected_protein = protein_options[selected_name] accession = selected_protein['accession'] # Show protein details st.markdown(f"**Accession:** {accession} | **Organism:** {selected_protein['organism']} | **Length:** {selected_protein['length']} aa") with st.expander("View sequence"): st.code(selected_protein['sequence'], language=None) # Analyze button - outside expander for reliability if accession in st.session_state.protein_analysis: st.markdown("---") st.markdown("#### πŸ€– AI Analysis") st.markdown(st.session_state.protein_analysis[accession]) else: if st.button(f"πŸ”¬ Analyze {accession}", key="analyze_protein_btn"): with st.spinner("Running sequence analysis with Claude..."): prompt = f"""Analyze this protein: Name: {selected_protein['protein_name']} Organism: {selected_protein['organism']} Length: {selected_protein['length']} amino acids Sequence (truncated): {selected_protein['sequence']} Provide: 1. Functional overview of this protein 2. Key domains/motifs likely present 3. Disease associations (if known) 4. Potential drug target considerations Use your knowledge of protein biology and the BioPython skill for analysis approaches.""" analysis = call_claude_with_skills( client, prompt, ["biopython", "uniprot"], task_context="Protein sequence analysis" ) if analysis: st.session_state.protein_analysis[accession] = analysis st.rerun() else: st.error("Failed to generate analysis. Check API key.") elif st.session_state.get("uniprot_searched", False): st.warning("No proteins found. Try a different search term.") else: # Paste Sequence sequence = st.text_area( "Paste Sequence (FASTA or raw)", height=150, placeholder=">sequence_name\nMKTVRQERLKSIVRILERSKEPVSGAQLA..." ) analysis_options = st.multiselect( "Analysis Options", ["Basic Statistics", "Composition Analysis", "Hydrophobicity Profile", "Secondary Structure Prediction"], default=["Basic Statistics"] ) if st.button("🧬 Analyze Sequence") and sequence: # Parse sequence clean_seq = "".join([line for line in sequence.split("\n") if not line.startswith(">")]) clean_seq = re.sub(r'[^A-Za-z]', '', clean_seq).upper() if len(clean_seq) < 10: st.error("Sequence too short. Please enter at least 10 residues.") return with st.spinner("Analyzing sequence..."): # Basic calculations seq_len = len(clean_seq) # Amino acid composition aa_counts = {aa: clean_seq.count(aa) for aa in "ACDEFGHIKLMNPQRSTVWY"} aa_freq = {aa: count/seq_len*100 for aa, count in aa_counts.items()} col1, col2, col3 = st.columns(3) with col1: st.metric("Length", f"{seq_len} aa") with col2: # Approximate MW avg_aa_mw = 110 # Da st.metric("Est. MW", f"{seq_len * avg_aa_mw / 1000:.1f} kDa") with col3: # Basic pI estimate (very rough) acidic = clean_seq.count('D') + clean_seq.count('E') basic = clean_seq.count('K') + clean_seq.count('R') + clean_seq.count('H') if basic > acidic: pi_est = "Basic (>7)" elif acidic > basic: pi_est = "Acidic (<7)" else: pi_est = "~Neutral" st.metric("Est. pI", pi_est) if "Composition Analysis" in analysis_options: st.markdown("#### Amino Acid Composition") fig = px.bar( x=list(aa_freq.keys()), y=list(aa_freq.values()), labels={"x": "Amino Acid", "y": "Frequency (%)"}, color=list(aa_freq.values()), color_continuous_scale="Viridis" ) fig.update_layout( template="plotly_dark", paper_bgcolor="rgba(0,0,0,0)", plot_bgcolor="rgba(0,0,0,0)" ) st.plotly_chart(fig, width='stretch') # AI Analysis prompt = f"""Analyze this protein sequence: Sequence length: {seq_len} amino acids Sequence: {clean_seq[:200]}{'...' if len(clean_seq) > 200 else ''} Amino acid composition highlights: - Most frequent: {max(aa_freq, key=aa_freq.get)} ({aa_freq[max(aa_freq, key=aa_freq.get)]:.1f}%) - Charged residues (D,E,K,R): {sum(aa_freq[aa] for aa in 'DEKR'):.1f}% - Hydrophobic (A,V,I,L,M,F,W): {sum(aa_freq[aa] for aa in 'AVILMFW'):.1f}% Requested analyses: {', '.join(analysis_options)} Provide insights on: 1. Likely protein type/family based on composition 2. Notable sequence features 3. Predicted structural characteristics 4. Suggestions for experimental validation""" analysis = call_claude_with_skills( client, prompt, ["biopython", "uniprot"], task_context="Sequence analysis" ) st.markdown("#### AI Analysis") st.markdown(analysis) # ============================================================================ # MODULE: LITERATURE SEARCH # ============================================================================ def render_literature_module(client: anthropic.Anthropic): """Literature Search Module""" st.markdown("### πŸ“š Literature Intelligence") st.markdown("Search PubMed and get AI-synthesized insights from scientific literature.") st.markdown( 'pubmed', unsafe_allow_html=True ) col1, col2 = st.columns([3, 1]) with col1: search_query = st.text_input( "Search Query", placeholder="e.g., CRISPR cancer therapy, single-cell RNA-seq tumor", key="lit_search" ) with col2: max_results = st.slider("Max Results", 5, 20, 10) search_type = st.radio( "Search Mode", ["Standard Search", "AI-Powered Synthesis"], horizontal=True ) if st.button("πŸ” Search Literature") and search_query: with st.spinner("Searching PubMed..."): articles = PubMedAPI.search_literature(search_query, max_results) if not articles: st.warning("No articles found. Try different search terms.") return st.success(f"Found {len(articles)} articles") if search_type == "Standard Search": for article in articles: with st.expander(f"πŸ“„ {article['title'][:80]}..."): st.markdown(f"**Authors:** {article['authors']}") st.markdown(f"**Journal:** {article['journal']}") st.markdown(f"**Date:** {article['pubdate']}") st.markdown(f"**PMID:** [{article['pmid']}]({article['url']})") else: # AI Synthesis st.markdown("#### πŸ“Š Article Results") # Show articles in compact format df = pd.DataFrame(articles) df = df[["pmid", "title", "journal", "pubdate"]] st.dataframe(df, width='stretch', hide_index=True) st.markdown("#### πŸ€– AI Synthesis") with st.spinner("Synthesizing findings..."): articles_text = "\n\n".join([ f"Title: {a['title']}\nAuthors: {a['authors']}\nJournal: {a['journal']} ({a['pubdate']})" for a in articles ]) prompt = f"""Based on these PubMed search results for "{search_query}": {articles_text} Provide a synthesis that includes: 1. **Key Themes**: What are the main research directions evident from these papers? 2. **Recent Advances**: What new findings or methods appear most significant? 3. **Research Gaps**: Based on the titles, what areas might need more investigation? 4. **Suggested Reading**: Which 2-3 papers seem most foundational for someone new to this topic? Be concise and scientifically rigorous. Note that you only have titles/metadata, not full texts.""" synthesis = call_claude_with_skills( client, prompt, ["pubmed"], task_context=f"Literature synthesis for: {search_query}" ) st.markdown(synthesis) # ============================================================================ # MODULE: SINGLE-CELL ANALYSIS (Demo) # ============================================================================ def render_single_cell_module(client: anthropic.Anthropic): """Single-Cell Analysis Demo Module""" st.markdown("### πŸ”¬ Single-Cell Analysis Demo") st.markdown("Interactive demonstration of single-cell RNA-seq analysis concepts.") st.markdown( 'scanpy', unsafe_allow_html=True ) st.info("πŸ’‘ This module demonstrates single-cell analysis concepts using simulated data. " "For real analysis, connect your own datasets.") # Generate demo data np.random.seed(42) n_cells = 500 n_genes = 50 # Simulate 4 cell clusters cluster_centers = [ np.array([2, 2]), np.array([-2, 2]), np.array([2, -2]), np.array([-2, -2]) ] cell_types = ["T cells", "B cells", "Monocytes", "NK cells"] umap_coords = [] cell_labels = [] for i, center in enumerate(cluster_centers): n_cluster = n_cells // 4 coords = np.random.randn(n_cluster, 2) * 0.5 + center umap_coords.extend(coords) cell_labels.extend([cell_types[i]] * n_cluster) umap_coords = np.array(umap_coords) # Create DataFrame df = pd.DataFrame({ "UMAP1": umap_coords[:, 0], "UMAP2": umap_coords[:, 1], "Cell Type": cell_labels, "n_genes": np.random.randint(500, 3000, len(cell_labels)), "n_counts": np.random.randint(1000, 10000, len(cell_labels)) }) # Visualization options color_by = st.selectbox( "Color by", ["Cell Type", "n_genes", "n_counts"] ) # Create plot if color_by == "Cell Type": fig = px.scatter( df, x="UMAP1", y="UMAP2", color="Cell Type", color_discrete_sequence=px.colors.qualitative.Set2, hover_data=["n_genes", "n_counts"], title="UMAP Projection - Cell Type Clusters" ) else: fig = px.scatter( df, x="UMAP1", y="UMAP2", color=color_by, color_continuous_scale="Viridis", hover_data=["Cell Type", "n_genes", "n_counts"], title=f"UMAP Projection - {color_by}" ) fig.update_layout( template="plotly_dark", paper_bgcolor="rgba(0,0,0,0)", plot_bgcolor="rgba(0,0,0,0)", height=500 ) fig.update_traces(marker=dict(size=6, opacity=0.7)) st.plotly_chart(fig, width='stretch') # Cluster statistics st.markdown("#### Cluster Statistics") stats = df.groupby("Cell Type").agg({ "n_genes": ["mean", "std"], "n_counts": ["mean", "std"], "UMAP1": "count" }).round(1) stats.columns = ["Avg Genes", "Std Genes", "Avg Counts", "Std Counts", "N Cells"] st.dataframe(stats, width='stretch') # AI interpretation if st.button("πŸ€– Get AI Interpretation"): with st.spinner("Generating interpretation..."): prompt = f"""Interpret this single-cell RNA-seq analysis result: Dataset: {n_cells} cells, {n_genes} variable genes Clusters identified: {', '.join(cell_types)} Cluster statistics: {stats.to_string()} Provide: 1. Quality assessment of this dataset 2. Biological interpretation of the cell type distribution 3. Suggested downstream analyses 4. Potential validation experiments This is simulated demonstration data, but analyze as if it were real.""" interpretation = call_claude_with_skills( client, prompt, ["scanpy"], task_context="Single-cell RNA-seq analysis interpretation" ) st.markdown(interpretation) # ============================================================================ # MODULE: PATIENT STRATIFICATION (Multi-Omics VAE) # ============================================================================ class MultiOmicsStratificationEngine: """ Multi-Omics Patient Stratification Engine Based on UCB Multi-Omics Stratification Agent architecture. Features: - Simulated multi-omics data with causal relationships - Gated encoders with attention mechanism - UMAP visualization of patient clusters - Survival analysis simulation - Biomarker discovery """ def __init__(self, n_samples: int = 300, n_genes: int = 500, seed: int = 2025): self.n_samples = n_samples self.n_genes = n_genes self.seed = seed self.data = None self.results = None def simulate_data(self) -> dict: """Simulate multi-omics data with biological relationships""" np.random.seed(self.seed) # Patient subtypes: 0=Non-Responder, 1=Responder, 2=High-Risk subtypes = np.random.choice([0, 1, 2], self.n_samples, p=[0.4, 0.4, 0.2]) # CNV: mostly diploid (0), some gains (+1) and losses (-1) cnv = np.random.choice([-1, 0, 1], size=(self.n_samples, self.n_genes), p=[0.05, 0.9, 0.05]) # Mutations: sparse binary with subtype-specific patterns mutations = np.zeros((self.n_samples, self.n_genes)) for i in range(self.n_samples): base_muts = np.random.random(self.n_genes) < 0.02 if subtypes[i] == 0: # Non-responder: resistance mutations base_muts[0:50] = np.random.random(50) < 0.3 elif subtypes[i] == 1: # Responder: sensitivity mutations base_muts[50:100] = np.random.random(50) < 0.3 else: # High-risk: inflammatory signature base_muts[100:150] = np.random.random(50) < 0.25 mutations[i] = base_muts.astype(int) # Gene Expression: baseline + CNV effect + subtype signature + noise gex = np.random.randn(self.n_samples, self.n_genes) * 0.5 gex += cnv * 0.8 # CNV affects expression # Subtype-specific expression signatures for i in range(self.n_samples): if subtypes[i] == 0: gex[i, 0:50] += 1.5 # Resistance genes up gex[i, 200:250] -= 1.0 # Immune genes down elif subtypes[i] == 1: gex[i, 50:100] += 1.2 # Sensitivity pathway gex[i, 200:250] += 1.5 # Immune activation else: gex[i, 100:150] += 2.0 # Inflammatory genes gex[i, 150:200] += 1.0 # Cytokine storm signature # Standardize GEX gex = (gex - gex.mean(axis=0)) / (gex.std(axis=0) + 1e-8) # Patient metadata patient_ids = [f"PT{i:04d}" for i in range(self.n_samples)] ages = np.random.normal(55, 12, self.n_samples).astype(int) ages = np.clip(ages, 25, 85) self.data = { 'gex': gex, 'mutations': mutations, 'cnv': cnv, 'subtypes': subtypes, 'patient_ids': patient_ids, 'ages': ages, 'n_samples': self.n_samples, 'n_genes': self.n_genes } return self.data def run_stratification(self) -> dict: """ Run patient stratification using PCA-based approach. (Simplified from full VAE for demo speed - real implementation uses PyTorch VAE) """ if self.data is None: self.simulate_data() from sklearn.decomposition import PCA from sklearn.preprocessing import StandardScaler from sklearn.ensemble import RandomForestClassifier from sklearn.model_selection import train_test_split # Combine modalities with learned weights (simulated attention) # In full model, this comes from attention mechanism gex_weight = 0.55 mut_weight = 0.30 cnv_weight = 0.15 # Create combined feature matrix combined = np.hstack([ self.data['gex'] * gex_weight, self.data['mutations'] * mut_weight, self.data['cnv'] * cnv_weight ]) # PCA for latent space (VAE substitute for demo) pca = PCA(n_components=32) latent = pca.fit_transform(combined) # UMAP for visualization from sklearn.manifold import TSNE # Using t-SNE as fallback if UMAP not available try: import umap.umap_ as umap_lib reducer = umap_lib.UMAP(n_neighbors=15, min_dist=0.1, random_state=42) embedding = reducer.fit_transform(latent) except ImportError: tsne = TSNE(n_components=2, random_state=42, perplexity=30) embedding = tsne.fit_transform(latent) # Train classifier X_train, X_test, y_train, y_test, idx_train, idx_test = train_test_split( latent, self.data['subtypes'], np.arange(self.n_samples), test_size=0.2, stratify=self.data['subtypes'], random_state=42 ) clf = RandomForestClassifier(n_estimators=100, random_state=42) clf.fit(X_train, y_train) # Predictions predictions = clf.predict(latent) probabilities = clf.predict_proba(latent) test_predictions = clf.predict(X_test) test_accuracy = (test_predictions == y_test).mean() # Biomarker discovery: feature importance from GEX gex_importance = np.abs(pca.components_[:, :self.n_genes]).mean(axis=0) top_biomarkers = np.argsort(gex_importance)[-20:] # Simulate survival times survival_times, events = self._simulate_survival(predictions) self.results = { 'embedding': embedding, 'latent': latent, 'predictions': predictions, 'probabilities': probabilities, 'true_labels': self.data['subtypes'], 'test_accuracy': test_accuracy, 'attention_weights': np.array([gex_weight, mut_weight, cnv_weight]), 'top_biomarkers': top_biomarkers, 'biomarker_importance': gex_importance, 'survival_times': survival_times, 'survival_events': events, 'idx_test': idx_test } return self.results def _simulate_survival(self, subtypes: np.ndarray) -> tuple: """Simulate survival times based on patient subtypes""" np.random.seed(self.seed + 1) times = [] events = [] for s in subtypes: if s == 0: # Non-responder t = np.random.exponential(6) elif s == 1: # Responder t = np.random.exponential(24) else: # High-risk t = np.random.exponential(10) # Censor at 24 months if t > 24: times.append(24) events.append(0) else: times.append(t) events.append(1) return np.array(times), np.array(events) def render_patient_stratification_module(client: anthropic.Anthropic): """Patient Stratification Module - Multi-Omics Analysis""" st.markdown("### 🎯 Patient Stratification") st.markdown("Multi-omics integration for clinical trial patient selection and risk stratification.") st.markdown( 'scanpy' 'biopython' 'pytorch', unsafe_allow_html=True ) # Configuration st.markdown("#### βš™οΈ Configuration") col1, col2, col3 = st.columns(3) with col1: n_patients = st.slider("Number of Patients", 100, 500, 300, step=50) with col2: n_genes = st.slider("Number of Genes", 200, 1000, 500, step=100) with col3: seed = st.number_input("Random Seed", value=2025, min_value=1) data_source = st.radio( "Data Source", ["🎲 Simulated Demo Data", "πŸ“ Upload Real Data (Coming Soon)"], horizontal=True ) if "Upload" in data_source: st.info("πŸ“€ File upload for real multi-omics data will be available in production. " "Supported formats: CSV (GEX matrix), VCF (mutations), SEG (CNV).") return if st.button("πŸš€ Run Stratification Pipeline", type="primary"): # Initialize engine engine = MultiOmicsStratificationEngine(n_samples=n_patients, n_genes=n_genes, seed=seed) with st.status("Running Multi-Omics Stratification Pipeline...", expanded=True) as status: st.write("🧬 Simulating multi-omics data...") data = engine.simulate_data() time.sleep(0.3) st.write(f" βœ“ Generated {n_patients} patients Γ— {n_genes} genes") st.write(f" βœ“ Modalities: GEX, Mutations, CNV") st.write("πŸ”§ Running stratification model...") results = engine.run_stratification() time.sleep(0.3) st.write(f" βœ“ Test accuracy: {results['test_accuracy']:.1%}") st.write("πŸ“Š Generating visualizations...") time.sleep(0.2) status.update(label="βœ… Pipeline Complete!", state="complete") # Store in session state st.session_state['strat_engine'] = engine st.session_state['strat_results'] = results st.session_state['strat_data'] = data # Display results if available if 'strat_results' in st.session_state: results = st.session_state['strat_results'] data = st.session_state['strat_data'] engine = st.session_state['strat_engine'] # Metrics row st.markdown("---") st.markdown("#### πŸ“ˆ Key Metrics") col1, col2, col3, col4 = st.columns(4) with col1: st.metric("Patients", data['n_samples']) with col2: st.metric("Test Accuracy", f"{results['test_accuracy']:.1%}") with col3: n_responders = (results['predictions'] == 1).sum() st.metric("Predicted Responders", f"{n_responders} ({n_responders/data['n_samples']:.0%})") with col4: n_highrisk = (results['predictions'] == 2).sum() st.metric("High-Risk Flagged", n_highrisk) # Tabs for visualizations tabs = st.tabs([ "πŸ—ΊοΈ Patient Map", "βš–οΈ Modality Attention", "πŸ“‰ Survival Analysis", "🧬 Biomarkers", "πŸ€– AI Insights" ]) # Tab 1: Patient Stratification Map (UMAP) with tabs[0]: st.markdown("##### Patient Stratification Map") st.caption("UMAP projection of learned patient representations") labels = ['Non-Responder', 'Responder', 'High-Risk'] colors = ['#808080', '#2ecc71', '#e74c3c'] df_umap = pd.DataFrame({ 'UMAP1': results['embedding'][:, 0], 'UMAP2': results['embedding'][:, 1], 'True Label': [labels[s] for s in data['subtypes']], 'Predicted': [labels[p] for p in results['predictions']], 'Patient ID': data['patient_ids'], 'Age': data['ages'], 'Responder Prob': results['probabilities'][:, 1] }) color_by = st.radio("Color by:", ["Predicted", "True Label", "Responder Prob"], horizontal=True) if color_by == "Responder Prob": fig = px.scatter( df_umap, x='UMAP1', y='UMAP2', color='Responder Prob', color_continuous_scale='RdYlGn', hover_data=['Patient ID', 'Age', 'True Label', 'Predicted'], title='Patient Stratification Map' ) else: fig = px.scatter( df_umap, x='UMAP1', y='UMAP2', color=color_by, color_discrete_map={ 'Non-Responder': '#808080', 'Responder': '#2ecc71', 'High-Risk': '#e74c3c' }, hover_data=['Patient ID', 'Age', 'Responder Prob'], title='Patient Stratification Map' ) fig.update_layout( template="plotly_dark", paper_bgcolor="rgba(0,0,0,0)", plot_bgcolor="rgba(0,0,0,0)", height=500 ) fig.update_traces(marker=dict(size=8, opacity=0.7, line=dict(width=0.5, color='white'))) st.plotly_chart(fig, width='stretch') # Confusion matrix st.markdown("##### Classification Performance") from sklearn.metrics import confusion_matrix, classification_report cm = confusion_matrix(data['subtypes'], results['predictions']) fig_cm = px.imshow( cm, labels=dict(x="Predicted", y="True", color="Count"), x=labels, y=labels, color_continuous_scale='Blues', text_auto=True ) fig_cm.update_layout( template="plotly_dark", paper_bgcolor="rgba(0,0,0,0)", height=350 ) st.plotly_chart(fig_cm, width='stretch') # Tab 2: Modality Attention with tabs[1]: st.markdown("##### Modality Importance") st.caption("Which omics data layer contributed most to stratification decisions?") attn = results['attention_weights'] fig_attn = go.Figure(data=[ go.Bar( x=['Gene Expression', 'Mutations', 'CNV'], y=attn, marker_color=['#4cc9f0', '#f72585', '#7209b7'], text=[f'{v:.0%}' for v in attn], textposition='outside' ) ]) fig_attn.update_layout( template="plotly_dark", paper_bgcolor="rgba(0,0,0,0)", plot_bgcolor="rgba(0,0,0,0)", yaxis_title="Attention Weight", yaxis_range=[0, 1], height=400, title="AI Attention: Which Data Drove the Decision?" ) st.plotly_chart(fig_attn, width='stretch') st.markdown(f""" **Interpretation:** - **Gene Expression** ({attn[0]:.0%}): Primary driver - captures pathway activity and treatment response signatures - **Mutations** ({attn[1]:.0%}): Key for identifying resistance/sensitivity markers - **CNV** ({attn[2]:.0%}): Supporting signal for dosage-sensitive genes """) # Tab 3: Survival Analysis with tabs[2]: st.markdown("##### Simulated Clinical Trial: Survival Analysis") st.caption("Kaplan-Meier curves showing progression-free survival by predicted group") fig_surv = go.Figure() for group_idx, (group_name, color) in enumerate(zip(labels, colors)): mask = results['predictions'] == group_idx g_times = results['survival_times'][mask] g_events = results['survival_events'][mask] # Sort by time sort_idx = np.argsort(g_times) g_times = g_times[sort_idx] g_events = g_events[sort_idx] # Kaplan-Meier calculation survival_prob = 1.0 probs_list = [1.0] ts = [0.0] n_at_risk = len(g_times) for i, t in enumerate(g_times): if g_events[i] == 1: survival_prob *= (1 - 1/max(n_at_risk, 1)) probs_list.append(survival_prob) ts.append(t) n_at_risk -= 1 fig_surv.add_trace(go.Scatter( x=ts, y=probs_list, mode='lines', name=f'{group_name} (n={mask.sum()})', line=dict(color=color, width=2, shape='hv') )) fig_surv.update_layout( template="plotly_dark", paper_bgcolor="rgba(0,0,0,0)", plot_bgcolor="rgba(0,0,0,0)", xaxis_title="Time (months)", yaxis_title="Progression-Free Survival", yaxis_range=[0, 1.05], xaxis_range=[0, 25], height=450, title="Kaplan-Meier Survival Curves by Predicted Subtype", legend=dict(x=0.7, y=0.95) ) st.plotly_chart(fig_surv, width='stretch') # Median survival times st.markdown("**Median Progression-Free Survival:**") col1, col2, col3 = st.columns(3) for i, (col, label, color) in enumerate(zip([col1, col2, col3], labels, colors)): mask = results['predictions'] == i median_pfs = np.median(results['survival_times'][mask]) with col: st.metric(label, f"{median_pfs:.1f} months") # Tab 4: Biomarkers with tabs[3]: st.markdown("##### Biomarker Discovery") st.caption("Genes most predictive of treatment response") top_genes = results['top_biomarkers'][-15:] importance = results['biomarker_importance'][top_genes] fig_bio = go.Figure(data=[ go.Bar( y=[f"Gene_{g:04d}" for g in top_genes], x=importance, orientation='h', marker_color='#2ecc71' ) ]) fig_bio.update_layout( template="plotly_dark", paper_bgcolor="rgba(0,0,0,0)", plot_bgcolor="rgba(0,0,0,0)", xaxis_title="Importance Score", height=450, title="Top 15 Predictive Biomarkers" ) st.plotly_chart(fig_bio, width='stretch') # Gene distribution by subtype st.markdown("##### Top Biomarker Expression by Subtype") top_gene = top_genes[-1] df_gene = pd.DataFrame({ 'Expression': data['gex'][:, top_gene], 'Subtype': [labels[s] for s in data['subtypes']] }) fig_box = px.box( df_gene, x='Subtype', y='Expression', color='Subtype', color_discrete_map={ 'Non-Responder': '#808080', 'Responder': '#2ecc71', 'High-Risk': '#e74c3c' }, title=f"Gene_{top_gene:04d} Expression by Subtype" ) fig_box.update_layout( template="plotly_dark", paper_bgcolor="rgba(0,0,0,0)", plot_bgcolor="rgba(0,0,0,0)", showlegend=False, height=350 ) st.plotly_chart(fig_box, width='stretch') # Tab 5: AI Insights with tabs[4]: st.markdown("##### πŸ€– Clinical Insights") if st.button("Generate Analysis", key="ai_strat"): with st.spinner("Analyzing the stratification results..."): # Build context subtype_counts = pd.Series(results['predictions']).value_counts().sort_index() top_genes = results['top_biomarkers'] prompt = f"""Analyze these patient stratification results from a multi-omics clinical trial analysis: **Dataset Summary:** - Total patients: {data['n_samples']} - Features: {data['n_genes']} genes across 3 modalities (GEX, Mutations, CNV) - Test accuracy: {results['test_accuracy']:.1%} **Stratification Results:** - Non-Responders: {subtype_counts.get(0, 0)} patients ({subtype_counts.get(0, 0)/data['n_samples']:.0%}) - Responders: {subtype_counts.get(1, 0)} patients ({subtype_counts.get(1, 0)/data['n_samples']:.0%}) - High-Risk: {subtype_counts.get(2, 0)} patients ({subtype_counts.get(2, 0)/data['n_samples']:.0%}) **Modality Importance:** - Gene Expression: {results['attention_weights'][0]:.0%} - Mutations: {results['attention_weights'][1]:.0%} - CNV: {results['attention_weights'][2]:.0%} **Survival Analysis:** - Responders show significantly longer progression-free survival - High-risk patients flagged for potential adverse events **Top Biomarker Genes:** Gene_{top_genes[-1]:04d}, Gene_{top_genes[-2]:04d}, Gene_{top_genes[-3]:04d} Provide clinical insights covering: 1. **Patient Selection**: Recommendations for trial enrollment criteria 2. **Risk Stratification**: How to use these results for patient monitoring 3. **Biomarker Validation**: Next steps for validating discovered markers 4. **Trial Design**: How this stratification could improve trial efficiency 5. **Regulatory Considerations**: Key points for FDA submission Be specific and actionable. This is for a precision medicine clinical trial context.""" analysis = call_claude_with_skills( client, prompt, ["scanpy", "biopython"], task_context="Multi-omics patient stratification for clinical trials" ) st.markdown(analysis) else: st.info("Click 'Generate AI Analysis' to get Claude's interpretation of the stratification results.") # ============================================================================ # MAIN APPLICATION # ============================================================================ def main(): # Header st.markdown('

🧬 BioAgent

', unsafe_allow_html=True) st.markdown('

AI-Powered Bioinformatics β€’ Drug Discovery β€’ Multi-Omics Analysis

', unsafe_allow_html=True) # Sidebar with st.sidebar: st.markdown("## βš™οΈ Configuration") # Try to get API key from environment variable (HuggingFace Spaces Docker) # or from Streamlit secrets, or let user input manually api_key_from_env = os.environ.get("ANTHROPIC_API_KEY") api_key_from_secrets = None try: api_key_from_secrets = st.secrets.get("ANTHROPIC_API_KEY") except Exception: pass # No secrets file exists if api_key_from_env: api_key = api_key_from_env st.success("βœ“ API key loaded from environment") elif api_key_from_secrets: api_key = api_key_from_secrets st.success("βœ“ API key loaded from secrets") else: api_key = st.text_input( "Anthropic API Key", type="password", help="Enter your Anthropic API key to enable AI features" ) if api_key: st.success("βœ“ API key configured") if api_key: client = anthropic.Anthropic(api_key=api_key) else: st.warning("⚠️ Enter API key for AI features") client = None st.markdown("---") st.markdown("## πŸ§ͺ Analysis Modules") module = st.radio( "Select Module", [ "πŸ’Š Drug Discovery", "🧬 Sequence Analysis", "🎯 Patient Stratification", "πŸ“š Literature Search", "πŸ”¬ Single-Cell Demo" ], label_visibility="collapsed" ) st.markdown("---") st.markdown("## πŸ“– About") st.markdown(""" **BioAgent** demonstrates Agentic bioinformatics workflows using: - Scientific Skills - Multi-Omics Data Integration - Real-time database APIs - Interactive visualizations - Anthropic Developed by Aravind [Risentia].""") st.markdown("---") # Skills status with st.expander("πŸ”§ Loaded Skills"): for key, info in SKILL_REGISTRY.items(): st.markdown(f"{info['icon']} **{key}**") st.caption(info['description']) # Main content if not client: st.markdown("""

πŸ‘‹ Welcome to BioAgent

Enter your Anthropic API key in the sidebar to enable AI-powered analysis features.

You can still explore the data integrations (PubChem, PubMed, UniProt) without an API key.

""", unsafe_allow_html=True) # Show demo even without API key st.markdown("### 🎯 Quick Demo: PubChem Search") demo_compound = st.text_input("Try searching a compound:", value="caffeine") if st.button("Search PubChem"): with st.spinner("Fetching from PubChem..."): compounds = PubChemAPI.search_compound(demo_compound) if compounds: for c in compounds: col1, col2 = st.columns([1, 2]) with col1: st.image(c["image_url"]) with col2: st.markdown(f"**CID:** {c['cid']}") st.markdown(f"**Formula:** {c['formula']}") st.markdown(f"**MW:** {c['mw']:.2f} g/mol") st.markdown(f"**LogP:** {c['logp']:.2f}") return # Render selected module if "Drug Discovery" in module: render_drug_discovery_module(client) elif "Sequence Analysis" in module: render_sequence_analysis_module(client) elif "Patient Stratification" in module: render_patient_stratification_module(client) elif "Literature" in module: render_literature_module(client) elif "Single-Cell" in module: render_single_cell_module(client) if __name__ == "__main__": main()