| """ |
| HABIT — Interactive Soil Water Retention Predictor |
| HuggingFace Space (Gradio) |
| |
| Downloads ensemble weights from huggingface.co/Teamrat/habit on startup, |
| then predicts water retention curves from user-supplied soil properties. |
| """ |
|
|
| import os |
|
|
| os.environ["TF_CPP_MIN_LOG_LEVEL"] = "3" |
| os.environ["KERAS_BACKEND"] = "tensorflow" |
|
|
| import time |
| import shutil |
| import numpy as np |
| import pandas as pd |
| import matplotlib |
|
|
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
| import gradio as gr |
| from huggingface_hub import hf_hub_download |
| import tensorflow as tf |
|
|
| |
| from habit_model import HABIT |
|
|
|
|
| |
| |
| |
|
|
| HF_REPO_ID = "Teamrat/habit" |
| NUM_MEMBERS = 20 |
|
|
| MODEL_CONFIG = { |
| "embedding_dim": 192, |
| "num_heads": 4, |
| "num_monotonic_basis": 40, |
| "dropout_rate": 0.15, |
| } |
|
|
| SCALER_PARAMS = { |
| "texture": {"center": [0.2712, 0.413, 0.172], "scale": [0.456, 0.4543, 0.183]}, |
| "bd": {"center": [1.4], "scale": [0.31]}, |
| "oc": {"center": [1.28], "scale": [1.9902]}, |
| "ksat": {"center": [2.1206], "scale": [1.5133]}, |
| } |
|
|
| DEFAULT_WP_KPA = np.array( |
| [1, 3, 6, 10, 33, 100, 300, 500, 1000, 5000, 10000, 15000], dtype=np.float64 |
| ) |
|
|
|
|
| |
| |
| |
|
|
|
|
| def download_and_load_ensemble(): |
| """Download weights from HF Hub and load all ensemble members.""" |
| cache_dir = os.path.join(os.path.expanduser("~"), ".cache", "habit-ptf", "weights") |
| os.makedirs(cache_dir, exist_ok=True) |
|
|
| models = [] |
| dummy = [ |
| np.zeros((1, 3), dtype=np.float32), |
| np.zeros((1, 1), dtype=np.float32), |
| np.zeros((1, 1), dtype=np.float32), |
| np.zeros((1, 1), dtype=np.float32), |
| np.ones((1, 4), dtype=np.float32), |
| np.zeros((1, 12), dtype=np.float32), |
| ] |
|
|
| for i in range(1, NUM_MEMBERS + 1): |
| name = f"member_{i:02d}.h5" |
| print(f"Loading {name}...", flush=True) |
|
|
| local_path = os.path.join(cache_dir, name) |
| if not os.path.exists(local_path): |
| downloaded = hf_hub_download( |
| repo_id=HF_REPO_ID, |
| filename=f"weights/{name}", |
| ) |
| |
| shutil.copy2(downloaded, local_path) |
|
|
| model = HABIT(**MODEL_CONFIG) |
| model(dummy, training=False) |
| model.load_weights(local_path) |
| models.append(model) |
|
|
| print(f"Loaded {len(models)}-member ensemble.", flush=True) |
| return models |
|
|
|
|
| print("Starting HABIT ensemble loading...") |
| t0 = time.time() |
| ENSEMBLE = download_and_load_ensemble() |
| print(f"Ensemble ready in {time.time() - t0:.1f}s") |
|
|
|
|
| |
| |
| |
|
|
|
|
| def robust_scale(values, center, scale): |
| return ((values - np.array(center)) / np.array(scale)).astype(np.float32) |
|
|
|
|
| def predict_retention(sand, silt, clay, bd, oc, ksat, wp_min, wp_max, n_points): |
| """Run ensemble prediction and return plot + table + CSV path.""" |
| |
| if sand is None or silt is None or clay is None: |
| return None, None, "Sand, silt, and clay are required." |
| sand_f, silt_f, clay_f = float(sand), float(silt), float(clay) |
| if sand_f + silt_f + clay_f < 1: |
| return None, None, "Texture fractions must sum to ~100% (or ~1.0)." |
|
|
| |
| if sand_f + silt_f + clay_f > 5: |
| sand_f, silt_f, clay_f = sand_f / 100, silt_f / 100, clay_f / 100 |
| total = sand_f + silt_f + clay_f |
| sand_f, silt_f, clay_f = sand_f / total, silt_f / total, clay_f / total |
|
|
| texture_sc = robust_scale( |
| np.array([[sand_f, silt_f, clay_f]]), |
| SCALER_PARAMS["texture"]["center"], |
| SCALER_PARAMS["texture"]["scale"], |
| ) |
|
|
| |
| mask = np.zeros((1, 4), dtype=np.float32) |
| mask[0, 0] = 1.0 |
|
|
| if bd is not None and bd > 0: |
| bd_sc = robust_scale( |
| np.array([[float(bd)]]), |
| SCALER_PARAMS["bd"]["center"], |
| SCALER_PARAMS["bd"]["scale"], |
| ) |
| mask[0, 1] = 1.0 |
| else: |
| bd_sc = np.zeros((1, 1), dtype=np.float32) |
|
|
| if oc is not None and oc > 0: |
| oc_val = float(oc) |
| if oc_val > 1.0: |
| oc_val /= 100 |
| oc_log = np.log1p(oc_val) |
| oc_sc = robust_scale( |
| np.array([[oc_log]]), |
| SCALER_PARAMS["oc"]["center"], |
| SCALER_PARAMS["oc"]["scale"], |
| ) |
| mask[0, 2] = 1.0 |
| else: |
| oc_sc = np.zeros((1, 1), dtype=np.float32) |
|
|
| if ksat is not None and ksat > 0: |
| ksat_log = np.log10(max(float(ksat), 1e-6)) |
| ksat_sc = robust_scale( |
| np.array([[ksat_log]]), |
| SCALER_PARAMS["ksat"]["center"], |
| SCALER_PARAMS["ksat"]["scale"], |
| ) |
| mask[0, 3] = 1.0 |
| else: |
| ksat_sc = np.zeros((1, 1), dtype=np.float32) |
|
|
| |
| stage_names = { |
| (1, 0, 0, 0): "Stage 0 — texture only", |
| (1, 1, 0, 0): "Stage 1 — texture + BD", |
| (1, 1, 1, 0): "Stage 2 — texture + BD + OC", |
| (1, 1, 1, 1): "Stage 3 — all properties", |
| } |
| mask_key = tuple(int(m) for m in mask[0]) |
| stage_label = stage_names.get(mask_key, f"Custom mask: {mask_key}") |
|
|
| |
| n_pts = int(n_points) if n_points else 50 |
| wp_kpa = np.logspace( |
| np.log10(max(float(wp_min), 0.1)), np.log10(float(wp_max)), n_pts |
| ) |
| wp_log = np.log10(wp_kpa).astype(np.float32).reshape(1, -1) |
|
|
| |
| inputs = [texture_sc, bd_sc, oc_sc, ksat_sc, mask, wp_log] |
| all_preds = [] |
| for model in ENSEMBLE: |
| pred = model(inputs, training=False).numpy() |
| all_preds.append(pred[0]) |
|
|
| all_preds = np.array(all_preds) |
| mean = np.mean(all_preds, axis=0) |
| std = np.std(all_preds, axis=0) |
| lower = np.percentile(all_preds, 2.5, axis=0) |
| upper = np.percentile(all_preds, 97.5, axis=0) |
|
|
| |
| fig, ax = plt.subplots(figsize=(8, 5)) |
| ax.fill_between( |
| wp_kpa, lower, upper, alpha=0.25, color="#2196F3", label="95% interval" |
| ) |
| ax.plot(wp_kpa, mean, color="#1565C0", linewidth=2, label="Ensemble mean") |
| for m in range(len(all_preds)): |
| ax.plot(wp_kpa, all_preds[m], color="#90CAF9", linewidth=0.4, alpha=0.6) |
| ax.set_xscale("log") |
| ax.set_xlabel("Water potential |ψ| (kPa)", fontsize=12) |
| ax.set_ylabel("Volumetric water content θ (cm³/cm³)", fontsize=12) |
| ax.set_title(f"HABIT Prediction — {stage_label}", fontsize=13) |
| ax.legend(loc="upper right") |
| ax.set_ylim(bottom=0) |
| ax.grid(True, alpha=0.3) |
| fig.tight_layout() |
|
|
| |
| standard_kpa = [1, 3, 6, 10, 33, 100, 300, 500, 1000, 5000, 10000, 15000] |
| standard_kpa = [p for p in standard_kpa if float(wp_min) <= p <= float(wp_max)] |
|
|
| table_rows = [] |
| for target_kpa in standard_kpa: |
| idx = np.argmin(np.abs(wp_kpa - target_kpa)) |
| table_rows.append( |
| { |
| "ψ (kPa)": int(target_kpa), |
| "θ mean": f"{mean[idx]:.4f}", |
| "θ std": f"{std[idx]:.4f}", |
| "θ lower 95%": f"{lower[idx]:.4f}", |
| "θ upper 95%": f"{upper[idx]:.4f}", |
| } |
| ) |
| table_df = pd.DataFrame(table_rows) |
|
|
| |
| full_df = pd.DataFrame( |
| { |
| "water_potential_kPa": wp_kpa, |
| "water_content_mean": mean, |
| "water_content_std": std, |
| "water_content_lower95": lower, |
| "water_content_upper95": upper, |
| } |
| ) |
| for m in range(len(all_preds)): |
| full_df[f"member_{m + 1:02d}"] = all_preds[m] |
|
|
| csv_path = "/tmp/habit_prediction.csv" |
| full_df.to_csv(csv_path, index=False) |
|
|
| return fig, table_df, csv_path |
|
|
|
|
| |
| |
| |
|
|
|
|
| def predict_from_csv(file): |
| """Run predictions for all soils in an uploaded CSV.""" |
| if file is None: |
| return None, None, "Please upload a CSV file." |
|
|
| df = pd.read_csv(file.name if hasattr(file, "name") else file) |
|
|
| cols_lower = {c.lower(): c for c in df.columns} |
| if not all(k in cols_lower for k in ["sand", "silt", "clay"]): |
| return None, None, f"CSV must have sand, silt, clay columns. Found: {list(df.columns)}" |
|
|
| results_all = [] |
| for idx, row in df.iterrows(): |
| sand = row[cols_lower["sand"]] |
| silt = row[cols_lower["silt"]] |
| clay = row[cols_lower["clay"]] |
| bd = row.get(cols_lower.get("bd")) if "bd" in cols_lower else None |
| oc = row.get(cols_lower.get("oc")) if "oc" in cols_lower else None |
| ksat = row.get(cols_lower.get("ksat")) if "ksat" in cols_lower else None |
| soil_id = row.get(cols_lower.get("soil_id", ""), idx + 1) |
|
|
| bd = None if bd is not None and (pd.isna(bd) or bd <= 0) else bd |
| oc = None if oc is not None and (pd.isna(oc) or oc <= 0) else oc |
| ksat = None if ksat is not None and (pd.isna(ksat) or ksat <= 0) else ksat |
|
|
| fig, table, csv_path = predict_retention( |
| sand, silt, clay, bd, oc, ksat, 1, 15000, 50 |
| ) |
| plt.close(fig) |
|
|
| pred = pd.read_csv(csv_path) |
| pred.insert(0, "soil_id", soil_id) |
| results_all.append(pred) |
|
|
| combined = pd.concat(results_all, ignore_index=True) |
| out_path = "/tmp/habit_batch_predictions.csv" |
| combined.to_csv(out_path, index=False) |
|
|
| summary = ( |
| combined.groupby("soil_id") |
| .agg( |
| n_points=("water_content_mean", "count"), |
| theta_sat=("water_content_mean", "max"), |
| theta_15000=("water_content_mean", "min"), |
| ) |
| .reset_index() |
| ) |
|
|
| return summary, out_path, f"Predicted {len(df)} soils successfully." |
|
|
|
|
| |
| |
| |
|
|
| EXAMPLE_SOILS = { |
| "Clay (heavy)": {"sand": 10, "silt": 30, "clay": 60, "bd": 1.2, "oc": 2.0, "ksat": 5}, |
| "Sandy loam": {"sand": 65, "silt": 25, "clay": 10, "bd": 1.5, "oc": 0.5, "ksat": 200}, |
| "Silt loam": {"sand": 15, "silt": 65, "clay": 20, "bd": 1.3, "oc": 1.5, "ksat": 25}, |
| "Loam (average)": {"sand": 40, "silt": 40, "clay": 20, "bd": 1.35, "oc": 1.2, "ksat": 50}, |
| "Sand (texture only)": {"sand": 90, "silt": 5, "clay": 5, "bd": None, "oc": None, "ksat": None}, |
| } |
|
|
|
|
| def load_preset(name): |
| if name and name in EXAMPLE_SOILS: |
| s = EXAMPLE_SOILS[name] |
| return ( |
| s["sand"], |
| s["silt"], |
| s["clay"], |
| s["bd"] if s["bd"] else 0, |
| s["oc"] if s["oc"] else 0, |
| s["ksat"] if s["ksat"] else 0, |
| ) |
| return gr.update(), gr.update(), gr.update(), gr.update(), gr.update(), gr.update() |
|
|
|
|
| with gr.Blocks( |
| title="HABIT — Soil Water Retention Predictor", |
| theme=gr.themes.Soft(), |
| ) as demo: |
|
|
| gr.Markdown( |
| """ |
| # HABIT — Soil Water Retention Predictor |
| **Hierarchical Attention-Based Inference with Transfer Learning** |
| |
| Predict soil water retention curves from basic soil properties using a 20-member |
| deep learning ensemble. Provide whatever properties you have — the model adapts automatically. |
| |
| *Ghezzehei TA (2025). Water Resources Research.* |
| | [Model weights](https://huggingface.co/Teamrat/habit) |
| | [pip install habit-ptf](https://pypi.org/project/habit-ptf/) |
| """ |
| ) |
|
|
| with gr.Tabs(): |
| |
| with gr.TabItem("Single Soil"): |
| with gr.Row(): |
| with gr.Column(scale=1): |
| gr.Markdown("### Soil Properties") |
|
|
| preset = gr.Dropdown( |
| choices=list(EXAMPLE_SOILS.keys()), |
| label="Load example soil", |
| interactive=True, |
| ) |
|
|
| gr.Markdown("**Texture** (required — % or fraction)") |
| with gr.Row(): |
| sand_in = gr.Number(label="Sand", value=40, precision=1) |
| silt_in = gr.Number(label="Silt", value=40, precision=1) |
| clay_in = gr.Number(label="Clay", value=20, precision=1) |
|
|
| gr.Markdown("**Optional properties** (leave at 0 to omit)") |
| bd_in = gr.Number( |
| label="Bulk density (g/cm³)", value=1.35, precision=2 |
| ) |
| oc_in = gr.Number( |
| label="Organic carbon (%)", value=1.2, precision=2 |
| ) |
| ksat_in = gr.Number( |
| label="Ksat (cm/day)", value=50, precision=1 |
| ) |
|
|
| gr.Markdown("**Water potential range**") |
| with gr.Row(): |
| wp_min_in = gr.Number(label="Min (kPa)", value=1, precision=0) |
| wp_max_in = gr.Number( |
| label="Max (kPa)", value=15000, precision=0 |
| ) |
| n_pts_in = gr.Number(label="Points", value=50, precision=0) |
|
|
| predict_btn = gr.Button("Predict", variant="primary", size="lg") |
|
|
| with gr.Column(scale=2): |
| plot_out = gr.Plot(label="Water Retention Curve") |
| table_out = gr.Dataframe(label="Predictions at Standard Tensions") |
| csv_out = gr.File(label="Download Full Results (CSV)") |
|
|
| preset.change( |
| fn=load_preset, |
| inputs=[preset], |
| outputs=[sand_in, silt_in, clay_in, bd_in, oc_in, ksat_in], |
| ) |
|
|
| predict_btn.click( |
| fn=predict_retention, |
| inputs=[ |
| sand_in, |
| silt_in, |
| clay_in, |
| bd_in, |
| oc_in, |
| ksat_in, |
| wp_min_in, |
| wp_max_in, |
| n_pts_in, |
| ], |
| outputs=[plot_out, table_out, csv_out], |
| ) |
|
|
| |
| with gr.TabItem("Batch (CSV Upload)"): |
| gr.Markdown( |
| """ |
| ### Batch Prediction |
| |
| Upload a CSV with columns: `sand`, `silt`, `clay` (required), |
| plus optional `bd`, `oc`, `ksat`, `soil_id`. |
| |
| Values can be percentages (0–100) or fractions (0–1). Missing optional |
| properties should be blank or 0. |
| """ |
| ) |
|
|
| csv_upload = gr.File(label="Upload CSV", file_types=[".csv"]) |
| batch_btn = gr.Button("Predict All", variant="primary") |
| batch_status = gr.Textbox(label="Status") |
| batch_summary = gr.Dataframe(label="Summary") |
| batch_download = gr.File(label="Download Results") |
|
|
| batch_btn.click( |
| fn=predict_from_csv, |
| inputs=[csv_upload], |
| outputs=[batch_summary, batch_download, batch_status], |
| ) |
|
|
| |
| with gr.TabItem("About"): |
| gr.Markdown( |
| """ |
| ### About HABIT |
| |
| HABIT is a deep learning model for predicting soil water retention curves |
| from basic soil properties. It uses a transformer-based architecture with: |
| |
| - **Property-specific encoders** for each soil property |
| - **Cross-attention layers** that learn interactions between properties |
| - **Monotonic output layer** ensuring physically correct behavior |
| (water content decreases with increasing tension) |
| - **Hierarchical training** so one model handles any combination of inputs |
| |
| #### Performance (test set, 95% CI from cluster bootstrap) |
| |
| | Inputs | R² | RMSE (cm³/cm³) | |
| |---|---|---| |
| | Texture only | 0.78 [0.74, 0.82] | 0.067 | |
| | + Bulk density | 0.85 [0.75, 0.91] | 0.056 | |
| | + Organic carbon | 0.86 [0.78, 0.92] | 0.052 | |
| | + Ksat | 0.92 [0.90, 0.94] | 0.043 | |
| |
| #### Python package |
| |
| ```bash |
| pip install habit-ptf |
| ``` |
| |
| ```python |
| from habit_ptf import load_ensemble |
| |
| predictor = load_ensemble() |
| result = predictor.predict(soil_dataframe) |
| ``` |
| |
| #### Citation |
| |
| Ghezzehei TA (2025). Interpretable Soil Water Retention Prediction Using |
| Hierarchical Attention Networks with Uncertainty Quantification. |
| *Water Resources Research*. |
| |
| #### License |
| |
| MIT (code and weights). Training data: CC BY 4.0. |
| """ |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| demo.launch() |
|
|