"""
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('', 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()