Team Ai
Apppublic

softwareDevelopment/Heterosis

sourceHugging Faceapache-2.0updated 7mo agoView on Hugging Face
0likes
app.py338 linesDownload Raw Back to root
1# -*- coding: utf-8 -*-2"""3Created on Wed Mar 18 10:29:52 20264 5@author: Ashmi6"""7 8# -*- coding: utf-8 -*-9"""10Heterosis Prediction - Gradio App for HuggingFace Hub11"""12 13import gradio as gr14import pandas as pd15import numpy as np16import os17import random18import matplotlib19matplotlib.use('Agg')20import matplotlib.pyplot as plt21import seaborn as sns22import tempfile23 24from sklearn.model_selection import KFold25from sklearn.preprocessing import StandardScaler26from sklearn.metrics import mean_squared_error, r2_score27from scipy.stats import pearsonr28from tensorflow.keras.models import Sequential29from tensorflow.keras.layers import Dense, Dropout, BatchNormalization, LeakyReLU30from tensorflow.keras import regularizers31from tensorflow.keras.optimizers import Adam32import tensorflow as tf33 34# ── Reproducibility ──────────────────────────────────────────────────────────35seed_value = 4236os.environ['PYTHONHASHSEED'] = str(seed_value)37os.environ['TF_DETERMINISTIC_OPS'] = '1'38random.seed(seed_value)39np.random.seed(seed_value)40tf.random.set_seed(seed_value)41 42 43# ── Core Model ────────────────────────────────────────────────────────────────44def build_model(input_dim, l1_reg=0.001, l2_reg=0.001, dropout_rate=0.2):45    model = Sequential()46    for units in [512, 512, 256, 128, 64, 32, 16, 8]:47        model.add(Dense(units, input_shape=(input_dim,) if units == 512 else [],48                        kernel_initializer='he_normal',49                        kernel_regularizer=regularizers.l1_l2(l1=l1_reg, l2=l2_reg)))50        model.add(BatchNormalization())51        model.add(Dropout(dropout_rate))52        model.add(LeakyReLU(alpha=0.1))53    model.add(Dense(1, activation="relu"))54    model.compile(loss='mse', optimizer=Adam(), metrics=['mae'])55    return model56 57 58def heterosis_predict(59    num_F1, num_P,60    train_pheno_file, train_geno_file, parent_geno_file,61    test_geno_file=None,62    l1_reg=0.001, l2_reg=0.001, dropout_rate=0.2,63    k_folds=5, epochs=100, batch_size=16,64    progress=gr.Progress(track_tqdm=True)65):66    logs = []67 68    def log(msg):69        logs.append(msg)70        return "\n".join(logs)71 72    # ── Load Data ──────────────────────────────────────────────────────────73    try:74        train_pheno = pd.read_csv(train_pheno_file.name)75        train_geno  = pd.read_csv(train_geno_file.name,  index_col=0)76        parent_geno_df = pd.read_csv(parent_geno_file.name, index_col=0)77    except Exception as e:78        return None, None, None, None, None, f"❌ File loading error: {e}"79 80    parent_index = parent_geno_df.index81    train_Y = np.array(pd.to_numeric(train_pheno.iloc[:, 1], errors="coerce"))82    train_X = train_geno.to_numpy(dtype=np.float32)83    parent_X_raw = parent_geno_df.to_numpy(dtype=np.float32)84 85    scaler = StandardScaler()86    train_X = scaler.fit_transform(train_X)87    parent_X = scaler.transform(parent_X_raw)88 89    # ── K-Fold CV ──────────────────────────────────────────────────────────90    kfold = KFold(n_splits=k_folds, shuffle=True, random_state=60)91    fold_metrics = []92    log(f"🔄 Starting {k_folds}-Fold Cross-Validation  (epochs={epochs}, batch={batch_size})\n")93 94    for fold, (train_idx, val_idx) in enumerate(kfold.split(train_X)):95        progress((fold) / k_folds, desc=f"Fold {fold+1}/{k_folds}")96        log(f"  ▶ Fold {fold+1}/{k_folds} ...")97        model = build_model(train_X.shape[1], l1_reg, l2_reg, dropout_rate)98        model.fit(train_X[train_idx], train_Y[train_idx],99                  validation_data=(train_X[val_idx], train_Y[val_idx]),100                  epochs=epochs, batch_size=batch_size, verbose=0)101 102        val_pred = model.predict(train_X[val_idx], verbose=0).flatten()103        mse  = mean_squared_error(train_Y[val_idx], val_pred)104        rmse = np.sqrt(mse)105        r2   = r2_score(train_Y[val_idx], val_pred)106        pearson_corr, _ = pearsonr(train_Y[val_idx], val_pred)107        fold_metrics.append((mse, rmse, r2, pearson_corr))108        log(f"    MSE={mse:.4f}  RMSE={rmse:.4f}  R²={r2:.4f}  Pearson r={pearson_corr:.4f}")109 110    avg_mse     = np.mean([m[0] for m in fold_metrics])111    avg_rmse    = np.mean([m[1] for m in fold_metrics])112    avg_r2      = np.mean([m[2] for m in fold_metrics])113    avg_pearson = np.mean([m[3] for m in fold_metrics])114    log(f"\n✅ Avg CV → MSE={avg_mse:.4f}  RMSE={avg_rmse:.4f}  R²={avg_r2:.4f}  Pearson r={avg_pearson:.4f}")115 116    # ── Final Model ────────────────────────────────────────────────────────117    progress(0.85, desc="Training final model …")118    log("\n🚀 Training final model on full dataset …")119    final_model = build_model(train_X.shape[1], l1_reg, l2_reg, dropout_rate)120    final_model.fit(train_X, train_Y, epochs=epochs, batch_size=batch_size, verbose=0)121 122    # ── Parent GEBVs ────────────────────────────────────────────────────────123    parent_GEBV_values = final_model.predict(parent_X, verbose=0).flatten()124    parent_GEBV_dict   = dict(zip(parent_index, parent_GEBV_values))125 126    # ── Hybrid Predictions ─────────────────────────────────────────────────127    hybrid_geno, F1_names, SCA_list = [], [], []128    Parent1_GEBV_list, Parent2_GEBV_list = [], []129 130    for i in range(len(parent_X) - 1):131        for j in range(i + 1, len(parent_X)):132            hybrid = (parent_X[i] + parent_X[j]) / 2133            hybrid_geno.append(hybrid)134            F1_names.append(f"{parent_index[i]} : {parent_index[j]}")135            P1 = parent_GEBV_dict[parent_index[i]]136            P2 = parent_GEBV_dict[parent_index[j]]137            Parent1_GEBV_list.append(P1)138            Parent2_GEBV_list.append(P2)139            SCA_list.append((P1 + P2) / 2)140 141    hybrid_geno = np.array(hybrid_geno)142    GEBV_pred   = final_model.predict(hybrid_geno, verbose=0).flatten()143 144    parent_mean  = np.mean(train_Y)145    better_parent = np.max(train_Y)146    MPH = ((GEBV_pred - parent_mean)   / parent_mean)   * 100147    BPH = ((GEBV_pred - better_parent) / better_parent) * 100148 149    hybrid_results = pd.DataFrame({150        "Hybrid":        F1_names,151        "GEBV":          GEBV_pred,152        "MPH":           MPH,153        "BPH":           BPH,154        "Parent1_GEBV":  Parent1_GEBV_list,155        "Parent2_GEBV":  Parent2_GEBV_list156    }).sort_values(by="GEBV", ascending=False)157 158    F1_out = hybrid_results.head(num_F1)159 160    parent_GEBV_df = (161        hybrid_results162        .assign(Parent1=[x.split(":")[0].strip() for x in hybrid_results["Hybrid"]],163                Parent2=[x.split(":")[1].strip() for x in hybrid_results["Hybrid"]])164        .melt(id_vars=["GEBV"], value_vars=["Parent1","Parent2"],165              var_name="Type", value_name="Parent")166        .groupby("Parent")167        .agg(GEBV=("GEBV","mean"))168        .reset_index()169        .sort_values(by="GEBV", ascending=False)170    )171    P_out = parent_GEBV_df.head(num_P)172 173    # ── Optional test predictions ──────────────────────────────────────────174    test_out = None175    if test_geno_file is not None:176        test_df  = pd.read_csv(test_geno_file.name, index_col=0)177        test_X   = scaler.transform(test_df.to_numpy(dtype=np.float32))178        test_GEBV = final_model.predict(test_X, verbose=0).flatten()179        test_out  = pd.DataFrame({"Test Genotype": test_df.index, "GEBV": test_GEBV})180 181    # ── Plot ───────────────────────────────────────────────────────────────182    progress(0.97, desc="Generating plots …")183    fig, axes = plt.subplots(1, 2, figsize=(14, 5))184    fig.patch.set_facecolor('#0f1117')185 186    # Panel 1 – Top F1 GEBV bar chart187    top_data = hybrid_results.head(min(20, len(hybrid_results)))188    colors   = plt.cm.viridis(np.linspace(0.3, 0.9, len(top_data)))189    axes[0].barh(range(len(top_data)), top_data["GEBV"].values, color=colors)190    axes[0].set_yticks(range(len(top_data)))191    axes[0].set_yticklabels(top_data["Hybrid"].values, fontsize=7, color='white')192    axes[0].set_xlabel("Predicted GEBV", color='white')193    axes[0].set_title("Top Hybrid Combinations", color='white', fontweight='bold')194    axes[0].set_facecolor('#1a1d2e')195    axes[0].tick_params(colors='white')196    for spine in axes[0].spines.values():197        spine.set_edgecolor('#444')198 199    # Panel 2 – GEBV distribution200    axes[1].hist(hybrid_results["GEBV"].values, bins=30,201                 color='#7c5cbf', edgecolor='#b388ff', alpha=0.85)202    axes[1].set_xlabel("Predicted GEBV", color='white')203    axes[1].set_ylabel("Count",          color='white')204    axes[1].set_title("GEBV Distribution Across All Hybrids", color='white', fontweight='bold')205    axes[1].set_facecolor('#1a1d2e')206    axes[1].tick_params(colors='white')207    for spine in axes[1].spines.values():208        spine.set_edgecolor('#444')209 210    cv_text = (f"CV Results  |  MSE: {avg_mse:.4f}  "211               f"RMSE: {avg_rmse:.4f}  R²: {avg_r2:.4f}  Pearson r: {avg_pearson:.4f}")212    fig.suptitle(cv_text, color='#aab4ff', fontsize=9, y=1.01)213    plt.tight_layout()214 215    tmp_plot = tempfile.NamedTemporaryFile(delete=False, suffix=".png")216    plt.savefig(tmp_plot.name, bbox_inches='tight',217                facecolor=fig.get_facecolor(), dpi=150)218    plt.close(fig)219 220    # ── Save CSVs to temp files ────────────────────────────────────────────221    def to_tmp_csv(df, name):222        path = os.path.join(tempfile.gettempdir(), name)223        df.to_csv(path, index=False)224        return path225 226    f1_csv     = to_tmp_csv(F1_out,          "Top_F1_Hybrids.csv")227    p_csv      = to_tmp_csv(P_out,           "Top_Parents.csv")228    all_csv    = to_tmp_csv(hybrid_results,  "All_Hybrid_Results.csv")229    test_csv   = to_tmp_csv(test_out, "Test_GEBV.csv") if test_out is not None else None230 231    log(f"\n✅ Done!  {len(hybrid_results)} hybrid combinations evaluated.")232    log(f"   Top {num_F1} F1 hybrids and Top {num_P} parents identified.")233 234    return (235        F1_out.round(4),236        P_out.round(4),237        test_out.round(4) if test_out is not None else pd.DataFrame({"Info": ["No test file provided"]}),238        tmp_plot.name,239        [f for f in [f1_csv, p_csv, all_csv, test_csv] if f],240        "\n".join(logs)241    )242 243 244# ── Gradio UI ─────────────────────────────────────────────────────────────────245css = """246body, .gradio-container {247    background: #0b0d18 !important;248    font-family: 'Courier New', monospace;249}250h1 { 251    background: linear-gradient(90deg, #7c5cbf, #4fa3e0);252    -webkit-background-clip: text;253    -webkit-text-fill-color: transparent;254    font-size: 2rem;255    font-weight: 900;256    letter-spacing: -0.5px;257}258.gr-button-primary {259    background: linear-gradient(135deg, #7c5cbf, #4fa3e0) !important;260    border: none !important;261    font-weight: 700 !important;262    letter-spacing: 0.5px !important;263}264.gr-box, .gr-panel { background: #13162a !important; border: 1px solid #2a2d4a !important; }265label { color: #aab4ff !important; font-size: 0.82rem !important; }266.gr-input, .gr-dropdown { background: #1a1d2e !important; color: #e0e6ff !important; border-color: #3a3d5a !important; }267"""268 269with gr.Blocks(css=css, title="Heterosis Prediction") as demo:270 271    gr.Markdown("""272# 🌿 Heterosis Prediction273### Deep Learning–based GEBV Prediction & Hybrid Screening274Upload your genotype/phenotype files, configure hyperparameters, and discover top-performing F1 hybrids.275""")276 277    with gr.Row():278        with gr.Column(scale=1):279            gr.Markdown("### 📂 Input Files")280            train_pheno_file = gr.File(label="Training Phenotype CSV  (col 0 = ID, col 1 = trait value)")281            train_geno_file  = gr.File(label="Training Genotype CSV   (row index = sample ID)")282            parent_geno_file = gr.File(label="Parent Genotype CSV     (row index = parent ID)")283            test_geno_file   = gr.File(label="Test Genotype CSV       (optional)")284 285            gr.Markdown("### ⚙️ Parameters")286            with gr.Row():287                num_F1   = gr.Slider(1, 50,  value=5,  step=1,  label="Top F1 hybrids to report")288                num_P    = gr.Slider(1, 30,  value=5,  step=1,  label="Top parents to report")289            with gr.Row():290                k_folds  = gr.Slider(2, 10,  value=5,  step=1,  label="K-Fold CV splits")291                epochs   = gr.Slider(10, 500, value=100, step=10, label="Epochs")292            with gr.Row():293                batch_size   = gr.Slider(8, 128, value=16, step=8,    label="Batch size")294                dropout_rate = gr.Slider(0.0, 0.5, value=0.2, step=0.05, label="Dropout rate")295            with gr.Row():296                l1_reg = gr.Number(value=0.001, label="L1 regularization")297                l2_reg = gr.Number(value=0.001, label="L2 regularization")298 299            run_btn = gr.Button("🚀  Run Prediction", variant="primary")300 301        with gr.Column(scale=2):302            gr.Markdown("### 📊 Results")303            with gr.Tabs():304                with gr.Tab("Top F1 Hybrids"):305                    f1_table = gr.DataFrame(label="Top F1 Hybrid Combinations")306                with gr.Tab("Top Parents"):307                    p_table  = gr.DataFrame(label="Top Parent Lines by Mean GEBV")308                with gr.Tab("Test Predictions"):309                    test_table = gr.DataFrame(label="Test Genotype GEBV Predictions")310                with gr.Tab("Plots"):311                    plot_out = gr.Image(label="Hybrid Performance Visualisation", type="filepath")312                with gr.Tab("Download CSVs"):313                    csv_out  = gr.Files(label="Download result files")314 315            gr.Markdown("### 📋 Training Log")316            log_box = gr.Textbox(label="", lines=12, max_lines=20,317                                 placeholder="Logs will appear here after training …")318 319    gr.Markdown("""320---321**Model architecture:** 8-layer DNN (512→512→256→128→64→32→16→8) with BatchNorm, LeakyReLU, Dropout & L1/L2 regularisation.  322**Metrics:** MSE · RMSE · R² · Pearson r evaluated via K-Fold CV.  323**Heterosis indices computed:** Mid-Parent Heterosis (MPH) · Better-Parent Heterosis (BPH).324""")325 326    run_btn.click(327        fn=heterosis_predict,328        inputs=[329            num_F1, num_P,330            train_pheno_file, train_geno_file, parent_geno_file, test_geno_file,331            l1_reg, l2_reg, dropout_rate,332            k_folds, epochs, batch_size333        ],334        outputs=[f1_table, p_table, test_table, plot_out, csv_out, log_box]335    )336 337if __name__ == "__main__":338    demo.launch()