softwareDevelopment/Heterosis
0
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()