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