File size: 5,914 Bytes
4f18e56 d4fdb7e ca8448a d4fdb7e 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 eaa8a05 4f18e56 ca8448a 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 a6ff6f3 4f18e56 eaa8a05 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 60537f6 d4fdb7e ca8448a d4fdb7e 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 d4fdb7e 4f18e56 ca8448a b0dd450 eaa8a05 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 | #!/usr/bin/env python3
"""
Flexynesis Tissue VAE — Web Application (compressed, real-model inference)
=========================================================================
Upload a bulk RNA-seq gene-expression matrix → UBERON tissue-of-origin
predictions + 121-dim latent embeddings.
Runs the trained supervised VAE via a 53 MB int8 TorchScript bundle
(vae_tissue_int8.torchscript.pt) — exact inference, no flexynesis dependency,
fits the Hugging Face free tier. Trained on 118,263 tissue-curated samples
(TCGA, GTEx, ARCHS4), 42 UBERON tissues, 94.9% balanced accuracy.
Author: Amit Pande, MDC Berlin/BIMSB
"""
import json
from pathlib import Path
from collections import Counter
import numpy as np
import pandas as pd
import streamlit as st
import torch
import joblib
MODEL_DIR = Path(".")
st.set_page_config(page_title="Flexynesis Tissue VAE", page_icon="🧬", layout="wide")
st.title("🧬 Flexynesis Tissue VAE")
st.markdown(
"Upload a bulk RNA-seq gene expression matrix to classify tissue-of-origin "
"using a supervised VAE trained on **118,263 tissue-curated samples** from TCGA, GTEx, and ARCHS4 "
"across **42 UBERON tissue categories** (94.9% balanced accuracy, 121-dim latent space)."
)
@st.cache_resource
def load_model():
model = torch.jit.load(str(MODEL_DIR / "vae_tissue_int8.torchscript.pt"))
model.eval()
art = joblib.load(MODEL_DIR / "vae_tissue.artifacts.joblib")
gene_list = list(art["feature_lists"]["gex"])
scaler = art["transforms"]["gex"]
label_mapping = {int(k): v for k, v in
json.loads((MODEL_DIR / "label_mapping.json").read_text()).items()}
return model, gene_list, scaler, label_mapping
try:
model, gene_list, scaler, label_mapping = load_model()
st.success(
f"✅ Model loaded · {len(gene_list):,} genes · "
f"{len(label_mapping)} tissue classes · exact VAE inference (int8 TorchScript)"
)
except Exception as e:
st.error(f"Model load error: {e}")
st.info("Ensure vae_tissue_int8.torchscript.pt, vae_tissue.artifacts.joblib, "
"and label_mapping.json are in the repository root.")
st.stop()
def orient_matrix(df):
in_cols = len(set(df.columns) & set(gene_list))
in_rows = len(set(df.index) & set(gene_list))
if in_rows > in_cols:
df = df.T
return df
st.markdown("---")
col1, col2 = st.columns([2, 1])
with col1:
uploaded = st.file_uploader(
"Upload CSV/TSV (genes × samples or samples × genes)",
type=["csv", "tsv", "txt"],
help="HGNC gene symbols. Log2-transformed expression values.")
with col2:
st.markdown(f"""
**Expected input:**
- HGNC gene symbols
- {len(gene_list):,} genes used by model
- Log2-transformed expression
""")
if uploaded:
sep = "\t" if uploaded.name.endswith((".tsv", ".txt")) else ","
df = pd.read_csv(uploaded, index_col=0, sep=sep)
st.write(f"**Uploaded:** {df.shape[0]:,} × {df.shape[1]:,}")
st.dataframe(df.iloc[:5, :5], use_container_width=True)
if st.button("🚀 Classify Tissues", type="primary"):
with st.spinner("Running the VAE..."):
df = orient_matrix(df)
overlap = len(set(df.columns) & set(gene_list))
st.write(f"Gene overlap: **{overlap:,}/{len(gene_list):,}** "
f"({100*overlap/len(gene_list):.1f}%)")
if overlap < 1000:
st.warning("Low gene overlap — results may be unreliable.")
# --- Input scale check: model expects log2(TPM) ---
common_for_scale = [g for g in gene_list if g in df.columns]
if common_for_scale:
vals = df[common_for_scale].to_numpy(dtype=float)
vmax = float(np.nanmax(vals))
p99 = float(np.nanpercentile(vals, 99))
st.caption(f"Input value range: 99th pct = {p99:.1f}, max = {vmax:.1f} "
"(log2(TPM) is typically 0–20).")
if vmax > 30 or p99 > 25:
st.warning(
"⚠️ Input values look larger than expected for **log2 scale**. "
"This model expects **log2-transformed expression** (e.g. log2(TPM+1)). "
"Raw TPM/counts will give unreliable predictions — please log2-transform first."
)
aligned = pd.DataFrame(0.0, index=df.index, columns=gene_list)
common = [g for g in gene_list if g in df.columns]
aligned[common] = df[common].values
aligned = aligned.fillna(0)
X = torch.tensor(scaler.transform(aligned.values), dtype=torch.float32)
with torch.no_grad():
logits = model(X)
nan_idx = {v: k for k, v in label_mapping.items()}.get("nan")
if nan_idx is not None:
logits[:, nan_idx] = -1e9
probs = torch.softmax(logits, dim=1)
pred_idx = logits.argmax(dim=1)
pred_labels = [label_mapping[int(i)] for i in pred_idx]
conf = probs.max(dim=1).values.numpy()
st.markdown("---")
st.subheader("📊 Results")
results = pd.DataFrame({
"Sample": df.index,
"Tissue": pred_labels,
"Confidence": [f"{c:.1%}" for c in conf],
})
st.dataframe(results, use_container_width=True, height=400)
st.subheader("Tissue Distribution")
st.bar_chart(pd.Series(pred_labels).value_counts())
st.download_button("📥 Predictions (CSV)", results.to_csv(index=False),
"flexynesis_predictions.csv", "text/csv")
st.markdown("---")
st.caption(
"Flexynesis Tissue VAE (v3) · Akalin Lab, MDC Berlin/BIMSB · "
"118,263 samples · 42 UBERON tissues · 94.9% balanced accuracy · "
"github.com/BIMSBbioinfo/flexynesis_tissue_vae_manuscript"
)
|