| |
| """ |
| 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.") |
|
|
| |
| 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" |
| ) |
|
|